Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 55 additions & 15 deletions cmd/api/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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)
Expand Down Expand Up @@ -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
}
Expand All @@ -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.
Expand Down
60 changes: 60 additions & 0 deletions downloadlimit/downloadlimit.go
Original file line number Diff line number Diff line change
@@ -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
}
160 changes: 160 additions & 0 deletions downloadlimit/downloadlimit_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions go.sum
Original file line number Diff line number Diff line change
Expand Up @@ -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=
Expand Down Expand Up @@ -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=
Expand Down
Loading
Loading