diff --git a/management/internals/server/boot.go b/management/internals/server/boot.go index ea999d82b..dbad2f0c8 100644 --- a/management/internals/server/boot.go +++ b/management/internals/server/boot.go @@ -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 }) diff --git a/management/server/http/handler.go b/management/server/http/handler.go index a57f44b3c..cd365ec13 100644 --- a/management/server/http/handler.go +++ b/management/server/http/handler.go @@ -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) } diff --git a/management/server/http/handlers/proxy/auth.go b/management/server/http/handlers/proxy/auth.go index 0f4b72e14..15327fd30 100644 --- a/management/server/http/handlers/proxy/auth.go +++ b/management/server/http/handlers/proxy/auth.go @@ -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, } } diff --git a/management/server/http/handlers/users/invites_handler.go b/management/server/http/handlers/users/invites_handler.go index 0f0f57c29..d96e58bde 100644 --- a/management/server/http/handlers/users/invites_handler.go +++ b/management/server/http/handlers/users/invites_handler.go @@ -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, diff --git a/management/server/http/middleware/auth_middleware.go b/management/server/http/middleware/auth_middleware.go index ba8c66241..1e831e6c3 100644 --- a/management/server/http/middleware/auth_middleware.go +++ b/management/server/http/middleware/auth_middleware.go @@ -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 { diff --git a/management/server/http/middleware/auth_middleware_test.go b/management/server/http/middleware/auth_middleware_test.go index a34554660..ceefc7c45 100644 --- a/management/server/http/middleware/auth_middleware_test.go +++ b/management/server/http/middleware/auth_middleware_test.go @@ -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, diff --git a/management/server/http/middleware/rate_limiter.go b/shared/ratelimit/rate_limiter.go similarity index 91% rename from management/server/http/middleware/rate_limiter.go rename to shared/ratelimit/rate_limiter.go index 6995e71bb..ff9a123a3 100644 --- a/management/server/http/middleware/rate_limiter.go +++ b/shared/ratelimit/rate_limiter.go @@ -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 { diff --git a/management/server/http/middleware/rate_limiter_test.go b/shared/ratelimit/rate_limiter_test.go similarity index 97% rename from management/server/http/middleware/rate_limiter_test.go rename to shared/ratelimit/rate_limiter_test.go index c647030c7..c43b8e03d 100644 --- a/management/server/http/middleware/rate_limiter_test.go +++ b/shared/ratelimit/rate_limiter_test.go @@ -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) { diff --git a/upload-server/server/local.go b/upload-server/server/local.go index 7db2740f9..859eeb9b2 100644 --- a/upload-server/server/local.go +++ b/upload-server/server/local.go @@ -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") diff --git a/upload-server/server/ratelimit.go b/upload-server/server/ratelimit.go index 081cf89c1..0a38e5cc3 100644 --- a/upload-server/server/ratelimit.go +++ b/upload-server/server/ratelimit.go @@ -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 } diff --git a/upload-server/server/ratelimit_test.go b/upload-server/server/ratelimit_test.go index 8414a1c09..18b5a34fd 100644 --- a/upload-server/server/ratelimit_test.go +++ b/upload-server/server/ratelimit_test.go @@ -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)) diff --git a/upload-server/server/s3.go b/upload-server/server/s3.go index ffc5df01f..55046fb9b 100644 --- a/upload-server/server/s3.go +++ b/upload-server/server/s3.go @@ -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 { diff --git a/upload-server/server/server.go b/upload-server/server/server.go index 607c5bdad..cd9b25c2e 100644 --- a/upload-server/server/server.go +++ b/upload-server/server/server.go @@ -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)