Files
VLoop/backend/internal/middleware/ratelimit/ratelimit.go
2026-05-20 16:34:46 +08:00

73 lines
1.5 KiB
Go

package ratelimit
import (
jwt "feedsystem_video_go/internal/middleware/jwt"
rediscache "feedsystem_video_go/internal/middleware/redis"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
)
type KeyFunc func(*gin.Context) (string, bool)
func Limit(
cache *rediscache.Client,
keyPrefix string,
maxRequests int64,
window time.Duration,
keyFunc KeyFunc,
) gin.HandlerFunc {
return func(c *gin.Context) {
if cache == nil || keyFunc == nil || maxRequests <= 0 || window <= 0 {
c.Next()
return
}
subject, ok := keyFunc(c)
if !ok {
c.Next()
return
}
key := buildKey(keyPrefix, subject)
count, err := cache.IncrementWithExpire(c.Request.Context(), key, window)
if err != nil {
c.Next()
return
}
if count > maxRequests {
c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{
"error": "too many requests",
})
return
}
c.Next()
}
}
func buildKey(keyPrefix, subject string) string {
keyPrefix = strings.TrimSpace(keyPrefix)
if keyPrefix == "" {
keyPrefix = "default"
}
return fmt.Sprintf("feedsystem:ratelimit:%s:%s", keyPrefix, strings.TrimSpace(subject))
}
func KeyByIP(c *gin.Context) (string, bool) {
ip := strings.TrimSpace(c.ClientIP())
if ip == "" {
return "", false
}
return ip, true
}
func KeyByAccount(c *gin.Context) (string, bool) {
accountID, err := jwt.GetAccountID(c)
if err != nil || accountID == 0 {
return "", false
}
return strconv.FormatUint(uint64(accountID), 10), true
}