diff --git a/middleware/ratelimit.go b/middleware/ratelimit.go new file mode 100644 index 00000000..bcd824fd --- /dev/null +++ b/middleware/ratelimit.go @@ -0,0 +1,51 @@ +package middleware + +import ( + "strconv" + "time" + + "github.com/cloudreve/Cloudreve/v4/application/dependency" + "github.com/cloudreve/Cloudreve/v4/pkg/serializer" + "github.com/gin-gonic/gin" +) + +const rateLimitPrefix = "rate_limit:" + +// RateLimit applies a fixed-window rate limit to requests sharing the same +// bucket key. Once `limit` requests arrive within `window`, further requests +// are rejected until the window expires. State lives in the shared KV store, +// so limits hold across cluster nodes when Redis is configured. +// +// The counter is approximate (read-then-write), which is acceptable for +// abuse prevention — it bounds attempts to roughly `limit` per window. +func RateLimit(limit int, window time.Duration, key func(*gin.Context) string) gin.HandlerFunc { + return func(c *gin.Context) { + kv := dependency.FromContext(c).KV() + bucket := rateLimitPrefix + key(c) + + count := 0 + if raw, ok := kv.Get(bucket); ok { + if v, ok := raw.(int); ok { + count = v + } + } + + if count >= limit { + c.Header("Retry-After", strconv.Itoa(int(window.Seconds()))) + c.JSON(200, serializer.NewError(serializer.CodeRateLimited, "Too many requests, please try again later.", nil)) + c.Abort() + return + } + + _ = kv.Set(bucket, count+1, int(window.Seconds())) + c.Next() + } +} + +// RateLimitByIP keys the rate limit on the client IP and a static bucket +// name, e.g. RateLimitByIP("login", 10, time.Minute). +func RateLimitByIP(bucket string, limit int, window time.Duration) gin.HandlerFunc { + return RateLimit(limit, window, func(c *gin.Context) string { + return bucket + ":" + c.ClientIP() + }) +} diff --git a/middleware/ratelimit_test.go b/middleware/ratelimit_test.go new file mode 100644 index 00000000..1b9a72dc --- /dev/null +++ b/middleware/ratelimit_test.go @@ -0,0 +1,123 @@ +package middleware + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/cloudreve/Cloudreve/v4/application/dependency" + "github.com/cloudreve/Cloudreve/v4/pkg/cache" + "github.com/gin-gonic/gin" +) + +var testEngine = func() *gin.Engine { + e := gin.New() + e.ContextWithFallback = true + return e +}() + +func newRateLimitContext(t *testing.T, dep dependency.Dep, remoteAddr string) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + w := httptest.NewRecorder() + c := gin.CreateTestContextOnly(w, testEngine) + req := httptest.NewRequest(http.MethodPost, "/session/token", nil) + req.RemoteAddr = remoteAddr + ":12345" + c.Request = req.WithContext(context.WithValue(req.Context(), dependency.DepCtx{}, dep)) + return c, w +} + +func TestRateLimitByIP(t *testing.T) { + gin.SetMode(gin.TestMode) + dep := dependency.NewDependency(dependency.WithKV(cache.NewMemoStore("", nil))) + + handler := RateLimitByIP("login", 3, time.Minute) + + for i := 0; i < 3; i++ { + c, _ := newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if c.IsAborted() { + t.Fatalf("request %d unexpectedly rejected", i+1) + } + } + + c, w := newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if !c.IsAborted() { + t.Fatal("request over limit was not rejected") + } + if w.Header().Get("Retry-After") != "60" { + t.Fatalf("missing/incorrect Retry-After: %q", w.Header().Get("Retry-After")) + } + + // A different IP has its own bucket. + c, _ = newRateLimitContext(t, dep, "192.0.2.2") + handler(c) + if c.IsAborted() { + t.Fatal("unrelated IP shared the bucket") + } + + // A different bucket name on the same IP has its own counter. + handler2 := RateLimitByIP("register", 1, time.Minute) + c, _ = newRateLimitContext(t, dep, "192.0.2.1") + handler2(c) + if c.IsAborted() { + t.Fatal("unrelated bucket shared the counter") + } +} + +func TestRateLimitExpiry(t *testing.T) { + gin.SetMode(gin.TestMode) + dep := dependency.NewDependency(dependency.WithKV(cache.NewMemoStore("", nil))) + + handler := RateLimitByIP("login", 1, time.Second) + + c, _ := newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if c.IsAborted() { + t.Fatal("first request rejected") + } + + c, _ = newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if !c.IsAborted() { + t.Fatal("second request not rejected") + } + + // Window expiry frees the bucket. MemoStore TTL is Unix-second granular + // (item valid while Expires >= now), so a 1s window can live ~2s. + time.Sleep(2100 * time.Millisecond) + + c, _ = newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if c.IsAborted() { + t.Fatal("request after window expiry still rejected") + } +} + +func TestRateLimitCountOverflow(t *testing.T) { + gin.SetMode(gin.TestMode) + kv := cache.NewMemoStore("", nil) + dep := dependency.NewDependency(dependency.WithKV(kv)) + + // A non-int value in the bucket is treated as a fresh window. + if err := kv.Set(rateLimitPrefix+"login:192.0.2.1", "garbage", 60); err != nil { + t.Fatal(err) + } + + handler := RateLimitByIP("login", 1, time.Minute) + c, _ := newRateLimitContext(t, dep, "192.0.2.1") + handler(c) + if c.IsAborted() { + t.Fatal("corrupt bucket state rejected request") + } + + raw, ok := kv.Get(rateLimitPrefix + "login:192.0.2.1") + if !ok { + t.Fatal("bucket not persisted") + } + if _, isInt := raw.(int); !isInt { + t.Fatalf("bucket value type drifted: %T", raw) + } +} diff --git a/pkg/request/ssrftest/ssrf_test.go b/pkg/request/ssrftest/ssrf_test.go index 995f5dd4..c23957c7 100644 --- a/pkg/request/ssrftest/ssrf_test.go +++ b/pkg/request/ssrftest/ssrf_test.go @@ -59,6 +59,16 @@ func TestValidateExternalURL_IPLiterals(t *testing.T) { "http://100.64.0.1/cgnat", "http://[::ffff:127.0.0.1]/v4mapped", "http://[::ffff:10.0.0.1]/v4mappedpriv", + // IPv4-in-IPv6 transition forms wrapping internal targets. + "http://[64:ff9b::127.0.0.1]/nat64-loopback", + "http://[64:ff9b::a9fe:a9fe]/nat64-metadata", + "http://[64:ff9b::169.254.169.254]/nat64-metadata-dec", + "http://[64:ff9b::a00:1]/nat64-private", + "http://[2002:7f00:0001::]/6to4-loopback", + "http://[2002:a9fe:a9fe::]/6to4-metadata", + "http://[2002:0a00:0001::]/6to4-private", + "http://[2001:0000:4136:e378:8000:63bf:f5ff:fffe]/teredo-10.0.0.1", + "http://[::127.0.0.1]/v4compatible", "http://224.0.0.1/multicast", } for _, raw := range cases { @@ -73,6 +83,9 @@ func TestValidateExternalURL_PublicIP(t *testing.T) { "http://1.1.1.1/", "https://8.8.8.8/", "http://[2606:4700:4700::1111]/", + // Transition forms wrapping *public* IPv4 remain allowed. + "http://[64:ff9b::808:808]/nat64-public", + "http://[2002:0808:0808::]/6to4-public", } for _, raw := range cases { err := request.ValidateExternalURL(ctx, raw, request.SSRFOptions{}) diff --git a/pkg/serializer/error.go b/pkg/serializer/error.go index 28666d33..bb030a43 100644 --- a/pkg/serializer/error.go +++ b/pkg/serializer/error.go @@ -257,6 +257,8 @@ const ( CodeAnonymouseAccessDenied = 40088 // CodeInsufficientScope OAuth token scope insufficient CodeInsufficientScope = 40089 + // CodeRateLimited 请求频率超限 + CodeRateLimited = 40090 // CodeDBError 数据库操作失败 CodeDBError = 50001 // CodeEncryptError 加密失败 diff --git a/pkg/util/common.go b/pkg/util/common.go index 1b5aa0c4..3791d7de 100644 --- a/pkg/util/common.go +++ b/pkg/util/common.go @@ -37,24 +37,24 @@ func RandStringRunes(n int) string { } func RandStringRunesCrypto(n int) string { - b := make([]rune, n) - for i := range b { - num, err := cryptoRand.Int(cryptoRand.Reader, big.NewInt(int64(len(RandomVariantAll)))) - if err != nil { - // fallback to math/rand on crypto failure - b[i] = RandomVariantAll[rand.Intn(len(RandomVariantAll))] - } else { - b[i] = RandomVariantAll[num.Int64()] - } - } - return string(b) + return randStringCrypto(n, RandomVariantAll) } // RandString returns random string in given length and variant func RandString(n int, variant []rune) string { + return randStringCrypto(n, variant) +} + +func randStringCrypto(n int, variant []rune) string { b := make([]rune, n) for i := range b { - b[i] = variant[rand.Intn(len(variant))] + num, err := cryptoRand.Int(cryptoRand.Reader, big.NewInt(int64(len(variant)))) + if err != nil { + // fallback to math/rand on crypto failure + b[i] = variant[rand.Intn(len(variant))] + } else { + b[i] = variant[num.Int64()] + } } return string(b) } diff --git a/routers/router.go b/routers/router.go index 24ad3606..0b124e06 100644 --- a/routers/router.go +++ b/routers/router.go @@ -2,6 +2,7 @@ package routers import ( "net/http" + "time" "github.com/cloudreve/Cloudreve/v4/application/constants" "github.com/cloudreve/Cloudreve/v4/application/dependency" @@ -289,6 +290,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { { // 用户登录 token.POST("", + middleware.RateLimitByIP("login", 10, time.Minute), middleware.CaptchaRequired(func(c *gin.Context) bool { return dep.SettingProvider().LoginCaptchaEnabled(c) }), @@ -298,11 +300,13 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { ) // 2-factor authentication token.POST("2fa", + middleware.RateLimitByIP("login_2fa", 10, time.Minute), controllers.FromJSON[usersvc.OtpValidationService](usersvc.OtpValidationParameterCtx{}), controllers.UserLogin2FAValidation, controllers.UserIssueToken, ) token.POST("refresh", + middleware.RateLimitByIP("token_refresh", 30, time.Minute), middleware.RequiredScopes(types.ScopeOfflineAccess), controllers.FromJSON[usersvc.RefreshTokenService](usersvc.RefreshTokenParameterCtx{}), controllers.UserRefreshToken, @@ -331,6 +335,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { controllers.UserSSOCallback, ) ssoRouter.POST("exchange", + middleware.RateLimitByIP("sso_exchange", 20, time.Minute), controllers.FromJSON[usersvc.SSOExchangeService](usersvc.SSOExchangeParameterCtx{}), controllers.UserSSOExchange, ) @@ -348,6 +353,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { controllers.GrantAppConsent, ) oauthRouter.POST("token", + middleware.RateLimitByIP("oauth_token", 20, time.Minute), controllers.FromForm[oauth.ExchangeTokenService](oauth.ExchangeTokenParamCtx{}), controllers.ExchangeToken, ) @@ -369,6 +375,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { { // WebAuthn login prepare authn.PUT("", + middleware.RateLimitByIP("authn", 20, time.Minute), middleware.IsFunctionEnabled(func(c *gin.Context) bool { return dep.SettingProvider().AuthnEnabled(c) }), @@ -376,6 +383,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { ) // WebAuthn finish login authn.POST("", + middleware.RateLimitByIP("authn", 20, time.Minute), middleware.IsFunctionEnabled(func(c *gin.Context) bool { return dep.SettingProvider().AuthnEnabled(c) }), @@ -391,6 +399,7 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { { // 用户注册 Done user.POST("", + middleware.RateLimitByIP("register", 5, time.Minute), middleware.IsFunctionEnabled(func(c *gin.Context) bool { return dep.SettingProvider().RegisterEnabled(c) }), @@ -402,12 +411,14 @@ func initMasterRouter(dep dependency.Dep) *gin.Engine { ) // 通过邮件里的链接重设密码 user.PATCH("reset/:id", + middleware.RateLimitByIP("reset_apply", 10, time.Minute), middleware.HashID(hashid.UserID), controllers.FromJSON[usersvc.UserResetService](usersvc.UserResetParameterCtx{}), controllers.UserReset, ) // 发送密码重设邮件 user.POST("reset", + middleware.RateLimitByIP("reset_mail", 5, time.Minute), middleware.CaptchaRequired(func(c *gin.Context) bool { return dep.SettingProvider().ForgotPasswordCaptchaEnabled(c) }),