add header auth cache to proxy

This commit is contained in:
pascal
2026-08-10 16:47:01 +02:00
parent f9abe2727f
commit ca6a71a0c8
8 changed files with 740 additions and 3 deletions
+64
View File
@@ -13,6 +13,7 @@ import (
"math"
"net"
"net/http"
"net/netip"
"net/url"
"os"
"strconv"
@@ -25,6 +26,7 @@ import (
log "github.com/sirupsen/logrus"
"golang.org/x/oauth2"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/shared/management/domain"
@@ -134,6 +136,10 @@ type ProxyServiceServer struct {
// initial snapshot delivery. Configurable via NB_PROXY_SNAPSHOT_BATCH_SIZE.
snapshotBatchSize int
authAttemptLimiter *authFailureLimiter
authClientLimiter *authFailureLimiter
authFailureMAC []byte
cancel context.CancelFunc
}
@@ -204,6 +210,10 @@ func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeT
snapshotBatchSize: snapshotBatchSizeFromEnv(),
cancel: cancel,
}
s.authAttemptLimiter = newAuthFailureLimiter()
s.authClientLimiter = newAuthClientLimiter()
s.authFailureMAC = make([]byte, sha256.Size)
_, _ = rand.Read(s.authFailureMAC)
go s.cleanupStaleProxies(ctx)
return s
}
@@ -1172,6 +1182,18 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
return nil, err
}
failureKey := s.authFailureKey(req)
limitFailures := failureKey != "" && s.authAttemptLimiter != nil && len(s.authFailureMAC) > 0
if limitFailures && s.authAttemptLimiter.isLimited(failureKey) {
return nil, status.Errorf(codes.ResourceExhausted, "too many failed authentication attempts for this credential, please try again later")
}
clientKey := s.authClientKey(ctx, req.GetId())
limitClient := clientKey != "" && s.authClientLimiter != nil
if limitClient && s.authClientLimiter.isLimited(clientKey) {
return nil, status.Errorf(codes.ResourceExhausted, "too many failed authentication attempts from this client, please try again later")
}
service, err := s.serviceManager.GetServiceByID(ctx, req.GetAccountId(), req.GetId())
if err != nil {
log.WithContext(ctx).Debugf("failed to get service from store: %v", err)
@@ -1179,6 +1201,14 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
}
authenticated, userId, method := s.authenticateRequest(ctx, req, service)
if !authenticated {
if limitFailures {
s.authAttemptLimiter.recordFailure(failureKey)
}
if limitClient {
s.authClientLimiter.recordFailure(clientKey)
}
}
// Non-OIDC schemes (PIN/Password/Header) authenticate against per-service
// secrets and have no user-level group context, so groups stay nil. Email
@@ -1194,6 +1224,40 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
}, nil
}
func (s *ProxyServiceServer) authClientKey(ctx context.Context, serviceID string) string {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return ""
}
values := md.Get(proxyauth.ClientIPMetadataKey)
if len(values) == 0 {
return ""
}
addr, err := netip.ParseAddr(strings.TrimSpace(values[0]))
if err != nil {
return ""
}
return serviceID + "|" + addr.Unmap().String()
}
func (s *ProxyServiceServer) authFailureKey(req *proto.AuthenticateRequest) string {
var secret string
switch v := req.GetRequest().(type) {
case *proto.AuthenticateRequest_Pin:
secret = "pin|" + v.Pin.GetPin()
case *proto.AuthenticateRequest_Password:
secret = "password|" + v.Password.GetPassword()
case *proto.AuthenticateRequest_HeaderAuth:
secret = "header|" + v.HeaderAuth.GetHeaderName() + "|" + v.HeaderAuth.GetHeaderValue()
default:
return ""
}
mac := hmac.New(sha256.New, s.authFailureMAC)
mac.Write([]byte(secret))
return req.GetId() + "|" + hex.EncodeToString(mac.Sum(nil))
}
func (s *ProxyServiceServer) authenticateRequest(ctx context.Context, req *proto.AuthenticateRequest, service *rpservice.Service) (bool, string, proxyauth.Method) {
switch v := req.GetRequest().(type) {
case *proto.AuthenticateRequest_Pin:
@@ -0,0 +1,189 @@
package grpc
import (
"context"
"crypto/rand"
"crypto/sha256"
"fmt"
"testing"
"time"
"github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/time/rate"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
proxyauth "github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/shared/hash/argon2id"
"github.com/netbirdio/netbird/shared/management/proto"
)
const authAttemptsHeaderName = "X-API-Key"
func newAuthAttemptsTestServer(t *testing.T) *ProxyServiceServer {
t.Helper()
firstHash, err := argon2id.Hash("first-key")
require.NoError(t, err)
secondHash, err := argon2id.Hash("second-key")
require.NoError(t, err)
svc := &rpservice.Service{
ID: "svc1",
Domain: "example.com",
Auth: rpservice.AuthConfig{
HeaderAuths: []*rpservice.HeaderAuthConfig{
{Enabled: true, Header: authAttemptsHeaderName, Value: firstHash},
{Enabled: true, Header: authAttemptsHeaderName, Value: secondHash},
},
},
}
ctrl := gomock.NewController(t)
mgr := rpservice.NewMockManager(ctrl)
mgr.EXPECT().GetServiceByID(gomock.Any(), gomock.Any(), gomock.Any()).Return(svc, nil).AnyTimes()
limiter := newAuthFailureLimiter()
t.Cleanup(limiter.stop)
clientLimiter := newAuthClientLimiter()
t.Cleanup(clientLimiter.stop)
mac := make([]byte, sha256.Size)
_, err = rand.Read(mac)
require.NoError(t, err)
return &ProxyServiceServer{
serviceManager: mgr,
authAttemptLimiter: limiter,
authClientLimiter: clientLimiter,
authFailureMAC: mac,
}
}
func clientIPContext(ip string) context.Context {
return metadata.NewIncomingContext(context.Background(), metadata.Pairs(proxyauth.ClientIPMetadataKey, ip))
}
func authAttemptsRequest(credential string) *proto.AuthenticateRequest {
return &proto.AuthenticateRequest{
Id: "svc1",
AccountId: "acc1",
Request: &proto.AuthenticateRequest_HeaderAuth{
HeaderAuth: &proto.HeaderAuthRequest{
HeaderName: authAttemptsHeaderName,
HeaderValue: credential,
},
},
}
}
func TestAuthenticate_ValidCredentialIsNeverRateLimited(t *testing.T) {
s := newAuthAttemptsTestServer(t)
for i := 0; i < proxyAuthFailureBurst*2; i++ {
resp, err := s.Authenticate(context.Background(), authAttemptsRequest("first-key"))
require.NoError(t, err, "a valid credential must never be throttled (attempt %d)", i)
require.True(t, resp.GetSuccess())
}
}
func TestAuthenticate_FailedCredentialIsRateLimited(t *testing.T) {
s := newAuthAttemptsTestServer(t)
for i := 0; i < proxyAuthFailureBurst; i++ {
resp, err := s.Authenticate(context.Background(), authAttemptsRequest("wrong-key"))
require.NoError(t, err, "attempt %d should be within the failure budget", i)
require.False(t, resp.GetSuccess())
}
_, err := s.Authenticate(context.Background(), authAttemptsRequest("wrong-key"))
require.Error(t, err)
assert.Equal(t, codes.ResourceExhausted, status.Code(err))
}
func TestAuthenticate_ThrottledCredentialDoesNotAffectOthers(t *testing.T) {
s := newAuthAttemptsTestServer(t)
for i := 0; i < proxyAuthFailureBurst+2; i++ {
_, _ = s.Authenticate(context.Background(), authAttemptsRequest("wrong-key"))
}
resp, err := s.Authenticate(context.Background(), authAttemptsRequest("first-key"))
require.NoError(t, err, "one throttled credential must not block a valid one")
assert.True(t, resp.GetSuccess())
resp, err = s.Authenticate(context.Background(), authAttemptsRequest("second-key"))
require.NoError(t, err)
assert.True(t, resp.GetSuccess())
_, err = s.Authenticate(context.Background(), authAttemptsRequest("another-wrong-key"))
require.NoError(t, err, "a different failing credential has its own budget")
}
func TestAuthenticate_DistinctCredentialsThrottledPerClient(t *testing.T) {
const budget = 3
s := newAuthAttemptsTestServer(t)
s.authClientLimiter.stop()
s.authClientLimiter = newAuthLimiter(rate.Every(time.Hour), budget)
t.Cleanup(s.authClientLimiter.stop)
ctx := clientIPContext("198.51.100.7")
for i := 0; i < budget; i++ {
resp, err := s.Authenticate(ctx, authAttemptsRequest(fmt.Sprintf("garbage-%d", i)))
require.NoError(t, err, "attempt %d should be within the client budget", i)
require.False(t, resp.GetSuccess())
}
_, err := s.Authenticate(ctx, authAttemptsRequest("garbage-final"))
require.Error(t, err, "a client rotating distinct credentials must be throttled")
assert.Equal(t, codes.ResourceExhausted, status.Code(err))
other := clientIPContext("198.51.100.8")
resp, err := s.Authenticate(other, authAttemptsRequest("first-key"))
require.NoError(t, err, "a different client must be unaffected")
assert.True(t, resp.GetSuccess())
}
func TestAuthenticate_OneStaleCredentialDoesNotExhaustSharedClientBudget(t *testing.T) {
s := newAuthAttemptsTestServer(t)
ctx := clientIPContext("198.51.100.9")
for i := 0; i < proxyAuthFailureBurst*4; i++ {
_, _ = s.Authenticate(ctx, authAttemptsRequest("stale-key"))
}
resp, err := s.Authenticate(ctx, authAttemptsRequest("first-key"))
require.NoError(t, err, "one client stuck on a stale key must not block others behind the same NAT")
assert.True(t, resp.GetSuccess())
}
func TestAuthenticate_ProxyWithoutClientIPIsNotClientLimited(t *testing.T) {
s := newAuthAttemptsTestServer(t)
for i := 0; i < proxyAuthFailureBurst*2; i++ {
_, _ = s.Authenticate(context.Background(), authAttemptsRequest(fmt.Sprintf("garbage-%d", i)))
}
resp, err := s.Authenticate(context.Background(), authAttemptsRequest("first-key"))
require.NoError(t, err, "an old proxy must not have its clients share one budget")
assert.True(t, resp.GetSuccess())
}
func TestAuthenticate_MalformedClientIPIsIgnored(t *testing.T) {
s := newAuthAttemptsTestServer(t)
ctx := clientIPContext("not-an-ip")
for i := 0; i < proxyAuthFailureBurst*2; i++ {
_, _ = s.Authenticate(ctx, authAttemptsRequest(fmt.Sprintf("garbage-%d", i)))
}
resp, err := s.Authenticate(ctx, authAttemptsRequest("first-key"))
require.NoError(t, err)
assert.True(t, resp.GetSuccess())
}
@@ -18,12 +18,16 @@ const (
proxyAuthLimiterCleanup = 5 * time.Minute
// proxyAuthLimiterTTL is how long a limiter is kept after the last failure.
proxyAuthLimiterTTL = 15 * time.Minute
proxyAuthClientBurst = 30
)
// defaultProxyAuthFailureRate is the token replenishment rate for failed auth attempts.
// One token every 12 seconds = 5 per minute.
var defaultProxyAuthFailureRate = rate.Every(12 * time.Second)
var defaultProxyAuthClientRate = rate.Limit(1)
// clientIP identifies a client by its IP address for rate limiting purposes.
type clientIP = string
@@ -37,6 +41,7 @@ type authFailureLimiter struct {
mu sync.Mutex
limiters map[clientIP]*limiterEntry
failureRate rate.Limit
burst int
cancel context.CancelFunc
}
@@ -45,10 +50,19 @@ func newAuthFailureLimiter() *authFailureLimiter {
}
func newAuthFailureLimiterWithRate(failureRate rate.Limit) *authFailureLimiter {
return newAuthLimiter(failureRate, proxyAuthFailureBurst)
}
func newAuthClientLimiter() *authFailureLimiter {
return newAuthLimiter(defaultProxyAuthClientRate, proxyAuthClientBurst)
}
func newAuthLimiter(failureRate rate.Limit, burst int) *authFailureLimiter {
ctx, cancel := context.WithCancel(context.Background())
l := &authFailureLimiter{
limiters: make(map[clientIP]*limiterEntry),
failureRate: failureRate,
burst: burst,
cancel: cancel,
}
go l.cleanupLoop(ctx)
@@ -77,7 +91,7 @@ func (l *authFailureLimiter) recordFailure(ip clientIP) {
entry, exists := l.limiters[ip]
if !exists {
entry = &limiterEntry{
limiter: rate.NewLimiter(l.failureRate, proxyAuthFailureBurst),
limiter: rate.NewLimiter(l.failureRate, l.burst),
}
l.limiters[ip] = entry
}