[management] move rate limiter to shared package (#7727)

This commit is contained in:
Pascal Fischer
2026-09-28 15:44:52 +02:00
committed by GitHub
parent 979571a99f
commit 5d92c89227
13 changed files with 81 additions and 58 deletions
+5 -5
View File
@@ -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
})
+3 -2
View File
@@ -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,
@@ -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 {
@@ -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) {
+2 -2
View File
@@ -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")
+7 -7
View File
@@ -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
}
+8 -8
View File
@@ -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))
+2 -2
View File
@@ -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 {
+3 -3
View File
@@ -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)