Merge remote-tracking branch 'origin/main' into fix_debug_upload_url_from_mgmt

# Conflicts:
#	management/server/activity/codes.go
#	management/server/store/sql_store.go
#	management/server/store/sql_store_test.go
#	upload-server/server/server.go
This commit is contained in:
riccardom
2026-09-28 10:56:27 +02:00
320 changed files with 25425 additions and 12010 deletions
@@ -0,0 +1,135 @@
package grpc
import (
"context"
"time"
"github.com/netbirdio/netbird/encryption"
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/shared/management/proto"
log "github.com/sirupsen/logrus"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
func PeerUpdateHandlerFactory(
peerKey wgtypes.Key,
updates chan *network_map.UpdateMessage,
secretsManager SecretsManager,
srv proto.ManagementService_SyncServer,
cleanupfunc func()) *PeerUpdateHandler {
return &PeerUpdateHandler{
peerKey: peerKey,
updates: updates,
secretsManager: secretsManager,
srv: srv,
encrypter: encryption.DefaultEncrypter{},
debouncer: NewUpdateDebouncer(1000 * time.Millisecond),
cleanupFunc: cleanupfunc,
}
}
// PeerUpdateHandler sends updates to the connected peer until the updates channel is closed.
// It implements a backpressure mechanism that sends the first update immediately,
// then debounces subsequent rapid updates, ensuring only the latest update is sent
// after a quiet period.
type PeerUpdateHandler struct {
peerKey wgtypes.Key
updates chan *network_map.UpdateMessage
appMetrics telemetry.AppMetrics
secretsManager SecretsManager
srv syncSender
encrypter encryption.Encrypter
debouncer Debouncer
cleanupFunc func()
}
func (pu *PeerUpdateHandler) WithMetrics(appMetrics telemetry.AppMetrics) *PeerUpdateHandler {
pu.appMetrics = appMetrics
return pu
}
//go:generate go tool mockgen -source=./peer_update_handler.go -destination=./sync_sender_mock.go -package=grpc
type syncSender interface {
Send(*proto.EncryptedMessage) error
Context() context.Context
}
func (pu *PeerUpdateHandler) HandleUpdates(ctx context.Context) error {
log.WithContext(ctx).Tracef("starting to handle updates for peer %s", pu.peerKey.String())
defer pu.debouncer.Stop()
for {
select {
// condition when there are some updates
// todo set the updates channel size to 1
case update, open := <-pu.updates:
if pu.appMetrics != nil {
pu.appMetrics.GRPCMetrics().UpdateChannelQueueLength(len(pu.updates) + 1)
}
if !open {
log.WithContext(ctx).Debugf("updates channel for peer %s was closed", pu.peerKey.String())
pu.cleanupFunc()
return nil
}
log.WithContext(ctx).Tracef("received an update for peer %s", pu.peerKey.String())
if pu.debouncer.ProcessUpdate(update) {
// Send immediately (first update or after quiet period)
if err := pu.SendUpdate(ctx, update); err != nil {
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", pu.peerKey.String(), err)
return err
}
}
// Timer expired - quiet period reached, send pending updates if any
case <-pu.debouncer.TimerChannel():
pendingUpdates := pu.debouncer.GetPendingUpdates()
if len(pendingUpdates) == 0 {
continue
}
log.WithContext(ctx).Debugf("sending %d debounced update(s) for peer %s", len(pendingUpdates), pu.peerKey.String())
for _, pendingUpdate := range pendingUpdates {
if err := pu.SendUpdate(ctx, pendingUpdate); err != nil {
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", pu.peerKey.String(), err)
return err
}
}
// condition when client <-> server connection has been terminated
case <-pu.srv.Context().Done():
// happens when connection drops, e.g. client disconnects
log.WithContext(ctx).Debugf("stream of peer %s has been closed", pu.peerKey.String())
pu.cleanupFunc()
return pu.srv.Context().Err()
}
}
}
func (pu *PeerUpdateHandler) SendUpdate(ctx context.Context, update *network_map.UpdateMessage) error {
key, err := pu.secretsManager.GetWGKey()
if err != nil {
pu.cleanupFunc()
return status.Errorf(codes.Internal, "failed processing update message")
}
encryptedResp, err := pu.encrypter.EncryptMessage(pu.peerKey, key, update.Update)
if err != nil {
pu.cleanupFunc()
return status.Errorf(codes.Internal, "failed processing update message")
}
err = pu.srv.Send(&proto.EncryptedMessage{
WgPubKey: key.PublicKey().String(),
Body: encryptedResp,
})
if err != nil {
pu.cleanupFunc()
return status.Errorf(codes.Internal, "failed sending update message")
}
log.WithContext(ctx).Tracef("sent an update to peer %s", pu.peerKey.String())
return nil
}
@@ -0,0 +1,155 @@
package grpc
import (
"context"
"fmt"
"sync"
"testing"
"time"
pb "github.com/golang/protobuf/proto" //nolint
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/stretchr/testify/assert"
"go.uber.org/mock/gomock"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
func TestSendPeerUpdates_FirstUpdate(t *testing.T) {
ctrl := gomock.NewController(t)
secretsManager := NewMockSecretsManager(ctrl)
updateDebouncer := NewMockDebouncer(ctrl)
syncSender := NewMocksyncSender(ctrl)
pu := PeerUpdateHandler{
peerKey: mustGenerateKey(t),
updates: make(chan *network_map.UpdateMessage),
secretsManager: secretsManager,
encrypter: testEncrypter{},
debouncer: updateDebouncer,
srv: syncSender,
cleanupFunc: func() {},
}
msg := network_map.UpdateMessage{
Update: &proto.SyncResponse{Version: 1},
}
timeCh := make(chan time.Time)
srvCtx := context.TODO()
srvKey := mustGenerateKey(t)
// mock a first update, should send it right away
updateDebouncer.EXPECT().ProcessUpdate(gomock.Eq(&msg)).Return(true)
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
secretsManager.EXPECT().GetWGKey().Return(srvKey, nil)
syncSender.EXPECT().Send(pbMatcher{x: &proto.EncryptedMessage{WgPubKey: srvKey.PublicKey().String(), Body: mustMarshal(t, &msg)}})
updateDebouncer.EXPECT().Stop()
var wg sync.WaitGroup
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
pu.updates <- &msg
close(pu.updates)
wg.Wait()
}
func TestSendPeerUpdates_TimerUpdate(t *testing.T) {
ctrl := gomock.NewController(t)
secretsManager := NewMockSecretsManager(ctrl)
updateDebouncer := NewMockDebouncer(ctrl)
syncSender := NewMocksyncSender(ctrl)
pu := PeerUpdateHandler{
peerKey: mustGenerateKey(t),
updates: make(chan *network_map.UpdateMessage),
secretsManager: secretsManager,
encrypter: testEncrypter{},
debouncer: updateDebouncer,
srv: syncSender,
cleanupFunc: func() {},
}
msg := network_map.UpdateMessage{
Update: &proto.SyncResponse{Version: 1},
}
timeCh := make(chan time.Time)
srvCtx := context.TODO()
srvKey := mustGenerateKey(t)
updateDebouncer.EXPECT().GetPendingUpdates().Return([]*network_map.UpdateMessage{&msg})
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
secretsManager.EXPECT().GetWGKey().Return(srvKey, nil)
syncSender.EXPECT().Send(pbMatcher{x: &proto.EncryptedMessage{WgPubKey: srvKey.PublicKey().String(), Body: mustMarshal(t, &msg)}})
updateDebouncer.EXPECT().Stop()
var wg sync.WaitGroup
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
timeCh <- time.Now()
close(pu.updates)
wg.Wait()
}
func TestSendPeerUpdates_ServerContextDone(t *testing.T) {
ctrl := gomock.NewController(t)
secretsManager := NewMockSecretsManager(ctrl)
updateDebouncer := NewMockDebouncer(ctrl)
syncSender := NewMocksyncSender(ctrl)
pu := PeerUpdateHandler{
peerKey: mustGenerateKey(t),
updates: make(chan *network_map.UpdateMessage),
secretsManager: secretsManager,
encrypter: testEncrypter{},
debouncer: updateDebouncer,
srv: syncSender,
cleanupFunc: func() {},
}
timeCh := make(chan time.Time)
srvCtx, cancel := context.WithCancel(context.TODO())
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
updateDebouncer.EXPECT().Stop()
var wg sync.WaitGroup
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
cancel()
wg.Wait()
}
func mustGenerateKey(t *testing.T) wgtypes.Key {
t.Helper()
k, err := wgtypes.GenerateKey()
assert.NoError(t, err)
return k
}
func mustMarshal(t *testing.T, msg *network_map.UpdateMessage) []byte {
t.Helper()
r, err := pb.Marshal(msg.Update)
assert.NoError(t, err)
return r
}
type testEncrypter struct{}
func (testEncrypter) EncryptMessage(remotePubKey wgtypes.Key, ourPrivateKey wgtypes.Key, message pb.Message) ([]byte, error) {
return pb.Marshal(message)
}
type pbMatcher struct {
x pb.Message
}
func (pbm pbMatcher) Matches(x any) bool {
msg, ok := x.(pb.Message)
if !ok {
return false
}
return pb.Equal(pbm.x, msg)
}
func (pbm pbMatcher) String() string {
return fmt.Sprintf("is equal to %s (%T)", pbm.x, pbm.x)
}
+17 -3
View File
@@ -102,7 +102,8 @@ type ProxyServiceServer struct {
mu sync.RWMutex
// Manager for reverse proxy operations
serviceManager rpservice.Manager
serviceManager rpservice.Manager
credentialLimits credentialVerificationLimiter
// agentNetworkSynth produces synthesised reverse-proxy services from
// Agent Network state. Optional — when nil the snapshot path only ships
// persisted services.
@@ -242,9 +243,10 @@ func (s *ProxyServiceServer) cleanupStaleProxies(ctx context.Context) {
}
}
// Close stops background goroutines.
// Close stops background goroutines and releases credential verification state.
func (s *ProxyServiceServer) Close() {
s.cancel()
s.credentialLimits.close()
}
// SetServiceManager sets the service manager. Must be called before serving.
@@ -412,6 +414,7 @@ func (s *ProxyServiceServer) SetProxyController(proxyController proxy.Controller
type proxyConnectParams struct {
proxyID string
address string
version string
capabilities *proto.ProxyCapabilities
}
@@ -422,6 +425,7 @@ func (s *ProxyServiceServer) GetMappingUpdate(req *proto.GetMappingUpdateRequest
return err
}
params.capabilities = req.GetCapabilities()
params.version = req.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
stream: stream,
@@ -455,6 +459,7 @@ func (s *ProxyServiceServer) SyncMappings(stream proto.ProxyService_SyncMappings
return err
}
params.capabilities = init.GetCapabilities()
params.version = init.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
syncStream: stream,
@@ -566,7 +571,7 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params
}
}
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, accountID, caps)
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, params.version, accountID, caps)
if err != nil {
cancel()
if accountID != nil {
@@ -1223,6 +1228,7 @@ func shallowCloneMapping(m *proto.ProxyMapping) *proto.ProxyMapping {
}
}
// Authenticate verifies service credentials and issues a session token.
func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
if err := enforceAccountScope(ctx, req.GetAccountId()); err != nil {
return nil, err
@@ -1234,6 +1240,14 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
return nil, status.Errorf(codes.FailedPrecondition, "get service from store: %v", err)
}
switch req.GetRequest().(type) {
case *proto.AuthenticateRequest_Pin, *proto.AuthenticateRequest_Password:
key := credentialVerificationKey{accountID: credentialAccountID(service.AccountID), serviceID: credentialServiceID(service.ID)}
if err := s.credentialLimits.allow(key); err != nil {
return nil, err
}
}
authenticated, userId, method := s.authenticateRequest(ctx, req, service)
// Non-OIDC schemes (PIN/Password/Header) authenticate against per-service
@@ -0,0 +1,93 @@
package grpc
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/shared/management/proto"
)
const (
versionTestProxyID = "proxy-a"
versionTestCluster = "cluster.example.com"
versionTestVersion = "0.60.0"
)
// hangupStream cancels its context on the first Send, emulating a proxy that
// disconnects right after receiving the initial snapshot. The legacy stream
// carries no proxy-to-management messages, so this is the only way for
// GetMappingUpdate to return.
type hangupStream struct {
recordingStream
ctx context.Context
cancel context.CancelFunc
}
func (s *hangupStream) Send(m *proto.GetMappingUpdateResponse) error {
s.cancel()
return s.recordingStream.Send(m)
}
func (s *hangupStream) Context() context.Context { return s.ctx }
// newVersionTestServer wires a server whose proxy manager only accepts a
// Connect carrying versionTestVersion, so a dropped or mangled version fails
// the test as an unexpected call.
func newVersionTestServer(t *testing.T) *ProxyServiceServer {
t.Helper()
ctrl := gomock.NewController(t)
svcMgr := rpservice.NewMockManager(ctrl)
svcMgr.EXPECT().GetGlobalServices(gomock.Any()).Return(nil, nil)
proxyMgr := proxy.NewMockManager(ctrl)
proxyMgr.EXPECT().
Connect(gomock.Any(), versionTestProxyID, gomock.Any(), versionTestCluster, gomock.Any(), versionTestVersion, gomock.Any(), gomock.Any()).
Return(&proxy.Proxy{ID: versionTestProxyID, Version: versionTestVersion}, nil)
proxyMgr.EXPECT().Disconnect(gomock.Any(), versionTestProxyID, gomock.Any()).Return(nil)
s := newSnapshotTestServer(t, 10)
s.serviceManager = svcMgr
s.proxyManager = proxyMgr
return s
}
func TestSyncMappings_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
// The init carries the version, the ack acknowledges the empty snapshot,
// and the exhausted fake stream then ends the RPC.
stream := &syncRecordingStream{
recvMsgs: []*proto.SyncMappingsRequest{
{Msg: &proto.SyncMappingsRequest_Init{Init: &proto.SyncMappingsInit{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}}},
ackMsg(),
},
}
err := s.SyncMappings(stream)
require.ErrorContains(t, err, "no more recv messages")
}
func TestGetMappingUpdate_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
stream := &hangupStream{ctx: ctx, cancel: cancel}
err := s.GetMappingUpdate(&proto.GetMappingUpdateRequest{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}, stream)
require.ErrorIs(t, err, context.Canceled)
}
@@ -0,0 +1,101 @@
package grpc
import (
"sync"
"time"
"golang.org/x/time/rate"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/types/known/durationpb"
)
const (
credentialVerificationInterval = 6 * time.Second
credentialVerificationBurst = 5
credentialVerificationMaxServices = 4096
credentialVerificationIdleTimeout = 15 * time.Minute
credentialVerificationCleanupInterval = time.Minute
)
type credentialAccountID string
type credentialServiceID string
type credentialVerificationKey struct {
accountID credentialAccountID
serviceID credentialServiceID
}
type credentialVerificationBudget struct {
limiter *rate.Limiter
lastUsed time.Time
}
// The zero value is ready to use. Budgets are local to this Management process;
// proxy replicas reaching this process share a service's verification budget.
type credentialVerificationLimiter struct {
mu sync.Mutex
now func() time.Time
services map[credentialVerificationKey]*credentialVerificationBudget
nextCleanup time.Time
closed bool
}
func (l *credentialVerificationLimiter) allow(key credentialVerificationKey) error {
l.mu.Lock()
defer l.mu.Unlock()
if l.closed {
return status.Error(codes.Unavailable, "credential verification is closed")
}
now := time.Now()
if l.now != nil {
now = l.now()
}
l.cleanup(now)
budget := l.services[key]
if budget == nil {
if len(l.services) >= credentialVerificationMaxServices {
return credentialVerificationThrottled(credentialVerificationCleanupInterval)
}
if l.services == nil {
l.services = make(map[credentialVerificationKey]*credentialVerificationBudget)
}
budget = &credentialVerificationBudget{limiter: rate.NewLimiter(rate.Every(credentialVerificationInterval), credentialVerificationBurst)}
l.services[key] = budget
}
budget.lastUsed = now
if budget.limiter.AllowN(now, 1) {
return nil
}
delay := max(time.Nanosecond, time.Duration((1-budget.limiter.TokensAt(now))*float64(credentialVerificationInterval)))
return credentialVerificationThrottled(delay)
}
func (l *credentialVerificationLimiter) cleanup(now time.Time) {
if now.Before(l.nextCleanup) {
return
}
l.nextCleanup = now.Add(credentialVerificationCleanupInterval)
for key, budget := range l.services {
if now.Sub(budget.lastUsed) >= credentialVerificationIdleTimeout {
delete(l.services, key)
}
}
}
func (l *credentialVerificationLimiter) close() {
l.mu.Lock()
defer l.mu.Unlock()
l.closed = true
l.services = nil
}
func credentialVerificationThrottled(delay time.Duration) error {
s := status.New(codes.ResourceExhausted, "too many credential verification attempts")
withRetry, err := s.WithDetails(&errdetails.RetryInfo{RetryDelay: durationpb.New(delay)})
if err != nil {
return s.Err()
}
return withRetry.Err()
}
@@ -0,0 +1,79 @@
package grpc
import (
"strconv"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
func TestCredentialVerificationRefillAndIsolation(t *testing.T) {
now := time.Now()
l := credentialVerificationLimiter{now: func() time.Time { return now }}
key := credentialVerificationKey{accountID: "account", serviceID: "service"}
for range credentialVerificationBurst {
require.NoError(t, l.allow(key))
}
err := l.allow(key)
require.Equal(t, codes.ResourceExhausted, status.Code(err), "the burst must be bounded")
now = now.Add(3 * time.Second)
err = l.allow(key)
require.Equal(t, codes.ResourceExhausted, status.Code(err), "a partially refilled token must not permit a check")
details := status.Convert(err).Details()
require.Len(t, details, 1, "throttling must provide RetryInfo")
retry, ok := details[0].(*errdetails.RetryInfo)
require.True(t, ok, "retry details must use the standard message")
assert.Equal(t, 3*time.Second, retry.RetryDelay.AsDuration(), "retry hint must reflect time until the next check")
now = now.Add(3 * time.Second)
require.NoError(t, l.allow(key))
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "only one check must refill every six seconds")
require.NoError(t, l.allow(credentialVerificationKey{accountID: "other-account", serviceID: key.serviceID}))
require.NoError(t, l.allow(credentialVerificationKey{accountID: key.accountID, serviceID: "other-service"}))
}
func TestCredentialVerificationCapacityAndExpiry(t *testing.T) {
now := time.Now()
l := credentialVerificationLimiter{now: func() time.Time { return now }}
for i := range credentialVerificationMaxServices {
require.NoError(t, l.allow(credentialVerificationKey{accountID: "account", serviceID: credentialServiceID(strconv.Itoa(i))}))
}
key := credentialVerificationKey{accountID: "account", serviceID: "new-service"}
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "capacity exhaustion must deny new checks")
now = now.Add(credentialVerificationIdleTimeout)
for range credentialVerificationBurst {
require.NoError(t, l.allow(key))
}
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "expiry must retain the normal burst bound")
}
func TestCredentialVerificationConcurrentChecksAndClose(t *testing.T) {
var l credentialVerificationLimiter
key := credentialVerificationKey{accountID: "account", serviceID: "service"}
var admitted atomic.Int32
var wg sync.WaitGroup
for range 100 {
wg.Go(func() {
if err := l.allow(key); err == nil {
admitted.Add(1)
} else {
assert.Equal(t, codes.ResourceExhausted, status.Code(err), "excess checks must be throttled")
}
})
}
wg.Wait()
assert.EqualValues(t, credentialVerificationBurst, admitted.Load(), "concurrent checks must share the burst")
for range 10 {
wg.Go(l.close)
wg.Go(func() { assert.Error(t, l.allow(key)) })
}
wg.Wait()
assert.Empty(t, l.services, "closing must release retained budgets")
assert.Equal(t, codes.Unavailable, status.Code(l.allow(key)), "checks after close must fail closed")
}
@@ -0,0 +1,18 @@
# Reverse proxy credential verification
The `ProxyService.Authenticate` RPC limits PIN and password checks before
verifying their Argon2 hashes. Both methods share one budget per account and
service: a burst of five checks, replenishing one check every six seconds
(ten per minute). Successful and failed checks consume the budget. Account
scope and service lookup run before the limiter.
Excess checks receive gRPC `ResourceExhausted` with a standard `RetryInfo` delay.
Updated proxies translate it to HTTP 429 and `Retry-After`. Older proxies show
an authentication-service error but cannot bypass the Management limit.
Budgets are held in memory per Management process and reset on restart. Proxy
replicas reaching the same Management process share its budgets. Multiple
Management processes have independent budgets; this is not a cluster-wide
limit. At most 4,096 service budgets are retained, with idle entries expiring
after fifteen minutes. Capacity exhaustion denies new checks until entries
expire. Closing the server releases the retained state.
@@ -0,0 +1,131 @@
package grpc_test
import (
"context"
"net"
"net/netip"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/peer"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
servicemanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service/manager"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/proto"
)
func credentialServer(t *testing.T) (*nbgrpc.ProxyServiceServer, context.Context, grpc.UnaryServerInterceptor) {
t.Helper()
ctx := context.Background()
s, err := store.NewStore(ctx, types.SqliteStoreEngine, t.TempDir(), nil, false)
require.NoError(t, err)
t.Cleanup(func() { assert.NoError(t, s.Close(ctx)) })
require.NoError(t, s.SaveAccount(ctx, &types.Account{Id: "account"}))
keys, err := sessionkey.GenerateKeyPair()
require.NoError(t, err)
for _, id := range []string{"service", "other-service"} {
svc := &service.Service{
ID: id, AccountID: "account", Name: id, Domain: id + ".example.com",
Enabled: true, SessionPrivateKey: keys.PrivateKey, SessionPublicKey: keys.PublicKey,
Auth: service.AuthConfig{
PinAuth: &service.PINAuthConfig{Enabled: true, Pin: "842716"},
PasswordAuth: &service.PasswordAuthConfig{Enabled: true, Password: "test-password"},
},
}
require.NoError(t, svc.Auth.HashSecrets())
require.NoError(t, s.CreateService(ctx, svc))
}
account := "account"
token, err := types.CreateNewProxyAccessToken("test proxy", time.Hour, &account, "admin")
require.NoError(t, err)
require.NoError(t, s.SaveProxyAccessToken(ctx, &token.ProxyAccessToken))
ctx = metadata.NewIncomingContext(ctx, metadata.Pairs("authorization", "Bearer "+string(token.PlainToken)))
ctx = peer.NewContext(ctx, &peer.Peer{Addr: net.TCPAddrFromAddrPort(netip.MustParseAddrPort("192.0.2.1:443"))})
server := nbgrpc.NewProxyServiceServer(nil, nil, nil, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
t.Cleanup(server.Close)
server.SetServiceManager(servicemanager.NewManager(s, nil, nil, nil, nil, nil))
interceptor, _, closeInterceptor := nbgrpc.NewProxyAuthInterceptors(s)
t.Cleanup(closeInterceptor)
return server, ctx, interceptor
}
func TestAuthenticateCredentialRateLimit(t *testing.T) {
server, ctx, interceptor := credentialServer(t)
authenticate := func(req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
response, err := interceptor(ctx, req, &grpc.UnaryServerInfo{FullMethod: "/management.ProxyService/Authenticate"}, func(ctx context.Context, req any) (any, error) {
return server.Authenticate(ctx, req.(*proto.AuthenticateRequest))
})
if err != nil {
return nil, err
}
return response.(*proto.AuthenticateResponse), nil
}
for i := range 5 {
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service"}
if i%2 == 0 {
req.Request = &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}
} else {
req.Request = &proto.AuthenticateRequest_Password{Password: &proto.PasswordRequest{Password: "wrong-password"}}
}
resp, err := authenticate(req)
require.NoError(t, err)
assert.False(t, resp.GetSuccess(), "incorrect PINs and passwords must be denied")
assert.Empty(t, resp.GetSessionToken(), "incorrect credentials must not issue a token")
}
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "842716"}}}
resp, err := authenticate(req)
assert.Nil(t, resp, "a throttled verification must not return a session")
require.Equal(t, codes.ResourceExhausted, status.Code(err), "PIN and password checks must share a service budget even with a valid proxy token")
details := status.Convert(err).Details()
require.Len(t, details, 1, "throttled responses must include a retry hint")
retry, ok := details[0].(*errdetails.RetryInfo)
require.True(t, ok, "the hint must use the standard RetryInfo message")
assert.Positive(t, retry.RetryDelay.AsDuration(), "the retry delay must be positive")
assert.LessOrEqual(t, retry.RetryDelay.AsDuration(), 6*time.Second, "the service must replenish one verification every six seconds")
req.AccountId = "another-account"
_, err = authenticate(req)
assert.Equal(t, codes.PermissionDenied, status.Code(err), "account scope must still be enforced before throttling")
req.AccountId = "account"
req.Id = "other-service"
resp, err = authenticate(req)
require.NoError(t, err)
assert.True(t, resp.GetSuccess(), "one service's throttle must not block another service")
assert.NotEmpty(t, resp.GetSessionToken(), "valid credentials on another service must issue a session")
}
func TestAuthenticateCredentialConcurrentLimit(t *testing.T) {
server, _, _ := credentialServer(t)
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}}
var checked, throttled atomic.Int32
var wg sync.WaitGroup
for range 20 {
wg.Go(func() {
resp, err := server.Authenticate(context.Background(), req)
switch status.Code(err) {
case codes.OK:
checked.Add(1)
assert.False(t, resp.GetSuccess(), "incorrect credentials must be denied")
case codes.ResourceExhausted:
throttled.Add(1)
default:
assert.NoError(t, err)
}
})
}
wg.Wait()
assert.EqualValues(t, 5, checked.Load(), "only the burst budget may reach concurrent credential verification")
assert.EqualValues(t, 15, throttled.Load(), "excess concurrent checks must be throttled")
}
+2 -86
View File
@@ -337,7 +337,8 @@ func (s *Server) Sync(req *proto.EncryptedMessage, srv proto.ManagementService_S
s.syncSem.Add(-1)
return s.handleUpdates(ctx, accountID, peerKey, peer, updates, srv, syncStart)
return PeerUpdateHandlerFactory(peerKey, updates, s.secretsManager, srv, func() { s.cancelPeerRoutines(ctx, accountID, peer, syncStart) }).
WithMetrics(s.appMetrics).HandleUpdates(ctx)
}
func (s *Server) handleHandshake(ctx context.Context, srv proto.ManagementService_JobServer) (wgtypes.Key, error) {
@@ -404,91 +405,6 @@ func (s *Server) sendJobsLoop(ctx context.Context, accountID string, peerKey wgt
}
}
// handleUpdates sends updates to the connected peer until the updates channel is closed.
// It implements a backpressure mechanism that sends the first update immediately,
// then debounces subsequent rapid updates, ensuring only the latest update is sent
// after a quiet period.
func (s *Server) handleUpdates(ctx context.Context, accountID string, peerKey wgtypes.Key, peer *nbpeer.Peer, updates chan *network_map.UpdateMessage, srv proto.ManagementService_SyncServer, streamStartTime time.Time) error {
log.WithContext(ctx).Tracef("starting to handle updates for peer %s", peerKey.String())
// Create a debouncer for this peer connection
debouncer := NewUpdateDebouncer(1000 * time.Millisecond)
defer debouncer.Stop()
for {
select {
// condition when there are some updates
// todo set the updates channel size to 1
case update, open := <-updates:
if s.appMetrics != nil {
s.appMetrics.GRPCMetrics().UpdateChannelQueueLength(len(updates) + 1)
}
if !open {
log.WithContext(ctx).Debugf("updates channel for peer %s was closed", peerKey.String())
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return nil
}
log.WithContext(ctx).Tracef("received an update for peer %s", peerKey.String())
if debouncer.ProcessUpdate(update) {
// Send immediately (first update or after quiet period)
if err := s.sendUpdate(ctx, accountID, peerKey, peer, update, srv, streamStartTime); err != nil {
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", peerKey.String(), err)
return err
}
}
// Timer expired - quiet period reached, send pending updates if any
case <-debouncer.TimerChannel():
pendingUpdates := debouncer.GetPendingUpdates()
if len(pendingUpdates) == 0 {
continue
}
log.WithContext(ctx).Debugf("sending %d debounced update(s) for peer %s", len(pendingUpdates), peerKey.String())
for _, pendingUpdate := range pendingUpdates {
if err := s.sendUpdate(ctx, accountID, peerKey, peer, pendingUpdate, srv, streamStartTime); err != nil {
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", peerKey.String(), err)
return err
}
}
// condition when client <-> server connection has been terminated
case <-srv.Context().Done():
// happens when connection drops, e.g. client disconnects
log.WithContext(ctx).Debugf("stream of peer %s has been closed", peerKey.String())
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return srv.Context().Err()
}
}
}
// sendUpdate encrypts the update message using the peer key and the server's wireguard key,
// then sends the encrypted message to the connected peer via the sync server.
func (s *Server) sendUpdate(ctx context.Context, accountID string, peerKey wgtypes.Key, peer *nbpeer.Peer, update *network_map.UpdateMessage, srv proto.ManagementService_SyncServer, streamStartTime time.Time) error {
key, err := s.secretsManager.GetWGKey()
if err != nil {
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return status.Errorf(codes.Internal, "failed processing update message")
}
encryptedResp, err := encryption.EncryptMessage(peerKey, key, update.Update)
if err != nil {
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return status.Errorf(codes.Internal, "failed processing update message")
}
err = srv.Send(&proto.EncryptedMessage{
WgPubKey: key.PublicKey().String(),
Body: encryptedResp,
})
if err != nil {
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return status.Errorf(codes.Internal, "failed sending update message")
}
log.WithContext(ctx).Tracef("sent an update to peer %s", peerKey.String())
return nil
}
// sendJob encrypts the update message using the peer key and the server's wireguard key,
// then sends the encrypted message to the connected peer via the sync server.
func (s *Server) sendJob(ctx context.Context, peerKey wgtypes.Key, job *job.Event, srv proto.ManagementService_JobServer) error {
@@ -0,0 +1,70 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./peer_update_handler.go
//
// Generated by this command:
//
// mockgen -source=./peer_update_handler.go -destination=./sync_sender_mock.go -package=grpc
//
// Package grpc is a generated GoMock package.
package grpc
import (
context "context"
reflect "reflect"
proto "github.com/netbirdio/netbird/shared/management/proto"
gomock "go.uber.org/mock/gomock"
)
// MocksyncSender is a mock of syncSender interface.
type MocksyncSender struct {
ctrl *gomock.Controller
recorder *MocksyncSenderMockRecorder
isgomock struct{}
}
// MocksyncSenderMockRecorder is the mock recorder for MocksyncSender.
type MocksyncSenderMockRecorder struct {
mock *MocksyncSender
}
// NewMocksyncSender creates a new mock instance.
func NewMocksyncSender(ctrl *gomock.Controller) *MocksyncSender {
mock := &MocksyncSender{ctrl: ctrl}
mock.recorder = &MocksyncSenderMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MocksyncSender) EXPECT() *MocksyncSenderMockRecorder {
return m.recorder
}
// Context mocks base method.
func (m *MocksyncSender) Context() context.Context {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Context")
ret0, _ := ret[0].(context.Context)
return ret0
}
// Context indicates an expected call of Context.
func (mr *MocksyncSenderMockRecorder) Context() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Context", reflect.TypeOf((*MocksyncSender)(nil).Context))
}
// Send mocks base method.
func (m *MocksyncSender) Send(arg0 *proto.EncryptedMessage) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Send", arg0)
ret0, _ := ret[0].(error)
return ret0
}
// Send indicates an expected call of Send.
func (mr *MocksyncSenderMockRecorder) Send(arg0 any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Send", reflect.TypeOf((*MocksyncSender)(nil).Send), arg0)
}
@@ -25,6 +25,8 @@ import (
const defaultDuration = 12 * time.Hour
// SecretsManager used to manage TURN and relay secrets
//
//go:generate go tool mockgen -source=./token_mgr.go -destination=./token_mgr_mock.go -package=grpc
type SecretsManager interface {
GenerateTurnToken() (*Token, error)
GenerateRelayToken() (*Token, error)
@@ -0,0 +1,111 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./token_mgr.go
//
// Generated by this command:
//
// mockgen -source=./token_mgr.go -destination=./token_mgr_mock.go -package=grpc
//
// Package grpc is a generated GoMock package.
package grpc
import (
context "context"
reflect "reflect"
gomock "go.uber.org/mock/gomock"
wgtypes "golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
// MockSecretsManager is a mock of SecretsManager interface.
type MockSecretsManager struct {
ctrl *gomock.Controller
recorder *MockSecretsManagerMockRecorder
isgomock struct{}
}
// MockSecretsManagerMockRecorder is the mock recorder for MockSecretsManager.
type MockSecretsManagerMockRecorder struct {
mock *MockSecretsManager
}
// NewMockSecretsManager creates a new mock instance.
func NewMockSecretsManager(ctrl *gomock.Controller) *MockSecretsManager {
mock := &MockSecretsManager{ctrl: ctrl}
mock.recorder = &MockSecretsManagerMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockSecretsManager) EXPECT() *MockSecretsManagerMockRecorder {
return m.recorder
}
// CancelRefresh mocks base method.
func (m *MockSecretsManager) CancelRefresh(peerKey string) {
m.ctrl.T.Helper()
m.ctrl.Call(m, "CancelRefresh", peerKey)
}
// CancelRefresh indicates an expected call of CancelRefresh.
func (mr *MockSecretsManagerMockRecorder) CancelRefresh(peerKey any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CancelRefresh", reflect.TypeOf((*MockSecretsManager)(nil).CancelRefresh), peerKey)
}
// GenerateRelayToken mocks base method.
func (m *MockSecretsManager) GenerateRelayToken() (*Token, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GenerateRelayToken")
ret0, _ := ret[0].(*Token)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GenerateRelayToken indicates an expected call of GenerateRelayToken.
func (mr *MockSecretsManagerMockRecorder) GenerateRelayToken() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GenerateRelayToken", reflect.TypeOf((*MockSecretsManager)(nil).GenerateRelayToken))
}
// GenerateTurnToken mocks base method.
func (m *MockSecretsManager) GenerateTurnToken() (*Token, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GenerateTurnToken")
ret0, _ := ret[0].(*Token)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GenerateTurnToken indicates an expected call of GenerateTurnToken.
func (mr *MockSecretsManagerMockRecorder) GenerateTurnToken() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GenerateTurnToken", reflect.TypeOf((*MockSecretsManager)(nil).GenerateTurnToken))
}
// GetWGKey mocks base method.
func (m *MockSecretsManager) GetWGKey() (wgtypes.Key, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetWGKey")
ret0, _ := ret[0].(wgtypes.Key)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetWGKey indicates an expected call of GetWGKey.
func (mr *MockSecretsManagerMockRecorder) GetWGKey() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetWGKey", reflect.TypeOf((*MockSecretsManager)(nil).GetWGKey))
}
// SetupRefresh mocks base method.
func (m *MockSecretsManager) SetupRefresh(ctx context.Context, accountID, peerKey string) {
m.ctrl.T.Helper()
m.ctrl.Call(m, "SetupRefresh", ctx, accountID, peerKey)
}
// SetupRefresh indicates an expected call of SetupRefresh.
func (mr *MockSecretsManagerMockRecorder) SetupRefresh(ctx, accountID, peerKey any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetupRefresh", reflect.TypeOf((*MockSecretsManager)(nil).SetupRefresh), ctx, accountID, peerKey)
}
@@ -6,6 +6,14 @@ import (
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
)
//go:generate go tool mockgen -source=./update_debouncer.go -destination=./update_debouncer_mock.go -package=grpc
type Debouncer interface {
Stop()
TimerChannel() <-chan time.Time
ProcessUpdate(update *network_map.UpdateMessage) bool
GetPendingUpdates() []*network_map.UpdateMessage
}
// UpdateDebouncer implements a backpressure mechanism that:
// - Sends the first update immediately
// - Coalesces rapid subsequent network map updates (only latest matters)
@@ -0,0 +1,96 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./update_debouncer.go
//
// Generated by this command:
//
// mockgen -source=./update_debouncer.go -destination=./update_debouncer_mock.go -package=grpc
//
// Package grpc is a generated GoMock package.
package grpc
import (
reflect "reflect"
time "time"
network_map "github.com/netbirdio/netbird/management/internals/controllers/network_map"
gomock "go.uber.org/mock/gomock"
)
// MockDebouncer is a mock of Debouncer interface.
type MockDebouncer struct {
ctrl *gomock.Controller
recorder *MockDebouncerMockRecorder
isgomock struct{}
}
// MockDebouncerMockRecorder is the mock recorder for MockDebouncer.
type MockDebouncerMockRecorder struct {
mock *MockDebouncer
}
// NewMockDebouncer creates a new mock instance.
func NewMockDebouncer(ctrl *gomock.Controller) *MockDebouncer {
mock := &MockDebouncer{ctrl: ctrl}
mock.recorder = &MockDebouncerMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockDebouncer) EXPECT() *MockDebouncerMockRecorder {
return m.recorder
}
// GetPendingUpdates mocks base method.
func (m *MockDebouncer) GetPendingUpdates() []*network_map.UpdateMessage {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetPendingUpdates")
ret0, _ := ret[0].([]*network_map.UpdateMessage)
return ret0
}
// GetPendingUpdates indicates an expected call of GetPendingUpdates.
func (mr *MockDebouncerMockRecorder) GetPendingUpdates() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPendingUpdates", reflect.TypeOf((*MockDebouncer)(nil).GetPendingUpdates))
}
// ProcessUpdate mocks base method.
func (m *MockDebouncer) ProcessUpdate(update *network_map.UpdateMessage) bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ProcessUpdate", update)
ret0, _ := ret[0].(bool)
return ret0
}
// ProcessUpdate indicates an expected call of ProcessUpdate.
func (mr *MockDebouncerMockRecorder) ProcessUpdate(update any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ProcessUpdate", reflect.TypeOf((*MockDebouncer)(nil).ProcessUpdate), update)
}
// Stop mocks base method.
func (m *MockDebouncer) Stop() {
m.ctrl.T.Helper()
m.ctrl.Call(m, "Stop")
}
// Stop indicates an expected call of Stop.
func (mr *MockDebouncerMockRecorder) Stop() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Stop", reflect.TypeOf((*MockDebouncer)(nil).Stop))
}
// TimerChannel mocks base method.
func (m *MockDebouncer) TimerChannel() <-chan time.Time {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "TimerChannel")
ret0, _ := ret[0].(<-chan time.Time)
return ret0
}
// TimerChannel indicates an expected call of TimerChannel.
func (mr *MockDebouncerMockRecorder) TimerChannel() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "TimerChannel", reflect.TypeOf((*MockDebouncer)(nil).TimerChannel))
}
@@ -570,7 +570,7 @@ func (m *testValidateSessionServiceManager) DeleteAccountCluster(_ context.Conte
type testValidateSessionProxyManager struct{}
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
return nil, nil
}