From 8ce959cca730cb26a690035b707627f0818bcf94 Mon Sep 17 00:00:00 2001 From: AiAe Date: Fri, 18 Sep 2026 22:59:49 +0300 Subject: [PATCH 1/5] New mapset download rate limit --- cmd/api/server.go | 5 +- downloadlimit/downloadlimit.go | 248 ++++++++++++++++++++ downloadlimit/downloadlimit_test.go | 345 ++++++++++++++++++++++++++++ go.mod | 2 + go.sum | 4 + handlers/download.go | 104 +++++++++ handlers/download_test.go | 308 +++++++++++++++++++++++++ 7 files changed, 1014 insertions(+), 2 deletions(-) create mode 100644 downloadlimit/downloadlimit.go create mode 100644 downloadlimit/downloadlimit_test.go create mode 100644 handlers/download_test.go diff --git a/cmd/api/server.go b/cmd/api/server.go index c67c74d..ff599ba 100644 --- a/cmd/api/server.go +++ b/cmd/api/server.go @@ -49,8 +49,9 @@ func initializeServer(port int) { // Initializes the rate limiter for the server func initializeRateLimiter(engine *gin.Engine) { 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{ diff --git a/downloadlimit/downloadlimit.go b/downloadlimit/downloadlimit.go new file mode 100644 index 0000000..a678842 --- /dev/null +++ b/downloadlimit/downloadlimit.go @@ -0,0 +1,248 @@ +package downloadlimit + +import ( + "context" + "errors" + "fmt" + "time" + + "github.com/redis/go-redis/v9" +) + +const ( + // DailyLimitBytes is the maximum number of mapset bytes a user can download + // during a single server-local calendar day. + DailyLimitBytes int64 = 1 << 30 // 1GB, easier test with 50MB max: 50 * 1024 * 1024 + + maxTransactionRetries = 32 + maxRetryBackoff = 10 * time.Millisecond +) + +var ( + errLimitExceeded = errors.New("daily download limit exceeded") + errTransactionRetriesExhausted = errors.New("download limit transaction retries exhausted") +) + +// Reservation identifies bytes that were added to a daily download counter. +// Its fields are intentionally private so callers can only release reservations +// that were returned by TryReserve. +type Reservation struct { + key string + byteCount int64 + expiresAt time.Time +} + +type transactionHook func(attempt int) error + +// TryReserve atomically reserves byteCount bytes from a user's server-local +// daily download allowance. The returned boolean is false when the reservation +// would exceed DailyLimitBytes. +func TryReserve(ctx context.Context, redisClient *redis.Client, userID int, byteCount int64) (*Reservation, bool, error) { + return tryReserveAt(ctx, redisClient, userID, byteCount, time.Now().In(time.Local), nil) +} + +func tryReserveAt( + ctx context.Context, + redisClient *redis.Client, + userID int, + byteCount int64, + now time.Time, + beforeCommit transactionHook, +) (*Reservation, bool, error) { + if redisClient == nil { + return nil, false, errors.New("redis client is nil") + } + + if userID <= 0 { + return nil, false, fmt.Errorf("invalid user id: %d", userID) + } + + if byteCount < 0 { + return nil, false, fmt.Errorf("invalid download byte count: %d", byteCount) + } + + key, expiresAt := quotaWindow(userID, now) + reservation := &Reservation{ + key: key, + byteCount: byteCount, + expiresAt: expiresAt, + } + + err := watchWithRetry(ctx, redisClient, key, func(tx *redis.Tx, attempt int) error { + current, err := tx.Get(ctx, key).Int64() + + if err == redis.Nil { + current = 0 + } else if err != nil { + return fmt.Errorf("read daily download counter: %w", err) + } + + if current < 0 { + return fmt.Errorf("daily download counter cannot be negative: %d", current) + } + + if beforeCommit != nil { + if err := beforeCommit(attempt); err != nil { + return err + } + } + + if current > DailyLimitBytes || byteCount > DailyLimitBytes-current { + // Execute a read-only transaction so WATCH can detect a concurrent + // release before we reject based on the value read above. + _, err = tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + pipe.Exists(ctx, key) + return nil + }) + + if err != nil { + return err + } + + return errLimitExceeded + } + + _, err = tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + pipe.IncrBy(ctx, key, byteCount) + pipe.ExpireAt(ctx, key, expiresAt) + return nil + }) + + return err + }) + + if errors.Is(err, errLimitExceeded) { + return nil, false, nil + } + + if err != nil { + return nil, false, fmt.Errorf("reserve daily download bytes: %w", err) + } + + return reservation, true, nil +} + +// Release removes a previously reserved byte count. It is safe to call after +// the reservation's key has expired; in that case it is a no-op. +func Release(ctx context.Context, redisClient *redis.Client, reservation *Reservation) error { + return release(ctx, redisClient, reservation, nil) +} + +func release( + ctx context.Context, + redisClient *redis.Client, + reservation *Reservation, + beforeCommit transactionHook, +) error { + if redisClient == nil { + return errors.New("redis client is nil") + } + + if reservation == nil { + return errors.New("download reservation is nil") + } + + err := watchWithRetry(ctx, redisClient, reservation.key, func(tx *redis.Tx, attempt int) error { + current, err := tx.Get(ctx, reservation.key).Int64() + + if err == redis.Nil { + return nil + } + + if err != nil { + return fmt.Errorf("read daily download counter: %w", err) + } + + if current < 0 { + return fmt.Errorf("daily download counter cannot be negative: %d", current) + } + + if current < reservation.byteCount { + return fmt.Errorf( + "daily download counter underflow: current=%d reservation=%d", + current, + reservation.byteCount, + ) + } + + if beforeCommit != nil { + if err := beforeCommit(attempt); err != nil { + return err + } + } + + _, err = tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + if current == reservation.byteCount { + pipe.Del(ctx, reservation.key) + return nil + } + + pipe.DecrBy(ctx, reservation.key, reservation.byteCount) + pipe.ExpireAt(ctx, reservation.key, reservation.expiresAt) + return nil + }) + + return err + }) + + if err != nil { + return fmt.Errorf("release daily download bytes: %w", err) + } + + return nil +} + +func watchWithRetry( + ctx context.Context, + redisClient *redis.Client, + key string, + operation func(tx *redis.Tx, attempt int) error, +) error { + for attempt := 0; attempt < maxTransactionRetries; attempt++ { + if err := ctx.Err(); err != nil { + return err + } + + err := redisClient.Watch(ctx, func(tx *redis.Tx) error { + return operation(tx, attempt) + }, key) + + if !errors.Is(err, redis.TxFailedErr) { + return err + } + + if attempt == maxTransactionRetries-1 { + break + } + + if err := waitForRetry(ctx, attempt); err != nil { + return err + } + } + + return fmt.Errorf("%w for key %q", errTransactionRetriesExhausted, key) +} + +func waitForRetry(ctx context.Context, attempt int) error { + delay := time.Duration(attempt+1) * time.Millisecond + + if delay > maxRetryBackoff { + delay = maxRetryBackoff + } + + timer := time.NewTimer(delay) + defer timer.Stop() + + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return 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..c0d6acf --- /dev/null +++ b/downloadlimit/downloadlimit_test.go @@ -0,0 +1,345 @@ +package downloadlimit + +import ( + "context" + "errors" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" +) + +func TestTryReserveAccumulatesAndEnforcesLimit(t *testing.T) { + server, client := newTestRedis(t) + ctx := context.Background() + now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) + server.SetTime(now) + + first, allowed, err := tryReserveAt(ctx, client, 42, 128, now, nil) + if err != nil { + t.Fatal(err) + } + + if !allowed || first == nil { + t.Fatal("expected first reservation to be allowed") + } + + second, allowed, err := tryReserveAt(ctx, client, 42, DailyLimitBytes-128, now, nil) + if err != nil { + t.Fatal(err) + } + + if !allowed || second == nil { + t.Fatal("expected reservation reaching the exact limit to be allowed") + } + + reservation, allowed, err := tryReserveAt(ctx, client, 42, 1, now, nil) + if err != nil { + t.Fatal(err) + } + + if allowed || reservation != nil { + t.Fatal("expected reservation above the limit to be rejected") + } + + key, expiresAt := quotaWindow(42, now) + assertCounterValue(t, ctx, client, key, DailyLimitBytes) + + if ttl := server.TTL(key); ttl != expiresAt.Sub(now) { + t.Fatalf("TTL = %v, want %v", ttl, expiresAt.Sub(now)) + } +} + +func TestTryReserveRetriesWatchConflict(t *testing.T) { + server, client := newTestRedis(t) + otherClient := redis.NewClient(client.Options()) + t.Cleanup(func() { _ = otherClient.Close() }) + + ctx := context.Background() + now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) + server.SetTime(now) + key, _ := quotaWindow(7, now) + hookCalls := 0 + + reservation, allowed, err := tryReserveAt(ctx, client, 7, 11, now, func(attempt int) error { + hookCalls++ + + if attempt == 0 { + return otherClient.IncrBy(ctx, key, 7).Err() + } + + return nil + }) + + if err != nil { + t.Fatal(err) + } + + if !allowed || reservation == nil { + t.Fatal("expected reservation to succeed after retry") + } + + if hookCalls < 2 { + t.Fatalf("hook called %d times, want at least 2", hookCalls) + } + + assertCounterValue(t, ctx, client, key, 18) +} + +func TestTryReserveRetriesStaleRejection(t *testing.T) { + server, client := newTestRedis(t) + otherClient := redis.NewClient(client.Options()) + t.Cleanup(func() { _ = otherClient.Close() }) + + ctx := context.Background() + now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) + server.SetTime(now) + fullReservation, allowed, err := tryReserveAt(ctx, client, 8, DailyLimitBytes, now, nil) + + if err != nil || !allowed { + t.Fatalf("initial reservation: allowed=%v err=%v", allowed, err) + } + + hookCalls := 0 + reservation, allowed, err := tryReserveAt(ctx, client, 8, 1, now, func(attempt int) error { + hookCalls++ + + if attempt == 0 { + return Release(ctx, otherClient, fullReservation) + } + + return nil + }) + + if err != nil { + t.Fatal(err) + } + + if !allowed || reservation == nil { + t.Fatal("expected reservation to be allowed after the concurrent release") + } + + if hookCalls < 2 { + t.Fatalf("hook called %d times, want at least 2", hookCalls) + } + + key, _ := quotaWindow(8, now) + assertCounterValue(t, ctx, client, key, 1) +} + +func TestTryReserveConcurrentRequestsDoNotExceedLimit(t *testing.T) { + server, client := newTestRedis(t) + ctx := context.Background() + now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) + server.SetTime(now) + 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 := tryReserveAt(ctx, client, 99, chunkSize, now, nil) + + 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 reservations = %d, want %d", got, allowedCount) + } + + key, _ := quotaWindow(99, now) + assertCounterValue(t, ctx, client, key, DailyLimitBytes) +} + +func TestTryReserveExpiresAtNextLocalMidnightAcrossDST(t *testing.T) { + server, client := newTestRedis(t) + ctx := context.Background() + location, err := time.LoadLocation("Europe/Sofia") + + if err != nil { + t.Fatal(err) + } + + now := time.Date(2026, time.March, 29, 0, 30, 0, 0, location) + server.SetTime(now) + _, allowed, err := tryReserveAt(ctx, client, 15, 1, now, nil) + + if err != nil { + t.Fatal(err) + } + + if !allowed { + t.Fatal("expected reservation to be allowed") + } + + key, 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) + } + + if ttl := server.TTL(key); ttl != expectedTTL { + t.Fatalf("TTL = %v, want %v", ttl, expectedTTL) + } +} + +func TestReleaseRestoresCounterAndHandlesExpiredReservation(t *testing.T) { + server, client := newTestRedis(t) + ctx := context.Background() + dayOne := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) + server.SetTime(dayOne) + + first, allowed, err := tryReserveAt(ctx, client, 21, 100, dayOne, nil) + if err != nil || !allowed { + t.Fatalf("first reservation: allowed=%v err=%v", allowed, err) + } + + _, allowed, err = tryReserveAt(ctx, client, 21, 25, dayOne, nil) + if err != nil || !allowed { + t.Fatalf("second reservation: allowed=%v err=%v", allowed, err) + } + + if err := Release(ctx, client, first); err != nil { + t.Fatal(err) + } + + dayOneKey, dayOneExpiry := quotaWindow(21, dayOne) + assertCounterValue(t, ctx, client, dayOneKey, 25) + + server.FastForward(dayOneExpiry.Sub(dayOne) + time.Second) + dayTwo := dayOneExpiry.Add(time.Hour) + server.SetTime(dayTwo) + _, allowed, err = tryReserveAt(ctx, client, 21, 200, dayTwo, nil) + if err != nil || !allowed { + t.Fatalf("day two reservation: allowed=%v err=%v", allowed, err) + } + + if err := Release(ctx, client, first); err != nil { + t.Fatal(err) + } + + dayTwoKey, _ := quotaWindow(21, dayTwo) + assertCounterValue(t, ctx, client, dayTwoKey, 200) +} + +func TestReleaseRetriesConcurrentCounterChange(t *testing.T) { + server, client := newTestRedis(t) + otherClient := redis.NewClient(client.Options()) + t.Cleanup(func() { _ = otherClient.Close() }) + + ctx := context.Background() + now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) + server.SetTime(now) + reservation, allowed, err := tryReserveAt(ctx, client, 31, 100, now, nil) + + if err != nil || !allowed { + t.Fatalf("initial reservation: allowed=%v err=%v", allowed, err) + } + + hookCalls := 0 + err = release(ctx, client, reservation, func(attempt int) error { + hookCalls++ + + if attempt == 0 { + _, allowed, err := tryReserveAt(ctx, otherClient, 31, 25, now, nil) + + if err != nil { + return err + } + + if !allowed { + return errors.New("concurrent reservation was unexpectedly rejected") + } + } + + return nil + }) + + if err != nil { + t.Fatal(err) + } + + if hookCalls < 2 { + t.Fatalf("hook called %d times, want at least 2", hookCalls) + } + + key, _ := quotaWindow(31, now) + assertCounterValue(t, ctx, client, key, 25) +} + +func TestTryReserveReturnsRedisAndCounterErrors(t *testing.T) { + server, client := newTestRedis(t) + ctx := context.Background() + now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) + key, _ := quotaWindow(55, now) + server.Set(key, "not-an-integer") + + if _, _, err := tryReserveAt(ctx, client, 55, 1, now, nil); err == nil { + t.Fatal("expected malformed counter to return an error") + } + + server.Close() + + if _, _, err := tryReserveAt(ctx, client, 56, 1, now, nil); err == nil { + t.Fatal("expected unavailable redis to return an error") + } +} + +func newTestRedis(t *testing.T) (*miniredis.Miniredis, *redis.Client) { + t.Helper() + + server := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{ + Addr: server.Addr(), + MaxRetries: -1, + DialTimeout: 100 * time.Millisecond, + ReadTimeout: 100 * time.Millisecond, + WriteTimeout: 100 * time.Millisecond, + }) + t.Cleanup(func() { _ = client.Close() }) + + return server, client +} + +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..1cf30fe 100644 --- a/go.mod +++ b/go.mod @@ -6,6 +6,7 @@ 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/alicebob/miniredis/v2 v2.35.0 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 @@ -68,6 +69,7 @@ require ( github.com/spf13/pflag v1.0.5 // indirect github.com/twitchyliquid64/golang-asm v0.15.1 // indirect github.com/ugorji/go/codec v1.2.12 // indirect + github.com/yuin/gopher-lua v1.1.1 // indirect go.opentelemetry.io/otel v1.30.0 // indirect go.opentelemetry.io/otel/metric v1.30.0 // indirect go.opentelemetry.io/otel/trace v1.30.0 // indirect diff --git a/go.sum b/go.sum index f26aab8..cffd9d7 100644 --- a/go.sum +++ b/go.sum @@ -22,6 +22,8 @@ github.com/JGLTechnologies/gin-rate-limit v1.5.4 h1:1hIaXIdGM9MZFZlXgjWJLpxaK0WH 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/alicebob/miniredis/v2 v2.35.0 h1:QwLphYqCEAo1eu1TqPRN2jgVMPBweeQcR21jeqDCONI= +github.com/alicebob/miniredis/v2 v2.35.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM= github.com/aws/aws-sdk-go v1.55.5 h1:KKUZBfBoyqy5d3swXyiC7Q76ic40rYcbqH7qjh59kzU= github.com/aws/aws-sdk-go v1.55.5/go.mod h1:eRwEWoyTWFMVYVQzKMNHWP5/RV4xIUGMQfXQHfHkpNU= github.com/aymerick/douceur v0.2.0 h1:Mv+mAeH1Q+n9Fr+oyamOlAkUNPWPlA8PPGR0QAaYuPk= @@ -208,6 +210,8 @@ github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08= github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE= github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg= +github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M= +github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw= go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.54.0 h1:TT4fX+nBOA/+LUkobKGW1ydGcn+G3vRw9+g5HwCphpk= go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.54.0/go.mod h1:L7UH0GbB0p47T4Rri3uHjbpCFYrVrwc1I25QhNPiGK8= go.opentelemetry.io/otel v1.30.0 h1:F2t8sK4qf1fAmY9ua4ohFS/K+FUuOPemHUIXHtktrts= diff --git a/handlers/download.go b/handlers/download.go index 79df1f8..ff8365c 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,11 +81,23 @@ func DownloadMapset(c *gin.Context) *APIError { return setFileContentLength(c, path) } + reservation, apiErr := reserveMapsetDownloadQuota(c, user.Id, path) + + if apiErr != nil { + return apiErr + } + if err := db.InsertMapsetDownload(&db.MapsetDownload{ UserId: user.Id, MapsetId: mapset.Id, Timestamp: time.Now().UnixMilli(), }); err != nil { + if reservation != nil { + if releaseErr := downloadlimit.Release(db.RedisCtx, db.Redis, reservation); releaseErr != nil { + logrus.Errorf("Error releasing download quota for user %d: %v", user.Id, releaseErr) + } + } + return APIErrorServerError("Error inserting mapset download into db", err) } @@ -183,4 +197,94 @@ func setFileContentLength(c *gin.Context, path string) *APIError { return nil } +func reserveMapsetDownloadQuota(c *gin.Context, userID int, path string) (*downloadlimit.Reservation, *APIError) { + fileInfo, err := os.Stat(path) + + if err != nil { + return nil, APIErrorServerError("Error getting mapset file information", err) + } + + byteCount, err := mapsetResponseByteCount(c, path, fileInfo) + + if err != nil { + return nil, APIErrorServerError("Error determining mapset response size", err) + } + + if byteCount == 0 { + return nil, nil + } + + reservation, allowed, err := downloadlimit.TryReserve(c.Request.Context(), db.Redis, userID, byteCount) + + if err != nil { + return nil, APIErrorServerError("Error checking mapset download rate limit", err) + } + + if !allowed { + return nil, &APIError{ + Status: http.StatusTooManyRequests, + Message: "Download rate limit has been reached", + } + } + + return reservation, nil +} + +// mapsetResponseByteCount asks net/http to evaluate the request as a HEAD +// response so quota accounting follows the same range and precondition rules as +// FileAttachment without reading the response body. +func mapsetResponseByteCount(c *gin.Context, path string, fileInfo os.FileInfo) (int64, error) { + file, err := os.Open(path) + + if err != nil { + return 0, err + } + + defer file.Close() + + request := c.Request.Clone(c.Request.Context()) + request.Method = http.MethodHead + response := &responseMetadataWriter{header: c.Writer.Header().Clone()} + http.ServeContent(response, request, fileInfo.Name(), fileInfo.ModTime(), file) + + if response.status != http.StatusOK && response.status != http.StatusPartialContent { + return 0, nil + } + + contentLength := response.header.Get("Content-Length") + + if contentLength == "" && response.status == http.StatusOK { + return fileInfo.Size(), nil + } + + byteCount, err := strconv.ParseInt(contentLength, 10, 64) + + if err != nil || byteCount < 0 { + return 0, fmt.Errorf("invalid response content length %q", contentLength) + } + + return byteCount, nil +} + +type responseMetadataWriter struct { + header http.Header + status int +} + +func (w *responseMetadataWriter) Header() http.Header { + return w.header +} + +func (w *responseMetadataWriter) WriteHeader(status int) { + if w.status == 0 { + w.status = status + } +} + +func (w *responseMetadataWriter) Write(body []byte) (int, error) { + if w.status == 0 { + w.WriteHeader(http.StatusOK) + } + return len(body), nil +} diff --git a/handlers/download_test.go b/handlers/download_test.go new file mode 100644 index 0000000..0246683 --- /dev/null +++ b/handlers/download_test.go @@ -0,0 +1,308 @@ +package handlers + +import ( + "context" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strconv" + "testing" + "time" + + "github.com/Quaver/api2/db" + "github.com/Quaver/api2/downloadlimit" + "github.com/alicebob/miniredis/v2" + "github.com/gin-gonic/gin" + "github.com/redis/go-redis/v9" +) + +func TestReserveMapsetDownloadQuotaUsesFileSize(t *testing.T) { + _, client := useTestDownloadRedis(t) + path := writeTestMapset(t, []byte("mapset-data")) + ctx := newDownloadTestContext() + + reservation, apiErr := reserveMapsetDownloadQuota(ctx, 101, path) + + if apiErr != nil { + t.Fatalf("API error = %#v", apiErr) + } + + if reservation == nil { + t.Fatal("expected a quota reservation") + } + + keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() + + if err != nil { + t.Fatal(err) + } + + if len(keys) != 1 { + t.Fatalf("quota keys = %v, want exactly one", keys) + } + + actual, err := client.Get(context.Background(), keys[0]).Int64() + + if err != nil { + t.Fatal(err) + } + + if expected := int64(len("mapset-data")); actual != expected { + t.Fatalf("reserved bytes = %d, want %d", actual, expected) + } +} + +func TestReserveMapsetDownloadQuotaUsesRangeSize(t *testing.T) { + _, client := useTestDownloadRedis(t) + path := writeTestMapset(t, []byte("mapset-data")) + ctx := newDownloadTestContext() + ctx.Request.Header.Set("Range", "bytes=0-0") + initialReservation, allowed, err := downloadlimit.TryReserve( + ctx.Request.Context(), + client, + 102, + downloadlimit.DailyLimitBytes-1, + ) + + if err != nil || !allowed || initialReservation == nil { + t.Fatalf("initial reservation: allowed=%v reservation=%v err=%v", allowed, initialReservation, err) + } + + reservation, apiErr := reserveMapsetDownloadQuota(ctx, 102, path) + + if apiErr != nil { + t.Fatalf("API error = %#v", apiErr) + } + + if reservation == nil { + t.Fatal("expected a quota reservation") + } + + keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() + + if err != nil { + t.Fatal(err) + } + + if len(keys) != 1 { + t.Fatalf("quota keys = %v, want exactly one", keys) + } + + actual, err := client.Get(context.Background(), keys[0]).Int64() + + if err != nil { + t.Fatal(err) + } + + if actual != downloadlimit.DailyLimitBytes { + t.Fatalf("counter = %d, want %d", actual, downloadlimit.DailyLimitBytes) + } +} + +func TestReserveMapsetDownloadQuotaSkipsNotModifiedResponse(t *testing.T) { + _, client := useTestDownloadRedis(t) + path := writeTestMapset(t, []byte("mapset-data")) + fileInfo, err := os.Stat(path) + + if err != nil { + t.Fatal(err) + } + + ctx := newDownloadTestContext() + ctx.Request.Header.Set("If-Modified-Since", fileInfo.ModTime().UTC().Format(http.TimeFormat)) + reservation, apiErr := reserveMapsetDownloadQuota(ctx, 103, path) + + if apiErr != nil { + t.Fatalf("API error = %#v", apiErr) + } + + if reservation != nil { + t.Fatal("expected no quota reservation for a not-modified response") + } + + keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() + + if err != nil { + t.Fatal(err) + } + + if len(keys) != 0 { + t.Fatalf("not-modified response created quota keys: %v", keys) + } +} + +func TestMapsetResponseByteCountMatchesServedBody(t *testing.T) { + path := writeTestMapset(t, []byte("mapset-data")) + fileInfo, err := os.Stat(path) + + if err != nil { + t.Fatal(err) + } + + tests := []struct { + name string + headers http.Header + }{ + {name: "full response"}, + {name: "single range", headers: http.Header{"Range": {"bytes=2-5"}}}, + {name: "multiple ranges", headers: http.Header{"Range": {"bytes=0-0,2-3"}}}, + { + name: "not modified", + headers: http.Header{ + "If-Modified-Since": {fileInfo.ModTime().UTC().Format(http.TimeFormat)}, + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + response := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(response) + ctx.Request = httptest.NewRequest(http.MethodGet, "/v2/download/mapset/1", nil) + ctx.Request.Header = test.headers.Clone() + byteCount, err := mapsetResponseByteCount(ctx, path, fileInfo) + + if err != nil { + t.Fatal(err) + } + + ctx.FileAttachment(path, "1.qp") + + if byteCount != int64(response.Body.Len()) { + t.Fatalf("reserved bytes = %d, served body bytes = %d", byteCount, response.Body.Len()) + } + }) + } +} + +func TestReserveMapsetDownloadQuotaReturnsTooManyRequests(t *testing.T) { + _, client := useTestDownloadRedis(t) + ctx := newDownloadTestContext() + initialReservation, allowed, err := downloadlimit.TryReserve( + ctx.Request.Context(), + client, + 202, + downloadlimit.DailyLimitBytes, + ) + + if err != nil || !allowed || initialReservation == nil { + t.Fatalf("initial reservation: allowed=%v reservation=%v err=%v", allowed, initialReservation, err) + } + + reservation, apiErr := reserveMapsetDownloadQuota(ctx, 202, writeTestMapset(t, []byte{1})) + + if reservation != nil { + t.Fatal("expected no reservation for a rejected download") + } + + 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) + } + + keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() + + if err != nil { + t.Fatal(err) + } + + if len(keys) != 1 { + t.Fatalf("quota keys = %v, want exactly one", keys) + } + + actual, err := client.Get(context.Background(), keys[0]).Int64() + + if err != nil { + t.Fatal(err) + } + + if actual != downloadlimit.DailyLimitBytes { + t.Fatalf("counter = %d, want %d", actual, downloadlimit.DailyLimitBytes) + } +} + +func TestReserveMapsetDownloadQuotaReturnsServerError(t *testing.T) { + server, _ := useTestDownloadRedis(t) + server.Close() + ctx := newDownloadTestContext() + + reservation, apiErr := reserveMapsetDownloadQuota(ctx, 303, writeTestMapset(t, []byte{1})) + + if reservation != nil { + t.Fatal("expected no reservation when redis is unavailable") + } + + if apiErr == nil || apiErr.Status != http.StatusInternalServerError { + t.Fatalf("API error = %#v, want status %d", apiErr, http.StatusInternalServerError) + } +} + +func TestSetFileContentLengthDoesNotReserveQuota(t *testing.T) { + _, client := useTestDownloadRedis(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")) + } + + keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() + + if err != nil { + t.Fatal(err) + } + + if len(keys) != 0 { + t.Fatalf("HEAD request created quota keys: %v", keys) + } +} + +func useTestDownloadRedis(t *testing.T) (*miniredis.Miniredis, *redis.Client) { + t.Helper() + + server := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{ + Addr: server.Addr(), + MaxRetries: -1, + DialTimeout: 100 * time.Millisecond, + ReadTimeout: 100 * time.Millisecond, + WriteTimeout: 100 * time.Millisecond, + }) + previousClient := db.Redis + db.Redis = client + + t.Cleanup(func() { + db.Redis = previousClient + _ = client.Close() + }) + + return server, client +} + +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 +} From bf23bfb6e4372238bc4e4cee49eea9ec5427ce7f Mon Sep 17 00:00:00 2001 From: AiAe Date: Sat, 19 Sep 2026 12:38:09 +0300 Subject: [PATCH 2/5] Rewrite everything --- downloadlimit/downloadlimit.go | 230 ++-------------------------- downloadlimit/downloadlimit_test.go | 214 +++----------------------- handlers/download.go | 91 ++--------- handlers/download_test.go | 157 ++----------------- 4 files changed, 57 insertions(+), 635 deletions(-) diff --git a/downloadlimit/downloadlimit.go b/downloadlimit/downloadlimit.go index a678842..1576b8d 100644 --- a/downloadlimit/downloadlimit.go +++ b/downloadlimit/downloadlimit.go @@ -2,243 +2,47 @@ package downloadlimit import ( "context" - "errors" "fmt" "time" "github.com/redis/go-redis/v9" ) -const ( - // DailyLimitBytes is the maximum number of mapset bytes a user can download - // during a single server-local calendar day. - DailyLimitBytes int64 = 1 << 30 // 1GB, easier test with 50MB max: 50 * 1024 * 1024 +// 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 - maxTransactionRetries = 32 - maxRetryBackoff = 10 * time.Millisecond -) - -var ( - errLimitExceeded = errors.New("daily download limit exceeded") - errTransactionRetriesExhausted = errors.New("download limit transaction retries exhausted") -) - -// Reservation identifies bytes that were added to a daily download counter. -// Its fields are intentionally private so callers can only release reservations -// that were returned by TryReserve. -type Reservation struct { - key string - byteCount int64 - expiresAt time.Time +// 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)) } -type transactionHook func(attempt int) error - -// TryReserve atomically reserves byteCount bytes from a user's server-local -// daily download allowance. The returned boolean is false when the reservation -// would exceed DailyLimitBytes. -func TryReserve(ctx context.Context, redisClient *redis.Client, userID int, byteCount int64) (*Reservation, bool, error) { - return tryReserveAt(ctx, redisClient, userID, byteCount, time.Now().In(time.Local), nil) -} - -func tryReserveAt( +func tryConsumeAt( ctx context.Context, redisClient *redis.Client, userID int, byteCount int64, now time.Time, - beforeCommit transactionHook, -) (*Reservation, bool, error) { - if redisClient == nil { - return nil, false, errors.New("redis client is nil") - } - - if userID <= 0 { - return nil, false, fmt.Errorf("invalid user id: %d", userID) - } - - if byteCount < 0 { - return nil, false, fmt.Errorf("invalid download byte count: %d", byteCount) - } - +) (bool, error) { key, expiresAt := quotaWindow(userID, now) - reservation := &Reservation{ - key: key, - byteCount: byteCount, - expiresAt: expiresAt, - } - - err := watchWithRetry(ctx, redisClient, key, func(tx *redis.Tx, attempt int) error { - current, err := tx.Get(ctx, key).Int64() - - if err == redis.Nil { - current = 0 - } else if err != nil { - return fmt.Errorf("read daily download counter: %w", err) - } - - if current < 0 { - return fmt.Errorf("daily download counter cannot be negative: %d", current) - } - - if beforeCommit != nil { - if err := beforeCommit(attempt); err != nil { - return err - } - } - - if current > DailyLimitBytes || byteCount > DailyLimitBytes-current { - // Execute a read-only transaction so WATCH can detect a concurrent - // release before we reject based on the value read above. - _, err = tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { - pipe.Exists(ctx, key) - return nil - }) - - if err != nil { - return err - } - - return errLimitExceeded - } - - _, err = tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { - pipe.IncrBy(ctx, key, byteCount) - pipe.ExpireAt(ctx, key, expiresAt) - return nil - }) - - return err - }) - - if errors.Is(err, errLimitExceeded) { - return nil, false, nil - } - - if err != nil { - return nil, false, fmt.Errorf("reserve daily download bytes: %w", err) - } - - return reservation, true, nil -} - -// Release removes a previously reserved byte count. It is safe to call after -// the reservation's key has expired; in that case it is a no-op. -func Release(ctx context.Context, redisClient *redis.Client, reservation *Reservation) error { - return release(ctx, redisClient, reservation, nil) -} - -func release( - ctx context.Context, - redisClient *redis.Client, - reservation *Reservation, - beforeCommit transactionHook, -) error { - if redisClient == nil { - return errors.New("redis client is nil") - } - - if reservation == nil { - return errors.New("download reservation is nil") - } - - err := watchWithRetry(ctx, redisClient, reservation.key, func(tx *redis.Tx, attempt int) error { - current, err := tx.Get(ctx, reservation.key).Int64() - - if err == redis.Nil { - return nil - } - - if err != nil { - return fmt.Errorf("read daily download counter: %w", err) - } - - if current < 0 { - return fmt.Errorf("daily download counter cannot be negative: %d", current) - } - - if current < reservation.byteCount { - return fmt.Errorf( - "daily download counter underflow: current=%d reservation=%d", - current, - reservation.byteCount, - ) - } - - if beforeCommit != nil { - if err := beforeCommit(attempt); err != nil { - return err - } - } - - _, err = tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { - if current == reservation.byteCount { - pipe.Del(ctx, reservation.key) - return nil - } - - pipe.DecrBy(ctx, reservation.key, reservation.byteCount) - pipe.ExpireAt(ctx, reservation.key, reservation.expiresAt) - return nil - }) - - return err - }) + used, err := redisClient.IncrBy(ctx, key, byteCount).Result() if err != nil { - return fmt.Errorf("release daily download bytes: %w", err) + return false, fmt.Errorf("increment daily download bytes: %w", err) } - return nil -} - -func watchWithRetry( - ctx context.Context, - redisClient *redis.Client, - key string, - operation func(tx *redis.Tx, attempt int) error, -) error { - for attempt := 0; attempt < maxTransactionRetries; attempt++ { - if err := ctx.Err(); err != nil { - return err - } - - err := redisClient.Watch(ctx, func(tx *redis.Tx) error { - return operation(tx, attempt) - }, key) - - if !errors.Is(err, redis.TxFailedErr) { - return err - } - - if attempt == maxTransactionRetries-1 { - break - } - - if err := waitForRetry(ctx, attempt); err != nil { - return err + if used > DailyLimitBytes { + if err := redisClient.DecrBy(ctx, key, byteCount).Err(); err != nil { + return false, fmt.Errorf("restore daily download bytes: %w", err) } } - return fmt.Errorf("%w for key %q", errTransactionRetriesExhausted, key) -} - -func waitForRetry(ctx context.Context, attempt int) error { - delay := time.Duration(attempt+1) * time.Millisecond - - if delay > maxRetryBackoff { - delay = maxRetryBackoff + if err := redisClient.ExpireAt(ctx, key, expiresAt).Err(); err != nil { + return false, fmt.Errorf("expire daily download counter: %w", err) } - timer := time.NewTimer(delay) - defer timer.Stop() - - select { - case <-ctx.Done(): - return ctx.Err() - case <-timer.C: - return nil - } + return used <= DailyLimitBytes, nil } func quotaWindow(userID int, now time.Time) (string, time.Time) { diff --git a/downloadlimit/downloadlimit_test.go b/downloadlimit/downloadlimit_test.go index c0d6acf..25cb87d 100644 --- a/downloadlimit/downloadlimit_test.go +++ b/downloadlimit/downloadlimit_test.go @@ -2,7 +2,6 @@ package downloadlimit import ( "context" - "errors" "sync" "sync/atomic" "testing" @@ -12,37 +11,29 @@ import ( "github.com/redis/go-redis/v9" ) -func TestTryReserveAccumulatesAndEnforcesLimit(t *testing.T) { +func TestTryConsumeAccumulatesAndEnforcesLimit(t *testing.T) { server, client := newTestRedis(t) ctx := context.Background() now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) server.SetTime(now) - first, allowed, err := tryReserveAt(ctx, client, 42, 128, now, nil) - if err != nil { - t.Fatal(err) - } - - if !allowed || first == nil { - t.Fatal("expected first reservation to be allowed") - } - - second, allowed, err := tryReserveAt(ctx, client, 42, DailyLimitBytes-128, now, nil) - if err != nil { - t.Fatal(err) + allowed, err := tryConsumeAt(ctx, client, 42, 128, now) + if err != nil || !allowed { + t.Fatalf("first download: allowed=%v err=%v", allowed, err) } - if !allowed || second == nil { - t.Fatal("expected reservation reaching the exact limit to be allowed") + allowed, err = tryConsumeAt(ctx, client, 42, DailyLimitBytes-128, now) + if err != nil || !allowed { + t.Fatalf("download reaching limit: allowed=%v err=%v", allowed, err) } - reservation, allowed, err := tryReserveAt(ctx, client, 42, 1, now, nil) + allowed, err = tryConsumeAt(ctx, client, 42, 1, now) if err != nil { t.Fatal(err) } - if allowed || reservation != nil { - t.Fatal("expected reservation above the limit to be rejected") + if allowed { + t.Fatal("expected download above the limit to be rejected") } key, expiresAt := quotaWindow(42, now) @@ -53,84 +44,7 @@ func TestTryReserveAccumulatesAndEnforcesLimit(t *testing.T) { } } -func TestTryReserveRetriesWatchConflict(t *testing.T) { - server, client := newTestRedis(t) - otherClient := redis.NewClient(client.Options()) - t.Cleanup(func() { _ = otherClient.Close() }) - - ctx := context.Background() - now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) - server.SetTime(now) - key, _ := quotaWindow(7, now) - hookCalls := 0 - - reservation, allowed, err := tryReserveAt(ctx, client, 7, 11, now, func(attempt int) error { - hookCalls++ - - if attempt == 0 { - return otherClient.IncrBy(ctx, key, 7).Err() - } - - return nil - }) - - if err != nil { - t.Fatal(err) - } - - if !allowed || reservation == nil { - t.Fatal("expected reservation to succeed after retry") - } - - if hookCalls < 2 { - t.Fatalf("hook called %d times, want at least 2", hookCalls) - } - - assertCounterValue(t, ctx, client, key, 18) -} - -func TestTryReserveRetriesStaleRejection(t *testing.T) { - server, client := newTestRedis(t) - otherClient := redis.NewClient(client.Options()) - t.Cleanup(func() { _ = otherClient.Close() }) - - ctx := context.Background() - now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) - server.SetTime(now) - fullReservation, allowed, err := tryReserveAt(ctx, client, 8, DailyLimitBytes, now, nil) - - if err != nil || !allowed { - t.Fatalf("initial reservation: allowed=%v err=%v", allowed, err) - } - - hookCalls := 0 - reservation, allowed, err := tryReserveAt(ctx, client, 8, 1, now, func(attempt int) error { - hookCalls++ - - if attempt == 0 { - return Release(ctx, otherClient, fullReservation) - } - - return nil - }) - - if err != nil { - t.Fatal(err) - } - - if !allowed || reservation == nil { - t.Fatal("expected reservation to be allowed after the concurrent release") - } - - if hookCalls < 2 { - t.Fatalf("hook called %d times, want at least 2", hookCalls) - } - - key, _ := quotaWindow(8, now) - assertCounterValue(t, ctx, client, key, 1) -} - -func TestTryReserveConcurrentRequestsDoNotExceedLimit(t *testing.T) { +func TestTryConsumeConcurrentRequestsDoNotExceedLimit(t *testing.T) { server, client := newTestRedis(t) ctx := context.Background() now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) @@ -150,7 +64,7 @@ func TestTryReserveConcurrentRequestsDoNotExceedLimit(t *testing.T) { defer waitGroup.Done() <-start - _, allowed, err := tryReserveAt(ctx, client, 99, chunkSize, now, nil) + allowed, err := tryConsumeAt(ctx, client, 99, chunkSize, now) if err != nil { errorsChannel <- err @@ -172,14 +86,14 @@ func TestTryReserveConcurrentRequestsDoNotExceedLimit(t *testing.T) { } if got := successful.Load(); got != allowedCount { - t.Fatalf("allowed reservations = %d, want %d", got, allowedCount) + t.Fatalf("allowed downloads = %d, want %d", got, allowedCount) } key, _ := quotaWindow(99, now) assertCounterValue(t, ctx, client, key, DailyLimitBytes) } -func TestTryReserveExpiresAtNextLocalMidnightAcrossDST(t *testing.T) { +func TestTryConsumeExpiresAtNextLocalMidnightAcrossDST(t *testing.T) { server, client := newTestRedis(t) ctx := context.Background() location, err := time.LoadLocation("Europe/Sofia") @@ -190,14 +104,10 @@ func TestTryReserveExpiresAtNextLocalMidnightAcrossDST(t *testing.T) { now := time.Date(2026, time.March, 29, 0, 30, 0, 0, location) server.SetTime(now) - _, allowed, err := tryReserveAt(ctx, client, 15, 1, now, nil) + allowed, err := tryConsumeAt(ctx, client, 15, 1, now) - if err != nil { - t.Fatal(err) - } - - if !allowed { - t.Fatal("expected reservation to be allowed") + if err != nil || !allowed { + t.Fatalf("download: allowed=%v err=%v", allowed, err) } key, expiresAt := quotaWindow(15, now) @@ -212,104 +122,20 @@ func TestTryReserveExpiresAtNextLocalMidnightAcrossDST(t *testing.T) { } } -func TestReleaseRestoresCounterAndHandlesExpiredReservation(t *testing.T) { - server, client := newTestRedis(t) - ctx := context.Background() - dayOne := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) - server.SetTime(dayOne) - - first, allowed, err := tryReserveAt(ctx, client, 21, 100, dayOne, nil) - if err != nil || !allowed { - t.Fatalf("first reservation: allowed=%v err=%v", allowed, err) - } - - _, allowed, err = tryReserveAt(ctx, client, 21, 25, dayOne, nil) - if err != nil || !allowed { - t.Fatalf("second reservation: allowed=%v err=%v", allowed, err) - } - - if err := Release(ctx, client, first); err != nil { - t.Fatal(err) - } - - dayOneKey, dayOneExpiry := quotaWindow(21, dayOne) - assertCounterValue(t, ctx, client, dayOneKey, 25) - - server.FastForward(dayOneExpiry.Sub(dayOne) + time.Second) - dayTwo := dayOneExpiry.Add(time.Hour) - server.SetTime(dayTwo) - _, allowed, err = tryReserveAt(ctx, client, 21, 200, dayTwo, nil) - if err != nil || !allowed { - t.Fatalf("day two reservation: allowed=%v err=%v", allowed, err) - } - - if err := Release(ctx, client, first); err != nil { - t.Fatal(err) - } - - dayTwoKey, _ := quotaWindow(21, dayTwo) - assertCounterValue(t, ctx, client, dayTwoKey, 200) -} - -func TestReleaseRetriesConcurrentCounterChange(t *testing.T) { - server, client := newTestRedis(t) - otherClient := redis.NewClient(client.Options()) - t.Cleanup(func() { _ = otherClient.Close() }) - - ctx := context.Background() - now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) - server.SetTime(now) - reservation, allowed, err := tryReserveAt(ctx, client, 31, 100, now, nil) - - if err != nil || !allowed { - t.Fatalf("initial reservation: allowed=%v err=%v", allowed, err) - } - - hookCalls := 0 - err = release(ctx, client, reservation, func(attempt int) error { - hookCalls++ - - if attempt == 0 { - _, allowed, err := tryReserveAt(ctx, otherClient, 31, 25, now, nil) - - if err != nil { - return err - } - - if !allowed { - return errors.New("concurrent reservation was unexpectedly rejected") - } - } - - return nil - }) - - if err != nil { - t.Fatal(err) - } - - if hookCalls < 2 { - t.Fatalf("hook called %d times, want at least 2", hookCalls) - } - - key, _ := quotaWindow(31, now) - assertCounterValue(t, ctx, client, key, 25) -} - -func TestTryReserveReturnsRedisAndCounterErrors(t *testing.T) { +func TestTryConsumeReturnsRedisErrors(t *testing.T) { server, client := newTestRedis(t) ctx := context.Background() now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) key, _ := quotaWindow(55, now) server.Set(key, "not-an-integer") - if _, _, err := tryReserveAt(ctx, client, 55, 1, now, nil); err == nil { + if _, err := tryConsumeAt(ctx, client, 55, 1, now); err == nil { t.Fatal("expected malformed counter to return an error") } server.Close() - if _, _, err := tryReserveAt(ctx, client, 56, 1, now, nil); err == nil { + if _, err := tryConsumeAt(ctx, client, 56, 1, now); err == nil { t.Fatal("expected unavailable redis to return an error") } } diff --git a/handlers/download.go b/handlers/download.go index ff8365c..4270975 100644 --- a/handlers/download.go +++ b/handlers/download.go @@ -81,9 +81,7 @@ func DownloadMapset(c *gin.Context) *APIError { return setFileContentLength(c, path) } - reservation, apiErr := reserveMapsetDownloadQuota(c, user.Id, path) - - if apiErr != nil { + if apiErr := enforceMapsetDownloadLimit(c, user.Id, path); apiErr != nil { return apiErr } @@ -92,12 +90,6 @@ func DownloadMapset(c *gin.Context) *APIError { MapsetId: mapset.Id, Timestamp: time.Now().UnixMilli(), }); err != nil { - if reservation != nil { - if releaseErr := downloadlimit.Release(db.RedisCtx, db.Redis, reservation); releaseErr != nil { - logrus.Errorf("Error releasing download quota for user %d: %v", user.Id, releaseErr) - } - } - return APIErrorServerError("Error inserting mapset download into db", err) } @@ -197,94 +189,29 @@ func setFileContentLength(c *gin.Context, path string) *APIError { return nil } -func reserveMapsetDownloadQuota(c *gin.Context, userID int, path string) (*downloadlimit.Reservation, *APIError) { +func enforceMapsetDownloadLimit(c *gin.Context, userID int, path string) *APIError { fileInfo, err := os.Stat(path) if err != nil { - return nil, APIErrorServerError("Error getting mapset file information", err) + return APIErrorServerError("Error getting mapset file information", err) } - byteCount, err := mapsetResponseByteCount(c, path, fileInfo) - - if err != nil { - return nil, APIErrorServerError("Error determining mapset response size", err) - } - - if byteCount == 0 { - return nil, nil + if fileInfo.Size() == 0 { + return nil } - reservation, allowed, err := downloadlimit.TryReserve(c.Request.Context(), db.Redis, userID, byteCount) + allowed, err := downloadlimit.TryConsume(c.Request.Context(), db.Redis, userID, fileInfo.Size()) if err != nil { - return nil, APIErrorServerError("Error checking mapset download rate limit", err) + return APIErrorServerError("Error checking mapset download rate limit", err) } if !allowed { - return nil, &APIError{ + return &APIError{ Status: http.StatusTooManyRequests, Message: "Download rate limit has been reached", } } - return reservation, nil -} - -// mapsetResponseByteCount asks net/http to evaluate the request as a HEAD -// response so quota accounting follows the same range and precondition rules as -// FileAttachment without reading the response body. -func mapsetResponseByteCount(c *gin.Context, path string, fileInfo os.FileInfo) (int64, error) { - file, err := os.Open(path) - - if err != nil { - return 0, err - } - - defer file.Close() - - request := c.Request.Clone(c.Request.Context()) - request.Method = http.MethodHead - response := &responseMetadataWriter{header: c.Writer.Header().Clone()} - http.ServeContent(response, request, fileInfo.Name(), fileInfo.ModTime(), file) - - if response.status != http.StatusOK && response.status != http.StatusPartialContent { - return 0, nil - } - - contentLength := response.header.Get("Content-Length") - - if contentLength == "" && response.status == http.StatusOK { - return fileInfo.Size(), nil - } - - byteCount, err := strconv.ParseInt(contentLength, 10, 64) - - if err != nil || byteCount < 0 { - return 0, fmt.Errorf("invalid response content length %q", contentLength) - } - - return byteCount, nil -} - -type responseMetadataWriter struct { - header http.Header - status int -} - -func (w *responseMetadataWriter) Header() http.Header { - return w.header -} - -func (w *responseMetadataWriter) WriteHeader(status int) { - if w.status == 0 { - w.status = status - } -} - -func (w *responseMetadataWriter) Write(body []byte) (int, error) { - if w.status == 0 { - w.WriteHeader(http.StatusOK) - } - - return len(body), nil + return nil } diff --git a/handlers/download_test.go b/handlers/download_test.go index 0246683..1bf54ed 100644 --- a/handlers/download_test.go +++ b/handlers/download_test.go @@ -17,21 +17,17 @@ import ( "github.com/redis/go-redis/v9" ) -func TestReserveMapsetDownloadQuotaUsesFileSize(t *testing.T) { +func TestEnforceMapsetDownloadLimitUsesFileSize(t *testing.T) { _, client := useTestDownloadRedis(t) path := writeTestMapset(t, []byte("mapset-data")) ctx := newDownloadTestContext() - reservation, apiErr := reserveMapsetDownloadQuota(ctx, 101, path) + apiErr := enforceMapsetDownloadLimit(ctx, 101, path) if apiErr != nil { t.Fatalf("API error = %#v", apiErr) } - if reservation == nil { - t.Fatal("expected a quota reservation") - } - keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() if err != nil { @@ -49,152 +45,25 @@ func TestReserveMapsetDownloadQuotaUsesFileSize(t *testing.T) { } if expected := int64(len("mapset-data")); actual != expected { - t.Fatalf("reserved bytes = %d, want %d", actual, expected) - } -} - -func TestReserveMapsetDownloadQuotaUsesRangeSize(t *testing.T) { - _, client := useTestDownloadRedis(t) - path := writeTestMapset(t, []byte("mapset-data")) - ctx := newDownloadTestContext() - ctx.Request.Header.Set("Range", "bytes=0-0") - initialReservation, allowed, err := downloadlimit.TryReserve( - ctx.Request.Context(), - client, - 102, - downloadlimit.DailyLimitBytes-1, - ) - - if err != nil || !allowed || initialReservation == nil { - t.Fatalf("initial reservation: allowed=%v reservation=%v err=%v", allowed, initialReservation, err) - } - - reservation, apiErr := reserveMapsetDownloadQuota(ctx, 102, path) - - if apiErr != nil { - t.Fatalf("API error = %#v", apiErr) - } - - if reservation == nil { - t.Fatal("expected a quota reservation") - } - - keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() - - if err != nil { - t.Fatal(err) - } - - if len(keys) != 1 { - t.Fatalf("quota keys = %v, want exactly one", keys) - } - - actual, err := client.Get(context.Background(), keys[0]).Int64() - - if err != nil { - t.Fatal(err) - } - - if actual != downloadlimit.DailyLimitBytes { - t.Fatalf("counter = %d, want %d", actual, downloadlimit.DailyLimitBytes) - } -} - -func TestReserveMapsetDownloadQuotaSkipsNotModifiedResponse(t *testing.T) { - _, client := useTestDownloadRedis(t) - path := writeTestMapset(t, []byte("mapset-data")) - fileInfo, err := os.Stat(path) - - if err != nil { - t.Fatal(err) - } - - ctx := newDownloadTestContext() - ctx.Request.Header.Set("If-Modified-Since", fileInfo.ModTime().UTC().Format(http.TimeFormat)) - reservation, apiErr := reserveMapsetDownloadQuota(ctx, 103, path) - - if apiErr != nil { - t.Fatalf("API error = %#v", apiErr) - } - - if reservation != nil { - t.Fatal("expected no quota reservation for a not-modified response") - } - - keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() - - if err != nil { - t.Fatal(err) - } - - if len(keys) != 0 { - t.Fatalf("not-modified response created quota keys: %v", keys) - } -} - -func TestMapsetResponseByteCountMatchesServedBody(t *testing.T) { - path := writeTestMapset(t, []byte("mapset-data")) - fileInfo, err := os.Stat(path) - - if err != nil { - t.Fatal(err) - } - - tests := []struct { - name string - headers http.Header - }{ - {name: "full response"}, - {name: "single range", headers: http.Header{"Range": {"bytes=2-5"}}}, - {name: "multiple ranges", headers: http.Header{"Range": {"bytes=0-0,2-3"}}}, - { - name: "not modified", - headers: http.Header{ - "If-Modified-Since": {fileInfo.ModTime().UTC().Format(http.TimeFormat)}, - }, - }, - } - - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - response := httptest.NewRecorder() - ctx, _ := gin.CreateTestContext(response) - ctx.Request = httptest.NewRequest(http.MethodGet, "/v2/download/mapset/1", nil) - ctx.Request.Header = test.headers.Clone() - byteCount, err := mapsetResponseByteCount(ctx, path, fileInfo) - - if err != nil { - t.Fatal(err) - } - - ctx.FileAttachment(path, "1.qp") - - if byteCount != int64(response.Body.Len()) { - t.Fatalf("reserved bytes = %d, served body bytes = %d", byteCount, response.Body.Len()) - } - }) + t.Fatalf("counted bytes = %d, want %d", actual, expected) } } -func TestReserveMapsetDownloadQuotaReturnsTooManyRequests(t *testing.T) { +func TestEnforceMapsetDownloadLimitReturnsTooManyRequests(t *testing.T) { _, client := useTestDownloadRedis(t) ctx := newDownloadTestContext() - initialReservation, allowed, err := downloadlimit.TryReserve( + allowed, err := downloadlimit.TryConsume( ctx.Request.Context(), client, 202, downloadlimit.DailyLimitBytes, ) - if err != nil || !allowed || initialReservation == nil { - t.Fatalf("initial reservation: allowed=%v reservation=%v err=%v", allowed, initialReservation, err) + if err != nil || !allowed { + t.Fatalf("initial usage: allowed=%v err=%v", allowed, err) } - reservation, apiErr := reserveMapsetDownloadQuota(ctx, 202, writeTestMapset(t, []byte{1})) - - if reservation != nil { - t.Fatal("expected no reservation for a rejected download") - } + apiErr := enforceMapsetDownloadLimit(ctx, 202, writeTestMapset(t, []byte{1})) if apiErr == nil || apiErr.Status != http.StatusTooManyRequests { t.Fatalf("API error = %#v, want status %d", apiErr, http.StatusTooManyRequests) @@ -225,23 +94,19 @@ func TestReserveMapsetDownloadQuotaReturnsTooManyRequests(t *testing.T) { } } -func TestReserveMapsetDownloadQuotaReturnsServerError(t *testing.T) { +func TestEnforceMapsetDownloadLimitReturnsServerError(t *testing.T) { server, _ := useTestDownloadRedis(t) server.Close() ctx := newDownloadTestContext() - reservation, apiErr := reserveMapsetDownloadQuota(ctx, 303, writeTestMapset(t, []byte{1})) - - if reservation != nil { - t.Fatal("expected no reservation when redis is unavailable") - } + apiErr := enforceMapsetDownloadLimit(ctx, 303, writeTestMapset(t, []byte{1})) if apiErr == nil || apiErr.Status != http.StatusInternalServerError { t.Fatalf("API error = %#v, want status %d", apiErr, http.StatusInternalServerError) } } -func TestSetFileContentLengthDoesNotReserveQuota(t *testing.T) { +func TestSetFileContentLengthDoesNotCountQuota(t *testing.T) { _, client := useTestDownloadRedis(t) path := writeTestMapset(t, []byte("head-response")) response := httptest.NewRecorder() From b064a55949c69eb66f5a7efcd62427651a3a1627 Mon Sep 17 00:00:00 2001 From: AiAe Date: Sat, 19 Sep 2026 12:48:10 +0300 Subject: [PATCH 3/5] Rewrite tests to drop usage of miniredis --- downloadlimit/downloadlimit_test.go | 95 +++++++++++------------- go.mod | 2 - go.sum | 4 - handlers/download_test.go | 110 ++++++++++++---------------- 4 files changed, 90 insertions(+), 121 deletions(-) diff --git a/downloadlimit/downloadlimit_test.go b/downloadlimit/downloadlimit_test.go index 25cb87d..b95faa4 100644 --- a/downloadlimit/downloadlimit_test.go +++ b/downloadlimit/downloadlimit_test.go @@ -7,27 +7,28 @@ import ( "testing" "time" - "github.com/alicebob/miniredis/v2" + "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) { - server, client := newTestRedis(t) + client, userID, now := useTestRedis(t) ctx := context.Background() - now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) - server.SetTime(now) - allowed, err := tryConsumeAt(ctx, client, 42, 128, now) + 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, 42, DailyLimitBytes-128, now) + 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, 42, 1, now) + allowed, err = tryConsumeAt(ctx, client, userID, 1, now) if err != nil { t.Fatal(err) } @@ -36,19 +37,17 @@ func TestTryConsumeAccumulatesAndEnforcesLimit(t *testing.T) { t.Fatal("expected download above the limit to be rejected") } - key, expiresAt := quotaWindow(42, now) + key, _ := quotaWindow(userID, now) assertCounterValue(t, ctx, client, key, DailyLimitBytes) - if ttl := server.TTL(key); ttl != expiresAt.Sub(now) { - t.Fatalf("TTL = %v, want %v", ttl, expiresAt.Sub(now)) + 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) { - server, client := newTestRedis(t) + client, userID, now := useTestRedis(t) ctx := context.Background() - now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) - server.SetTime(now) const workerCount = 24 const allowedCount = 8 chunkSize := DailyLimitBytes / allowedCount @@ -64,7 +63,7 @@ func TestTryConsumeConcurrentRequestsDoNotExceedLimit(t *testing.T) { defer waitGroup.Done() <-start - allowed, err := tryConsumeAt(ctx, client, 99, chunkSize, now) + allowed, err := tryConsumeAt(ctx, client, userID, chunkSize, now) if err != nil { errorsChannel <- err @@ -89,13 +88,11 @@ func TestTryConsumeConcurrentRequestsDoNotExceedLimit(t *testing.T) { t.Fatalf("allowed downloads = %d, want %d", got, allowedCount) } - key, _ := quotaWindow(99, now) + key, _ := quotaWindow(userID, now) assertCounterValue(t, ctx, client, key, DailyLimitBytes) } -func TestTryConsumeExpiresAtNextLocalMidnightAcrossDST(t *testing.T) { - server, client := newTestRedis(t) - ctx := context.Background() +func TestQuotaWindowEndsAtNextLocalMidnightAcrossDST(t *testing.T) { location, err := time.LoadLocation("Europe/Sofia") if err != nil { @@ -103,57 +100,49 @@ func TestTryConsumeExpiresAtNextLocalMidnightAcrossDST(t *testing.T) { } now := time.Date(2026, time.March, 29, 0, 30, 0, 0, location) - server.SetTime(now) - allowed, err := tryConsumeAt(ctx, client, 15, 1, now) - - if err != nil || !allowed { - t.Fatalf("download: allowed=%v err=%v", allowed, err) - } - - key, expiresAt := quotaWindow(15, now) + _, 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) } - - if ttl := server.TTL(key); ttl != expectedTTL { - t.Fatalf("TTL = %v, want %v", ttl, expectedTTL) - } } func TestTryConsumeReturnsRedisErrors(t *testing.T) { - server, client := newTestRedis(t) - ctx := context.Background() - now := time.Date(2026, time.September, 18, 12, 0, 0, 0, time.UTC) - key, _ := quotaWindow(55, now) - server.Set(key, "not-an-integer") + 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 := tryConsumeAt(ctx, client, 55, 1, now); err == nil { - t.Fatal("expected malformed counter to return an error") + if _, err := TryConsume(context.Background(), client, 55, 1); err == nil { + t.Fatal("expected Redis error") } +} - server.Close() +func useTestRedis(t *testing.T) (*redis.Client, int, time.Time) { + t.Helper() - if _, err := tryConsumeAt(ctx, client, 56, 1, now); err == nil { - t.Fatal("expected unavailable redis to return an error") + if config.Instance == nil { + if err := config.Load("../config.json"); err != nil { + t.Fatal(err) + } } -} -func newTestRedis(t *testing.T) (*miniredis.Miniredis, *redis.Client) { - t.Helper() + db.InitializeRedis() + if err := db.Redis.Ping(context.Background()).Err(); err != nil { + t.Fatal(err) + } - server := miniredis.RunT(t) - client := redis.NewClient(&redis.Options{ - Addr: server.Addr(), - MaxRetries: -1, - DialTimeout: 100 * time.Millisecond, - ReadTimeout: 100 * time.Millisecond, - WriteTimeout: 100 * time.Millisecond, - }) - t.Cleanup(func() { _ = client.Close() }) + 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 server, client + return db.Redis, userID, now } func assertCounterValue(t *testing.T, ctx context.Context, client *redis.Client, key string, expected int64) { diff --git a/go.mod b/go.mod index 1cf30fe..eb25056 100644 --- a/go.mod +++ b/go.mod @@ -6,7 +6,6 @@ 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/alicebob/miniredis/v2 v2.35.0 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 @@ -69,7 +68,6 @@ require ( github.com/spf13/pflag v1.0.5 // indirect github.com/twitchyliquid64/golang-asm v0.15.1 // indirect github.com/ugorji/go/codec v1.2.12 // indirect - github.com/yuin/gopher-lua v1.1.1 // indirect go.opentelemetry.io/otel v1.30.0 // indirect go.opentelemetry.io/otel/metric v1.30.0 // indirect go.opentelemetry.io/otel/trace v1.30.0 // indirect diff --git a/go.sum b/go.sum index cffd9d7..f26aab8 100644 --- a/go.sum +++ b/go.sum @@ -22,8 +22,6 @@ github.com/JGLTechnologies/gin-rate-limit v1.5.4 h1:1hIaXIdGM9MZFZlXgjWJLpxaK0WH 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/alicebob/miniredis/v2 v2.35.0 h1:QwLphYqCEAo1eu1TqPRN2jgVMPBweeQcR21jeqDCONI= -github.com/alicebob/miniredis/v2 v2.35.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM= github.com/aws/aws-sdk-go v1.55.5 h1:KKUZBfBoyqy5d3swXyiC7Q76ic40rYcbqH7qjh59kzU= github.com/aws/aws-sdk-go v1.55.5/go.mod h1:eRwEWoyTWFMVYVQzKMNHWP5/RV4xIUGMQfXQHfHkpNU= github.com/aymerick/douceur v0.2.0 h1:Mv+mAeH1Q+n9Fr+oyamOlAkUNPWPlA8PPGR0QAaYuPk= @@ -210,8 +208,6 @@ github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08= github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE= github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg= -github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M= -github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw= go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.54.0 h1:TT4fX+nBOA/+LUkobKGW1ydGcn+G3vRw9+g5HwCphpk= go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.54.0/go.mod h1:L7UH0GbB0p47T4Rri3uHjbpCFYrVrwc1I25QhNPiGK8= go.opentelemetry.io/otel v1.30.0 h1:F2t8sK4qf1fAmY9ua4ohFS/K+FUuOPemHUIXHtktrts= diff --git a/handlers/download_test.go b/handlers/download_test.go index 1bf54ed..92db106 100644 --- a/handlers/download_test.go +++ b/handlers/download_test.go @@ -2,44 +2,35 @@ 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/alicebob/miniredis/v2" "github.com/gin-gonic/gin" "github.com/redis/go-redis/v9" ) +var downloadTestUserSequence atomic.Int64 + func TestEnforceMapsetDownloadLimitUsesFileSize(t *testing.T) { - _, client := useTestDownloadRedis(t) + client, userID := useTestDownloadRedis(t) path := writeTestMapset(t, []byte("mapset-data")) ctx := newDownloadTestContext() - apiErr := enforceMapsetDownloadLimit(ctx, 101, path) - - if apiErr != nil { + if apiErr := enforceMapsetDownloadLimit(ctx, userID, path); apiErr != nil { t.Fatalf("API error = %#v", apiErr) } - keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() - - if err != nil { - t.Fatal(err) - } - - if len(keys) != 1 { - t.Fatalf("quota keys = %v, want exactly one", keys) - } - - actual, err := client.Get(context.Background(), keys[0]).Int64() - + actual, err := client.Get(context.Background(), downloadLimitKey(userID)).Int64() if err != nil { t.Fatal(err) } @@ -50,12 +41,12 @@ func TestEnforceMapsetDownloadLimitUsesFileSize(t *testing.T) { } func TestEnforceMapsetDownloadLimitReturnsTooManyRequests(t *testing.T) { - _, client := useTestDownloadRedis(t) + client, userID := useTestDownloadRedis(t) ctx := newDownloadTestContext() allowed, err := downloadlimit.TryConsume( ctx.Request.Context(), client, - 202, + userID, downloadlimit.DailyLimitBytes, ) @@ -63,7 +54,7 @@ func TestEnforceMapsetDownloadLimitReturnsTooManyRequests(t *testing.T) { t.Fatalf("initial usage: allowed=%v err=%v", allowed, err) } - apiErr := enforceMapsetDownloadLimit(ctx, 202, writeTestMapset(t, []byte{1})) + 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) @@ -73,18 +64,7 @@ func TestEnforceMapsetDownloadLimitReturnsTooManyRequests(t *testing.T) { t.Fatalf("message = %q", apiErr.Message) } - keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() - - if err != nil { - t.Fatal(err) - } - - if len(keys) != 1 { - t.Fatalf("quota keys = %v, want exactly one", keys) - } - - actual, err := client.Get(context.Background(), keys[0]).Int64() - + actual, err := client.Get(context.Background(), downloadLimitKey(userID)).Int64() if err != nil { t.Fatal(err) } @@ -95,19 +75,28 @@ func TestEnforceMapsetDownloadLimitReturnsTooManyRequests(t *testing.T) { } func TestEnforceMapsetDownloadLimitReturnsServerError(t *testing.T) { - server, _ := useTestDownloadRedis(t) - server.Close() - ctx := newDownloadTestContext() + 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(ctx, 303, writeTestMapset(t, []byte{1})) + 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 TestSetFileContentLengthDoesNotCountQuota(t *testing.T) { - _, client := useTestDownloadRedis(t) +func TestSetFileContentLength(t *testing.T) { path := writeTestMapset(t, []byte("head-response")) response := httptest.NewRecorder() ctx, _ := gin.CreateTestContext(response) @@ -120,38 +109,35 @@ func TestSetFileContentLengthDoesNotCountQuota(t *testing.T) { if got := response.Header().Get("Content-Length"); got != strconv.Itoa(len("head-response")) { t.Fatalf("Content-Length = %q, want %d", got, len("head-response")) } +} - keys, err := client.Keys(context.Background(), "quaver:download_rate_limit:*").Result() +func useTestDownloadRedis(t *testing.T) (*redis.Client, int) { + t.Helper() - if err != nil { - t.Fatal(err) + if config.Instance == nil { + if err := config.Load("../config.json"); err != nil { + t.Fatal(err) + } } - if len(keys) != 0 { - t.Fatalf("HEAD request created quota keys: %v", keys) + db.InitializeRedis() + if err := db.Redis.Ping(context.Background()).Err(); err != nil { + t.Fatal(err) } -} -func useTestDownloadRedis(t *testing.T) (*miniredis.Miniredis, *redis.Client) { - t.Helper() + userID := int(time.Now().UnixNano() + downloadTestUserSequence.Add(1)) + key := downloadLimitKey(userID) + t.Cleanup(func() { _ = db.Redis.Del(context.Background(), key).Err() }) - server := miniredis.RunT(t) - client := redis.NewClient(&redis.Options{ - Addr: server.Addr(), - MaxRetries: -1, - DialTimeout: 100 * time.Millisecond, - ReadTimeout: 100 * time.Millisecond, - WriteTimeout: 100 * time.Millisecond, - }) - previousClient := db.Redis - db.Redis = client - - t.Cleanup(func() { - db.Redis = previousClient - _ = client.Close() - }) + return db.Redis, userID +} - return server, client +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 { From fd8d815f33d181fbad05d8cd122cb1852caa70ca Mon Sep 17 00:00:00 2001 From: AiAe Date: Sat, 19 Sep 2026 13:33:16 +0300 Subject: [PATCH 4/5] Use redis_rate for rate limit --- cmd/api/server.go | 65 +++++++++++++++++++++++++++++++++++++---------- go.mod | 2 +- go.sum | 4 +-- 3 files changed, 55 insertions(+), 16 deletions(-) diff --git a/cmd/api/server.go b/cmd/api/server.go index ff599ba..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,19 +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/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 } @@ -74,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/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= From 3d0f7088be5e91c830bf8b904711f9d7b3d14031 Mon Sep 17 00:00:00 2001 From: AiAe Date: Sat, 19 Sep 2026 20:16:25 +0300 Subject: [PATCH 5/5] Add mutex to tryConsumeAt --- downloadlimit/downloadlimit.go | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/downloadlimit/downloadlimit.go b/downloadlimit/downloadlimit.go index 1576b8d..d229a71 100644 --- a/downloadlimit/downloadlimit.go +++ b/downloadlimit/downloadlimit.go @@ -3,6 +3,7 @@ package downloadlimit import ( "context" "fmt" + "sync" "time" "github.com/redis/go-redis/v9" @@ -12,6 +13,10 @@ import ( // 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) { @@ -25,6 +30,9 @@ func tryConsumeAt( 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()