mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-28 17:49:08 +02:00
[management] move rate limiter to shared package (#7727)
This commit is contained in:
@@ -38,11 +38,11 @@ import (
|
||||
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
||||
nbContext "github.com/netbirdio/netbird/management/server/context"
|
||||
nbhttp "github.com/netbirdio/netbird/management/server/http"
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/management/server/idp"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/telemetry"
|
||||
mgmtProto "github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
"github.com/netbirdio/netbird/util/crypt"
|
||||
)
|
||||
|
||||
@@ -171,10 +171,10 @@ func (s *BaseServer) Router() *mux.Router {
|
||||
})
|
||||
}
|
||||
|
||||
func (s *BaseServer) RateLimiter() *middleware.APIRateLimiter {
|
||||
return Create(s, func() *middleware.APIRateLimiter {
|
||||
cfg, enabled := middleware.RateLimiterConfigFromEnv()
|
||||
limiter := middleware.NewAPIRateLimiter(cfg)
|
||||
func (s *BaseServer) RateLimiter() *ratelimit.APIRateLimiter {
|
||||
return Create(s, func() *ratelimit.APIRateLimiter {
|
||||
cfg, enabled := ratelimit.RateLimiterConfigFromEnv()
|
||||
limiter := ratelimit.NewAPIRateLimiter(cfg)
|
||||
limiter.SetEnabled(enabled)
|
||||
return limiter
|
||||
})
|
||||
|
||||
@@ -58,10 +58,11 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/networks/resources"
|
||||
"github.com/netbirdio/netbird/management/server/networks/routers"
|
||||
"github.com/netbirdio/netbird/management/server/telemetry"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
)
|
||||
|
||||
// NewAPIHandler creates the Management service HTTP API handler registering all the available endpoints.
|
||||
func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *middleware.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager) (http.Handler, error) {
|
||||
func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *ratelimit.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager) (http.Handler, error) {
|
||||
|
||||
// Register bypass paths for unauthenticated endpoints
|
||||
if err := bypass.AddBypassPath("/api/instance"); err != nil {
|
||||
@@ -84,7 +85,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou
|
||||
|
||||
if rateLimiter == nil {
|
||||
log.Warn("NewAPIHandler: nil rate limiter, rate limiting disabled")
|
||||
rateLimiter = middleware.NewAPIRateLimiter(nil)
|
||||
rateLimiter = ratelimit.NewAPIRateLimiter(nil)
|
||||
rateLimiter.SetEnabled(false)
|
||||
}
|
||||
|
||||
|
||||
@@ -16,21 +16,21 @@ import (
|
||||
"golang.org/x/oauth2"
|
||||
|
||||
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
)
|
||||
|
||||
// AuthCallbackHandler handles OAuth callbacks for proxy authentication.
|
||||
type AuthCallbackHandler struct {
|
||||
proxyService *nbgrpc.ProxyServiceServer
|
||||
rateLimiter *middleware.APIRateLimiter
|
||||
rateLimiter *ratelimit.APIRateLimiter
|
||||
trustedProxies []netip.Prefix
|
||||
}
|
||||
|
||||
// NewAuthCallbackHandler creates a new OAuth callback handler.
|
||||
func NewAuthCallbackHandler(proxyService *nbgrpc.ProxyServiceServer, trustedProxies []netip.Prefix) *AuthCallbackHandler {
|
||||
rateLimiterConfig := &middleware.RateLimiterConfig{
|
||||
rateLimiterConfig := &ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 10,
|
||||
Burst: 15,
|
||||
CleanupInterval: 5 * time.Minute,
|
||||
@@ -39,7 +39,7 @@ func NewAuthCallbackHandler(proxyService *nbgrpc.ProxyServiceServer, trustedProx
|
||||
|
||||
return &AuthCallbackHandler{
|
||||
proxyService: proxyService,
|
||||
rateLimiter: middleware.NewAPIRateLimiter(rateLimiterConfig),
|
||||
rateLimiter: ratelimit.NewAPIRateLimiter(rateLimiterConfig),
|
||||
trustedProxies: trustedProxies,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,15 +11,15 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
"github.com/netbirdio/netbird/shared/management/http/util"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
)
|
||||
|
||||
// publicInviteRateLimiter limits public invite requests by IP address to prevent brute-force attacks
|
||||
var publicInviteRateLimiter = middleware.NewAPIRateLimiter(&middleware.RateLimiterConfig{
|
||||
var publicInviteRateLimiter = ratelimit.NewAPIRateLimiter(&ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 10, // 10 attempts per minute per IP
|
||||
Burst: 5, // Allow burst of 5 requests
|
||||
CleanupInterval: 10 * time.Minute,
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/auth"
|
||||
"github.com/netbirdio/netbird/shared/management/http/util"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
)
|
||||
|
||||
type EnsureAccountFunc func(ctx context.Context, userAuth auth.UserAuth) (string, string, error)
|
||||
@@ -33,7 +34,7 @@ type AuthMiddleware struct {
|
||||
ensureAccount EnsureAccountFunc
|
||||
getUserFromUserAuth GetUserFromUserAuthFunc
|
||||
syncUserJWTGroups SyncUserJWTGroupsFunc
|
||||
rateLimiter *APIRateLimiter
|
||||
rateLimiter *ratelimit.APIRateLimiter
|
||||
patUsageTracker *PATUsageTracker
|
||||
isValidChildAccount IsValidChildAccountFunc
|
||||
}
|
||||
@@ -44,7 +45,7 @@ func NewAuthMiddleware(
|
||||
ensureAccount EnsureAccountFunc,
|
||||
syncUserJWTGroups SyncUserJWTGroupsFunc,
|
||||
getUserFromUserAuth GetUserFromUserAuthFunc,
|
||||
rateLimiter *APIRateLimiter,
|
||||
rateLimiter *ratelimit.APIRateLimiter,
|
||||
meter metric.Meter,
|
||||
isValidChildAccount IsValidChildAccountFunc,
|
||||
) *AuthMiddleware {
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/util"
|
||||
nbauth "github.com/netbirdio/netbird/shared/auth"
|
||||
nbjwt "github.com/netbirdio/netbird/shared/auth/jwt"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -196,7 +197,7 @@ func TestAuthMiddleware_Handler(t *testing.T) {
|
||||
GetPATInfoFunc: mockGetAccountInfoFromPAT,
|
||||
}
|
||||
|
||||
disabledLimiter := NewAPIRateLimiter(nil)
|
||||
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
|
||||
disabledLimiter.SetEnabled(false)
|
||||
authMiddleware := NewAuthMiddleware(
|
||||
mockAuth,
|
||||
@@ -260,7 +261,7 @@ func TestAuthMiddleware_SyncUserJWTGroupsDetachedFromRequestCancellation(t *test
|
||||
GetPATInfoFunc: mockGetAccountInfoFromPAT,
|
||||
}
|
||||
|
||||
disabledLimiter := NewAPIRateLimiter(nil)
|
||||
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
|
||||
disabledLimiter.SetEnabled(false)
|
||||
|
||||
authMiddleware := NewAuthMiddleware(
|
||||
@@ -311,7 +312,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
|
||||
t.Run("PAT Token Rate Limiting - Burst Works", func(t *testing.T) {
|
||||
// Configure rate limiter: 10 requests per minute with burst of 5
|
||||
rateLimitConfig := &RateLimiterConfig{
|
||||
rateLimitConfig := &ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 10,
|
||||
Burst: 5,
|
||||
CleanupInterval: 5 * time.Minute,
|
||||
@@ -329,7 +330,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
|
||||
return &types.User{}, nil
|
||||
},
|
||||
NewAPIRateLimiter(rateLimitConfig),
|
||||
ratelimit.NewAPIRateLimiter(rateLimitConfig),
|
||||
nil,
|
||||
func(_ context.Context, _, _, _ string) bool { return false },
|
||||
)
|
||||
@@ -364,7 +365,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
|
||||
t.Run("PAT Token Rate Limiting - Rate Limit Enforced", func(t *testing.T) {
|
||||
// Configure very low rate limit: 1 request per minute
|
||||
rateLimitConfig := &RateLimiterConfig{
|
||||
rateLimitConfig := &ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 1,
|
||||
Burst: 1,
|
||||
CleanupInterval: 5 * time.Minute,
|
||||
@@ -382,7 +383,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
|
||||
return &types.User{}, nil
|
||||
},
|
||||
NewAPIRateLimiter(rateLimitConfig),
|
||||
ratelimit.NewAPIRateLimiter(rateLimitConfig),
|
||||
nil,
|
||||
func(_ context.Context, _, _, _ string) bool { return false },
|
||||
)
|
||||
@@ -408,7 +409,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
|
||||
t.Run("Bearer Token Not Rate Limited", func(t *testing.T) {
|
||||
// Configure strict rate limit
|
||||
rateLimitConfig := &RateLimiterConfig{
|
||||
rateLimitConfig := &ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 1,
|
||||
Burst: 1,
|
||||
CleanupInterval: 5 * time.Minute,
|
||||
@@ -426,7 +427,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
|
||||
return &types.User{}, nil
|
||||
},
|
||||
NewAPIRateLimiter(rateLimitConfig),
|
||||
ratelimit.NewAPIRateLimiter(rateLimitConfig),
|
||||
nil,
|
||||
func(_ context.Context, _, _, _ string) bool { return false },
|
||||
)
|
||||
@@ -453,7 +454,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
|
||||
t.Run("PAT Token Rate Limiting Per Token", func(t *testing.T) {
|
||||
// Configure rate limiter
|
||||
rateLimitConfig := &RateLimiterConfig{
|
||||
rateLimitConfig := &ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 1,
|
||||
Burst: 1,
|
||||
CleanupInterval: 5 * time.Minute,
|
||||
@@ -471,7 +472,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
|
||||
return &types.User{}, nil
|
||||
},
|
||||
NewAPIRateLimiter(rateLimitConfig),
|
||||
ratelimit.NewAPIRateLimiter(rateLimitConfig),
|
||||
nil,
|
||||
func(_ context.Context, _, _, _ string) bool { return false },
|
||||
)
|
||||
@@ -518,7 +519,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
|
||||
t.Run("Rate Limiter Cleanup", func(t *testing.T) {
|
||||
// Configure rate limiter with short cleanup interval and TTL for testing
|
||||
rateLimitConfig := &RateLimiterConfig{
|
||||
rateLimitConfig := &ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 60,
|
||||
Burst: 1,
|
||||
CleanupInterval: 100 * time.Millisecond,
|
||||
@@ -536,7 +537,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
|
||||
return &types.User{}, nil
|
||||
},
|
||||
NewAPIRateLimiter(rateLimitConfig),
|
||||
ratelimit.NewAPIRateLimiter(rateLimitConfig),
|
||||
nil,
|
||||
func(_ context.Context, _, _, _ string) bool { return false },
|
||||
)
|
||||
@@ -578,7 +579,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("Terraform User Agent Not Rate Limited", func(t *testing.T) {
|
||||
rateLimitConfig := &RateLimiterConfig{
|
||||
rateLimitConfig := &ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 1,
|
||||
Burst: 1,
|
||||
CleanupInterval: 5 * time.Minute,
|
||||
@@ -596,7 +597,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
|
||||
return &types.User{}, nil
|
||||
},
|
||||
NewAPIRateLimiter(rateLimitConfig),
|
||||
ratelimit.NewAPIRateLimiter(rateLimitConfig),
|
||||
nil,
|
||||
func(_ context.Context, _, _, _ string) bool { return false },
|
||||
)
|
||||
@@ -634,7 +635,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("Non-Terraform User Agent With PAT Is Rate Limited", func(t *testing.T) {
|
||||
rateLimitConfig := &RateLimiterConfig{
|
||||
rateLimitConfig := &ratelimit.RateLimiterConfig{
|
||||
RequestsPerMinute: 1,
|
||||
Burst: 1,
|
||||
CleanupInterval: 5 * time.Minute,
|
||||
@@ -652,7 +653,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
|
||||
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
|
||||
return &types.User{}, nil
|
||||
},
|
||||
NewAPIRateLimiter(rateLimitConfig),
|
||||
ratelimit.NewAPIRateLimiter(rateLimitConfig),
|
||||
nil,
|
||||
func(_ context.Context, _, _, _ string) bool { return false },
|
||||
)
|
||||
@@ -740,7 +741,7 @@ func TestAuthMiddleware_Handler_Child(t *testing.T) {
|
||||
GetPATInfoFunc: mockGetAccountInfoFromPAT,
|
||||
}
|
||||
|
||||
disabledLimiter := NewAPIRateLimiter(nil)
|
||||
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
|
||||
disabledLimiter.SetEnabled(false)
|
||||
authMiddleware := NewAuthMiddleware(
|
||||
mockAuth,
|
||||
|
||||
+21
-3
@@ -1,7 +1,8 @@
|
||||
package middleware
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
@@ -13,7 +14,6 @@ import (
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/time/rate"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/http/util"
|
||||
"github.com/netbirdio/netbird/trustedproxy"
|
||||
)
|
||||
|
||||
@@ -264,13 +264,31 @@ func (rl *APIRateLimiter) Middleware(next http.Handler) http.Handler {
|
||||
}
|
||||
clientIP := getClientIP(r, rl.config.TrustedProxies)
|
||||
if !rl.Allow(clientIP) {
|
||||
util.WriteErrorResponse("rate limit exceeded, please try again later", http.StatusTooManyRequests, w)
|
||||
writeTooManyRequests(w)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// errorResponse is the JSON body of a rejected request.
|
||||
type errorResponse struct {
|
||||
Message string `json:"message"`
|
||||
Code int `json:"code"`
|
||||
}
|
||||
|
||||
// writeTooManyRequests writes a JSON error response with status 429 Too Many Requests.
|
||||
func writeTooManyRequests(w http.ResponseWriter) {
|
||||
w.Header().Set("Content-Type", "application/json; charset=UTF-8")
|
||||
w.WriteHeader(http.StatusTooManyRequests)
|
||||
if err := json.NewEncoder(w).Encode(errorResponse{
|
||||
Message: "rate limit exceeded, please try again later",
|
||||
Code: http.StatusTooManyRequests,
|
||||
}); err != nil {
|
||||
log.Debugf("writing rate limit response: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// getClientIP extracts the client IP address from the request. Forwarding headers
|
||||
// are used only when the request arrives from a trusted proxy.
|
||||
func getClientIP(r *http.Request, trusted *trustedproxy.List) string {
|
||||
+3
-1
@@ -1,4 +1,4 @@
|
||||
package middleware
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
@@ -66,6 +66,8 @@ func TestAPIRateLimiter_Middleware(t *testing.T) {
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
assert.Equal(t, http.StatusTooManyRequests, rr.Code)
|
||||
assert.Equal(t, "application/json; charset=UTF-8", rr.Header().Get("Content-Type"))
|
||||
assert.JSONEq(t, `{"message":"rate limit exceeded, please try again later","code":429}`, rr.Body.String())
|
||||
}
|
||||
|
||||
func TestAPIRateLimiter_Middleware_DifferentIPs(t *testing.T) {
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
|
||||
"github.com/netbirdio/netbird/upload-server/types"
|
||||
)
|
||||
@@ -28,7 +28,7 @@ type local struct {
|
||||
signer *signer
|
||||
}
|
||||
|
||||
func configureLocalHandlers(mux *http.ServeMux, limiter *middleware.APIRateLimiter) error {
|
||||
func configureLocalHandlers(mux *http.ServeMux, limiter *ratelimit.APIRateLimiter) error {
|
||||
envURL, ok := os.LookupEnv("SERVER_URL")
|
||||
if !ok {
|
||||
return fmt.Errorf("SERVER_URL environment variable is required")
|
||||
|
||||
@@ -5,27 +5,27 @@ import (
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
)
|
||||
|
||||
const defaultUploadBurst = 100
|
||||
|
||||
func newRateLimiter() *middleware.APIRateLimiter {
|
||||
cfg, enabled := middleware.RateLimiterConfigFromEnv()
|
||||
if os.Getenv(middleware.RateLimitingBurstEnv) == "" {
|
||||
func newRateLimiter() *ratelimit.APIRateLimiter {
|
||||
cfg, enabled := ratelimit.RateLimiterConfigFromEnv()
|
||||
if os.Getenv(ratelimit.RateLimitingBurstEnv) == "" {
|
||||
cfg.Burst = defaultUploadBurst
|
||||
}
|
||||
|
||||
// Rate limiting is enabled by default unless explicitly disabled
|
||||
if os.Getenv(middleware.RateLimitingEnabledEnv) == "" {
|
||||
if os.Getenv(ratelimit.RateLimitingEnabledEnv) == "" {
|
||||
enabled = true
|
||||
}
|
||||
|
||||
limiter := middleware.NewAPIRateLimiter(cfg)
|
||||
limiter := ratelimit.NewAPIRateLimiter(cfg)
|
||||
limiter.SetEnabled(enabled)
|
||||
|
||||
log.Infof("Upload URL rate limiting: enabled=%t rate=%.0f/min burst=%d trusted_proxies=%q",
|
||||
limiter.Enabled(), cfg.RequestsPerMinute, cfg.Burst, os.Getenv(middleware.RateLimitingTrustedProxiesEnv))
|
||||
limiter.Enabled(), cfg.RequestsPerMinute, cfg.Burst, os.Getenv(ratelimit.RateLimitingTrustedProxiesEnv))
|
||||
|
||||
return limiter
|
||||
}
|
||||
|
||||
@@ -7,11 +7,11 @@ import (
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
"github.com/netbirdio/netbird/upload-server/types"
|
||||
)
|
||||
|
||||
func newTestRateLimiter(t *testing.T) *middleware.APIRateLimiter {
|
||||
func newTestRateLimiter(t *testing.T) *ratelimit.APIRateLimiter {
|
||||
t.Helper()
|
||||
|
||||
limiter := newRateLimiter()
|
||||
@@ -32,8 +32,8 @@ func getUploadURL(t *testing.T, mux *http.ServeMux) int {
|
||||
}
|
||||
|
||||
func Test_GetUploadURLIsRateLimited(t *testing.T) {
|
||||
t.Setenv(middleware.RateLimitingBurstEnv, "2")
|
||||
t.Setenv(middleware.RateLimitingRPMEnv, "1")
|
||||
t.Setenv(ratelimit.RateLimitingBurstEnv, "2")
|
||||
t.Setenv(ratelimit.RateLimitingRPMEnv, "1")
|
||||
mux, _ := newLocalMux(t)
|
||||
|
||||
require.Equal(t, http.StatusOK, getUploadURL(t, mux))
|
||||
@@ -42,8 +42,8 @@ func Test_GetUploadURLIsRateLimited(t *testing.T) {
|
||||
}
|
||||
|
||||
func Test_RateLimitingIsOnByDefault(t *testing.T) {
|
||||
t.Setenv(middleware.RateLimitingEnabledEnv, "")
|
||||
t.Setenv(middleware.RateLimitingBurstEnv, "1")
|
||||
t.Setenv(ratelimit.RateLimitingEnabledEnv, "")
|
||||
t.Setenv(ratelimit.RateLimitingBurstEnv, "1")
|
||||
mux, _ := newLocalMux(t)
|
||||
|
||||
require.Equal(t, http.StatusOK, getUploadURL(t, mux))
|
||||
@@ -51,8 +51,8 @@ func Test_RateLimitingIsOnByDefault(t *testing.T) {
|
||||
}
|
||||
|
||||
func Test_RateLimitingCanBeDisabled(t *testing.T) {
|
||||
t.Setenv(middleware.RateLimitingEnabledEnv, "false")
|
||||
t.Setenv(middleware.RateLimitingBurstEnv, "1")
|
||||
t.Setenv(ratelimit.RateLimitingEnabledEnv, "false")
|
||||
t.Setenv(ratelimit.RateLimitingBurstEnv, "1")
|
||||
mux, _ := newLocalMux(t)
|
||||
|
||||
require.Equal(t, http.StatusOK, getUploadURL(t, mux))
|
||||
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
|
||||
"github.com/netbirdio/netbird/upload-server/types"
|
||||
)
|
||||
@@ -23,7 +23,7 @@ type sThree struct {
|
||||
presignClient *s3.PresignClient
|
||||
}
|
||||
|
||||
func configureS3Handlers(mux *http.ServeMux, limiter *middleware.APIRateLimiter) error {
|
||||
func configureS3Handlers(mux *http.ServeMux, limiter *ratelimit.APIRateLimiter) error {
|
||||
bucket := os.Getenv(bucketVar)
|
||||
region, ok := os.LookupEnv("AWS_REGION")
|
||||
if !ok {
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"github.com/google/uuid"
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/http/middleware"
|
||||
"github.com/netbirdio/netbird/shared/ratelimit"
|
||||
"github.com/netbirdio/netbird/upload-server/types"
|
||||
)
|
||||
|
||||
@@ -21,7 +21,7 @@ const (
|
||||
|
||||
type Server struct {
|
||||
srv *http.Server
|
||||
limiter *middleware.APIRateLimiter
|
||||
limiter *ratelimit.APIRateLimiter
|
||||
}
|
||||
|
||||
func NewServer() *Server {
|
||||
@@ -63,7 +63,7 @@ func (s *Server) Stop() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func configureMux(mux *http.ServeMux) (*middleware.APIRateLimiter, error) {
|
||||
func configureMux(mux *http.ServeMux) (*ratelimit.APIRateLimiter, error) {
|
||||
limiter := newRateLimiter()
|
||||
|
||||
_, ok := os.LookupEnv(bucketVar)
|
||||
|
||||
Reference in New Issue
Block a user