diff --git a/cmd/api/server.go b/cmd/api/server.go index c67c74d..fc6e3ff 100644 --- a/cmd/api/server.go +++ b/cmd/api/server.go @@ -3,12 +3,13 @@ package main import ( "context" "fmt" - ratelimit "github.com/JGLTechnologies/gin-rate-limit" "github.com/Quaver/api2/config" + "github.com/Quaver/api2/db" "github.com/Quaver/api2/handlers" "github.com/Quaver/api2/middleware" "github.com/gin-contrib/cors" "github.com/gin-gonic/gin" + redisrate "github.com/go-redis/redis_rate/v10" "github.com/sirupsen/logrus" "net/http" "os" @@ -18,6 +19,14 @@ import ( "time" ) +const apiRateLimitKeyPrefix = "quaver:api_rate_limit:" + +var apiRequestRateLimit = redisrate.PerMinute(100) + +type apiRateLimiter interface { + Allow(ctx context.Context, key string, limit redisrate.Limit) (*redisrate.Result, error) +} + // Starts the server on a given port func initializeServer(port int) { gin.SetMode(gin.ReleaseMode) @@ -48,18 +57,20 @@ func initializeServer(port int) { // Initializes the rate limiter for the server func initializeRateLimiter(engine *gin.Engine) { + engine.Use(newRateLimitMiddleware(redisrate.NewLimiter(db.Redis))) +} + +func newRateLimitMiddleware(limiter apiRateLimiter) gin.HandlerFunc { rateLimitBypassRoutes := map[string]struct{}{ - "/v2/mapset/search": {}, - "/v2/map/:id": {}, + "/v2/mapset/search": {}, + "/v2/map/:id": {}, + "/v2/download/mapset/:id": {}, } - store := ratelimit.InMemoryStore(&ratelimit.InMemoryOptions{ - Rate: time.Minute, - Limit: 100, - }) + return func(c *gin.Context) { + clientIP := c.ClientIP() - engine.Use(func(c *gin.Context) { - if !config.Instance.IsProduction || slices.Contains(config.Instance.Server.RateLimitIpWhitelist, c.ClientIP()) { + if !config.Instance.IsProduction || slices.Contains(config.Instance.Server.RateLimitIpWhitelist, clientIP) { c.Next() return } @@ -73,19 +84,48 @@ func initializeRateLimiter(engine *gin.Engine) { } } - info := store.Limit(c.ClientIP(), c) - c.Header("X-Rate-Limit-Limit", fmt.Sprintf("%d", info.Limit)) - c.Header("X-Rate-Limit-Remaining", fmt.Sprintf("%v", info.RemainingHits)) - c.Header("X-Rate-Limit-Reset", fmt.Sprintf("%d", info.ResetTime.Unix())) + result, err := limiter.Allow( + c.Request.Context(), + apiRateLimitKeyPrefix+clientIP, + apiRequestRateLimit, + ) + + if err != nil { + // Rate limiting should not make the entire API unavailable when Redis + // is temporarily unreachable. Authentication and individual handlers + // can still apply their own Redis failure policies. + logrus.Errorf("Error checking API rate limit: %v", err) + c.Next() + return + } + + resetAfter := result.ResetAfter + + if resetAfter < 0 { + resetAfter = 0 + } - if info.RateLimited { + c.Header("X-Rate-Limit-Limit", fmt.Sprintf("%d", result.Limit.Rate)) + c.Header("X-Rate-Limit-Remaining", fmt.Sprintf("%d", result.Remaining)) + c.Header("X-Rate-Limit-Reset", fmt.Sprintf("%d", time.Now().Add(resetAfter).Unix())) + + if result.Allowed == 0 { + c.Header("Retry-After", fmt.Sprintf("%d", retryAfterSeconds(result.RetryAfter))) c.JSON(http.StatusTooManyRequests, gin.H{"error": "Too many requests"}) c.Abort() return } c.Next() - }) + } +} + +func retryAfterSeconds(retryAfter time.Duration) int64 { + if retryAfter <= 0 { + return 1 + } + + return int64((retryAfter + time.Second - 1) / time.Second) } // Initializes all the routes for the server. diff --git a/downloadlimit/downloadlimit.go b/downloadlimit/downloadlimit.go new file mode 100644 index 0000000..d229a71 --- /dev/null +++ b/downloadlimit/downloadlimit.go @@ -0,0 +1,60 @@ +package downloadlimit + +import ( + "context" + "fmt" + "sync" + "time" + + "github.com/redis/go-redis/v9" +) + +// DailyLimitBytes is the maximum number of mapset bytes a user can download +// during a single server-local calendar day. +const DailyLimitBytes int64 = 1 << 30 // 1GB max: 1 << 30, easier test with 50MB max: 50 * 1024 * 1024 + +// downloadLimitMu keeps the Redis increment, rollback, and expiry operations from +// interleaving. +var downloadLimitMu sync.Mutex + +// TryConsume adds byteCount to a user's daily download total. It returns false +// and restores the counter when the new total exceeds DailyLimitBytes. +func TryConsume(ctx context.Context, redisClient *redis.Client, userID int, byteCount int64) (bool, error) { + return tryConsumeAt(ctx, redisClient, userID, byteCount, time.Now().In(time.Local)) +} + +func tryConsumeAt( + ctx context.Context, + redisClient *redis.Client, + userID int, + byteCount int64, + now time.Time, +) (bool, error) { + downloadLimitMu.Lock() + defer downloadLimitMu.Unlock() + + key, expiresAt := quotaWindow(userID, now) + used, err := redisClient.IncrBy(ctx, key, byteCount).Result() + + if err != nil { + return false, fmt.Errorf("increment daily download bytes: %w", err) + } + + if used > DailyLimitBytes { + if err := redisClient.DecrBy(ctx, key, byteCount).Err(); err != nil { + return false, fmt.Errorf("restore daily download bytes: %w", err) + } + } + + if err := redisClient.ExpireAt(ctx, key, expiresAt).Err(); err != nil { + return false, fmt.Errorf("expire daily download counter: %w", err) + } + + return used <= DailyLimitBytes, nil +} + +func quotaWindow(userID int, now time.Time) (string, time.Time) { + expiresAt := time.Date(now.Year(), now.Month(), now.Day()+1, 0, 0, 0, 0, now.Location()) + key := fmt.Sprintf("quaver:download_rate_limit:%s:%d", now.Format("2006-01-02"), userID) + return key, expiresAt +} diff --git a/downloadlimit/downloadlimit_test.go b/downloadlimit/downloadlimit_test.go new file mode 100644 index 0000000..b95faa4 --- /dev/null +++ b/downloadlimit/downloadlimit_test.go @@ -0,0 +1,160 @@ +package downloadlimit + +import ( + "context" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/Quaver/api2/config" + "github.com/Quaver/api2/db" + "github.com/redis/go-redis/v9" +) + +var testUserSequence atomic.Int64 + +func TestTryConsumeAccumulatesAndEnforcesLimit(t *testing.T) { + client, userID, now := useTestRedis(t) + ctx := context.Background() + + allowed, err := tryConsumeAt(ctx, client, userID, 128, now) + if err != nil || !allowed { + t.Fatalf("first download: allowed=%v err=%v", allowed, err) + } + + allowed, err = tryConsumeAt(ctx, client, userID, DailyLimitBytes-128, now) + if err != nil || !allowed { + t.Fatalf("download reaching limit: allowed=%v err=%v", allowed, err) + } + + allowed, err = tryConsumeAt(ctx, client, userID, 1, now) + if err != nil { + t.Fatal(err) + } + + if allowed { + t.Fatal("expected download above the limit to be rejected") + } + + key, _ := quotaWindow(userID, now) + assertCounterValue(t, ctx, client, key, DailyLimitBytes) + + if ttl, err := client.TTL(ctx, key).Result(); err != nil || ttl <= 0 { + t.Fatalf("TTL = %v, err=%v", ttl, err) + } +} + +func TestTryConsumeConcurrentRequestsDoNotExceedLimit(t *testing.T) { + client, userID, now := useTestRedis(t) + ctx := context.Background() + const workerCount = 24 + const allowedCount = 8 + chunkSize := DailyLimitBytes / allowedCount + start := make(chan struct{}) + errorsChannel := make(chan error, workerCount) + var successful atomic.Int64 + var waitGroup sync.WaitGroup + + for range workerCount { + waitGroup.Add(1) + + go func() { + defer waitGroup.Done() + <-start + + allowed, err := tryConsumeAt(ctx, client, userID, chunkSize, now) + + if err != nil { + errorsChannel <- err + return + } + + if allowed { + successful.Add(1) + } + }() + } + + close(start) + waitGroup.Wait() + close(errorsChannel) + + for err := range errorsChannel { + t.Fatal(err) + } + + if got := successful.Load(); got != allowedCount { + t.Fatalf("allowed downloads = %d, want %d", got, allowedCount) + } + + key, _ := quotaWindow(userID, now) + assertCounterValue(t, ctx, client, key, DailyLimitBytes) +} + +func TestQuotaWindowEndsAtNextLocalMidnightAcrossDST(t *testing.T) { + location, err := time.LoadLocation("Europe/Sofia") + + if err != nil { + t.Fatal(err) + } + + now := time.Date(2026, time.March, 29, 0, 30, 0, 0, location) + _, expiresAt := quotaWindow(15, now) + expectedTTL := 22*time.Hour + 30*time.Minute + + if expiresAt.Sub(now) != expectedTTL { + t.Fatalf("midnight duration = %v, want %v", expiresAt.Sub(now), expectedTTL) + } +} + +func TestTryConsumeReturnsRedisErrors(t *testing.T) { + client := redis.NewClient(&redis.Options{ + Addr: "127.0.0.1:0", + MaxRetries: -1, + DialTimeout: 10 * time.Millisecond, + ReadTimeout: 10 * time.Millisecond, + WriteTimeout: 10 * time.Millisecond, + }) + t.Cleanup(func() { _ = client.Close() }) + + if _, err := TryConsume(context.Background(), client, 55, 1); err == nil { + t.Fatal("expected Redis error") + } +} + +func useTestRedis(t *testing.T) (*redis.Client, int, time.Time) { + t.Helper() + + if config.Instance == nil { + if err := config.Load("../config.json"); err != nil { + t.Fatal(err) + } + } + + db.InitializeRedis() + if err := db.Redis.Ping(context.Background()).Err(); err != nil { + t.Fatal(err) + } + + now := time.Now().In(time.Local) + userID := int(now.UnixNano() + testUserSequence.Add(1)) + key, _ := quotaWindow(userID, now) + t.Cleanup(func() { _ = db.Redis.Del(context.Background(), key).Err() }) + + return db.Redis, userID, now +} + +func assertCounterValue(t *testing.T, ctx context.Context, client *redis.Client, key string, expected int64) { + t.Helper() + + actual, err := client.Get(ctx, key).Int64() + + if err != nil { + t.Fatal(err) + } + + if actual != expected { + t.Fatalf("counter = %d, want %d", actual, expected) + } +} diff --git a/go.mod b/go.mod index eb25056..5d99754 100644 --- a/go.mod +++ b/go.mod @@ -5,13 +5,13 @@ go 1.23.2 require ( github.com/Azure/azure-pipeline-go v0.2.3 github.com/Azure/azure-storage-blob-go v0.15.0 - github.com/JGLTechnologies/gin-rate-limit v1.5.4 github.com/aws/aws-sdk-go v1.55.5 github.com/disgoorg/disgo v0.18.13 github.com/elastic/go-elasticsearch/v8 v8.15.0 github.com/gabriel-vasile/mimetype v1.4.5 github.com/gin-contrib/cors v1.7.2 github.com/gin-gonic/gin v1.10.0 + github.com/go-redis/redis_rate/v10 v10.0.1 github.com/go-resty/resty/v2 v2.15.3 github.com/golang-jwt/jwt/v5 v5.2.1 github.com/golang-migrate/migrate/v4 v4.18.1 diff --git a/go.sum b/go.sum index f26aab8..1adc7fd 100644 --- a/go.sum +++ b/go.sum @@ -18,8 +18,6 @@ github.com/Azure/go-autorest/logger v0.2.1 h1:IG7i4p/mDa2Ce4TRyAO8IHnVhAVF3RFU+Z github.com/Azure/go-autorest/logger v0.2.1/go.mod h1:T9E3cAhj2VqvPOtCYAvby9aBXkZmbF5NWuPV8+WeEW8= github.com/Azure/go-autorest/tracing v0.6.0 h1:TYi4+3m5t6K48TGI9AUdb+IzbnSxvnvUMfuitfgcfuo= github.com/Azure/go-autorest/tracing v0.6.0/go.mod h1:+vhtPC754Xsa23ID7GlGsrdKBpUA79WCAKPPZVC2DeU= -github.com/JGLTechnologies/gin-rate-limit v1.5.4 h1:1hIaXIdGM9MZFZlXgjWJLpxaK0WHEa5MeloK49nmQsc= -github.com/JGLTechnologies/gin-rate-limit v1.5.4/go.mod h1:mGEhNzlHEg/Tk+KH/mKylZLTfDjACnx7MVYaAlj07eU= github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY= github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU= github.com/aws/aws-sdk-go v1.55.5 h1:KKUZBfBoyqy5d3swXyiC7Q76ic40rYcbqH7qjh59kzU= @@ -91,6 +89,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.22.1 h1:40JcKH+bBNGFczGuoBYgX4I6m/i27HYW8P9FDk5PbgA= github.com/go-playground/validator/v10 v10.22.1/go.mod h1:dbuPbCMFw/DrkbEynArYaCwl3amGuJotoKCe95atGMM= +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-resty/resty/v2 v2.15.3 h1:bqff+hcqAflpiF591hhJzNdkRsFhlB96CYfBwSFvql8= github.com/go-resty/resty/v2 v2.15.3/go.mod h1:0fHAoK7JoBy/Ch36N8VFeMsK7xQOHhvWaC3iOktwmIU= github.com/go-sql-driver/mysql v1.7.0/go.mod h1:OXbVy3sEdcQ2Doequ6Z5BW6fXNQTmx+9S1MCJN5yJMI= diff --git a/handlers/download.go b/handlers/download.go index 79df1f8..4270975 100644 --- a/handlers/download.go +++ b/handlers/download.go @@ -4,11 +4,13 @@ import ( "errors" "fmt" "github.com/Quaver/api2/db" + "github.com/Quaver/api2/downloadlimit" "github.com/Quaver/api2/files" "github.com/Quaver/api2/tools" "github.com/gin-gonic/gin" "github.com/sirupsen/logrus" "gorm.io/gorm" + "net/http" "os" "strconv" "time" @@ -79,6 +81,10 @@ func DownloadMapset(c *gin.Context) *APIError { return setFileContentLength(c, path) } + if apiErr := enforceMapsetDownloadLimit(c, user.Id, path); apiErr != nil { + return apiErr + } + if err := db.InsertMapsetDownload(&db.MapsetDownload{ UserId: user.Id, MapsetId: mapset.Id, @@ -183,4 +189,29 @@ func setFileContentLength(c *gin.Context, path string) *APIError { return nil } +func enforceMapsetDownloadLimit(c *gin.Context, userID int, path string) *APIError { + fileInfo, err := os.Stat(path) + + if err != nil { + return APIErrorServerError("Error getting mapset file information", err) + } + + if fileInfo.Size() == 0 { + return nil + } + allowed, err := downloadlimit.TryConsume(c.Request.Context(), db.Redis, userID, fileInfo.Size()) + + if err != nil { + return APIErrorServerError("Error checking mapset download rate limit", err) + } + + if !allowed { + return &APIError{ + Status: http.StatusTooManyRequests, + Message: "Download rate limit has been reached", + } + } + + return nil +} diff --git a/handlers/download_test.go b/handlers/download_test.go new file mode 100644 index 0000000..92db106 --- /dev/null +++ b/handlers/download_test.go @@ -0,0 +1,159 @@ +package handlers + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strconv" + "sync/atomic" + "testing" + "time" + + "github.com/Quaver/api2/config" + "github.com/Quaver/api2/db" + "github.com/Quaver/api2/downloadlimit" + "github.com/gin-gonic/gin" + "github.com/redis/go-redis/v9" +) + +var downloadTestUserSequence atomic.Int64 + +func TestEnforceMapsetDownloadLimitUsesFileSize(t *testing.T) { + client, userID := useTestDownloadRedis(t) + path := writeTestMapset(t, []byte("mapset-data")) + ctx := newDownloadTestContext() + + if apiErr := enforceMapsetDownloadLimit(ctx, userID, path); apiErr != nil { + t.Fatalf("API error = %#v", apiErr) + } + + actual, err := client.Get(context.Background(), downloadLimitKey(userID)).Int64() + if err != nil { + t.Fatal(err) + } + + if expected := int64(len("mapset-data")); actual != expected { + t.Fatalf("counted bytes = %d, want %d", actual, expected) + } +} + +func TestEnforceMapsetDownloadLimitReturnsTooManyRequests(t *testing.T) { + client, userID := useTestDownloadRedis(t) + ctx := newDownloadTestContext() + allowed, err := downloadlimit.TryConsume( + ctx.Request.Context(), + client, + userID, + downloadlimit.DailyLimitBytes, + ) + + if err != nil || !allowed { + t.Fatalf("initial usage: allowed=%v err=%v", allowed, err) + } + + apiErr := enforceMapsetDownloadLimit(ctx, userID, writeTestMapset(t, []byte{1})) + + if apiErr == nil || apiErr.Status != http.StatusTooManyRequests { + t.Fatalf("API error = %#v, want status %d", apiErr, http.StatusTooManyRequests) + } + + if apiErr.Message != "Download rate limit has been reached" { + t.Fatalf("message = %q", apiErr.Message) + } + + actual, err := client.Get(context.Background(), downloadLimitKey(userID)).Int64() + if err != nil { + t.Fatal(err) + } + + if actual != downloadlimit.DailyLimitBytes { + t.Fatalf("counter = %d, want %d", actual, downloadlimit.DailyLimitBytes) + } +} + +func TestEnforceMapsetDownloadLimitReturnsServerError(t *testing.T) { + previousClient := db.Redis + client := redis.NewClient(&redis.Options{ + Addr: "127.0.0.1:0", + MaxRetries: -1, + DialTimeout: 10 * time.Millisecond, + ReadTimeout: 10 * time.Millisecond, + WriteTimeout: 10 * time.Millisecond, + }) + db.Redis = client + t.Cleanup(func() { + db.Redis = previousClient + _ = client.Close() + }) + + apiErr := enforceMapsetDownloadLimit(newDownloadTestContext(), 303, writeTestMapset(t, []byte{1})) + + if apiErr == nil || apiErr.Status != http.StatusInternalServerError { + t.Fatalf("API error = %#v, want status %d", apiErr, http.StatusInternalServerError) + } +} + +func TestSetFileContentLength(t *testing.T) { + path := writeTestMapset(t, []byte("head-response")) + response := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(response) + ctx.Request = httptest.NewRequest(http.MethodHead, "/v2/download/mapset/1", nil) + + if apiErr := setFileContentLength(ctx, path); apiErr != nil { + t.Fatalf("API error = %#v", apiErr) + } + + if got := response.Header().Get("Content-Length"); got != strconv.Itoa(len("head-response")) { + t.Fatalf("Content-Length = %q, want %d", got, len("head-response")) + } +} + +func useTestDownloadRedis(t *testing.T) (*redis.Client, int) { + t.Helper() + + if config.Instance == nil { + if err := config.Load("../config.json"); err != nil { + t.Fatal(err) + } + } + + db.InitializeRedis() + if err := db.Redis.Ping(context.Background()).Err(); err != nil { + t.Fatal(err) + } + + userID := int(time.Now().UnixNano() + downloadTestUserSequence.Add(1)) + key := downloadLimitKey(userID) + t.Cleanup(func() { _ = db.Redis.Del(context.Background(), key).Err() }) + + return db.Redis, userID +} + +func downloadLimitKey(userID int) string { + return fmt.Sprintf( + "quaver:download_rate_limit:%s:%d", + time.Now().In(time.Local).Format("2006-01-02"), + userID, + ) +} + +func newDownloadTestContext() *gin.Context { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Request = httptest.NewRequest(http.MethodGet, "/v2/download/mapset/1", nil) + return ctx +} + +func writeTestMapset(t *testing.T, contents []byte) string { + t.Helper() + + path := filepath.Join(t.TempDir(), "mapset.qp") + + if err := os.WriteFile(path, contents, 0o600); err != nil { + t.Fatal(err) + } + + return path +}