mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-07 05:59:06 +02:00
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:
@@ -245,7 +245,7 @@ func computeMode(t *testing.T, ctx context.Context, mode Mode, nmData *networkma
|
||||
peerGroups := maps.Keys(nmData.GetPeerGroups(peerID))
|
||||
resp := mgmtgrpc.ToComponentSyncResponse(ctx, nil, nil, nil, peer, nil, nil, components, nil,
|
||||
dnsDomain, nil, nmData.AccountSettings, nil, peerGroups, dnsFwdPort)
|
||||
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain)
|
||||
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain, false)
|
||||
require.NoError(t, err, "expand envelope")
|
||||
return res.NetworkMap
|
||||
default:
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
@@ -17,6 +18,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
rpproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
@@ -29,6 +31,9 @@ import (
|
||||
const (
|
||||
testAccountID = "acc-1"
|
||||
testUserID = "user-bob"
|
||||
// testClusterAddress is the shared proxy cluster the settings tests pin
|
||||
// their gateway to; the fixture seeds a connected private-capable proxy for it.
|
||||
testClusterAddress = "eu.proxy.netbird.io"
|
||||
)
|
||||
|
||||
// agentNetworkHandlerFixture builds a real agentnetwork.Manager with
|
||||
@@ -75,6 +80,12 @@ func newAgentNetworkHandlerFixture(t *testing.T) *agentNetworkHandlerFixture {
|
||||
manager := agentnetwork.NewManager(st, perms, accounts, nil)
|
||||
h := &handler{manager: manager}
|
||||
|
||||
// The labeled bootstrap validates its proxy_address against the live
|
||||
// clusters, so seed the shared cluster these tests pin to as a real,
|
||||
// private-capable one — the wire-shape assertions then run through the
|
||||
// validated path rather than the "nothing connected yet" carve-out.
|
||||
seedSharedPrivateCluster(t, st, testClusterAddress)
|
||||
|
||||
router := mux.NewRouter()
|
||||
router.HandleFunc("/agent-network/providers", h.createProvider).Methods("POST")
|
||||
router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET")
|
||||
@@ -268,3 +279,21 @@ func TestConsumptionHandler_PopulatedAccountListsRows(t *testing.T) {
|
||||
assert.Equal(t, groupRow.WindowStartUtc, userRow.WindowStartUtc,
|
||||
"rows recorded in the same window must share the aligned window_start_utc")
|
||||
}
|
||||
|
||||
// seedSharedPrivateCluster registers a connected, NetBird-operated proxy
|
||||
// with private capabilities (the `private` capability) so
|
||||
// clusterAddr is a cluster any account may pin its agent-network gateway to.
|
||||
func seedSharedPrivateCluster(t *testing.T, st store.Store, clusterAddr string) {
|
||||
t.Helper()
|
||||
private := true
|
||||
now := time.Now().UTC()
|
||||
require.NoError(t, st.SaveProxy(context.Background(), &rpproxy.Proxy{
|
||||
ID: "shared-proxy-" + clusterAddr,
|
||||
SessionID: "shared-session",
|
||||
ClusterAddress: clusterAddr,
|
||||
LastSeen: now,
|
||||
ConnectedAt: &now,
|
||||
Status: rpproxy.StatusConnected,
|
||||
Capabilities: rpproxy.Capabilities{Private: &private},
|
||||
}), "seeding the shared proxy cluster must succeed")
|
||||
}
|
||||
|
||||
@@ -1036,6 +1036,18 @@ func (m *managerImpl) bootstrapSelfAddressed(ctx context.Context, settings *type
|
||||
if err != nil {
|
||||
return status.Errorf(status.InvalidArgument, "invalid endpoint: %s", err)
|
||||
}
|
||||
if err := m.requireHostNotForeign(ctx, settings.AccountID, hostname); err != nil {
|
||||
return err
|
||||
}
|
||||
// Another account's labeled pin beneath this hostname makes it their
|
||||
// cluster: a proxy serving them there would never serve this endpoint.
|
||||
// The domain unique index already arbitrates two endpoints on one name.
|
||||
if err := m.requireNotClaimedByOtherAccount(ctx, settings.AccountID, hostname, m.store.HasGatewayClusterPinnedByOtherAccount); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.validateGatewayCluster(ctx, settings.AccountID, hostname); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
settings.Domain = hostname
|
||||
settings.ProxyAddress = hostname
|
||||
@@ -1054,6 +1066,99 @@ func (m *managerImpl) bootstrapSelfAddressed(ctx context.Context, settings *type
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateGatewayCluster rejects a bootstrap pinned to a cluster that cannot
|
||||
// serve the account's gateway — a labeled endpoint beneath the cluster and a
|
||||
// self-addressed one on the very address a proxy declares alike, since the
|
||||
// service behind either is the same private one.
|
||||
//
|
||||
// The synthesised gateway service is unconditionally private
|
||||
// (buildAccountService): agents reach it over the WireGuard tunnel and are
|
||||
// authorised by ValidateTunnelPeer against the policies' source groups, and
|
||||
// its single target is the cluster itself with DirectUpstream. Only a cluster
|
||||
// with private capabilities can serve that. Management reports it per cluster
|
||||
// as the `private` capability, the same flag the dashboard renders as
|
||||
// supports_private when it gates NetBird-only services.
|
||||
//
|
||||
// Without this check the bootstrap happily pins to any cluster the caller
|
||||
// names, including one without private capabilities — and the endpoint it
|
||||
// allocates is immutable, so the account is left with a dead gateway that only
|
||||
// a DeleteSettings/re-bootstrap can undo.
|
||||
//
|
||||
// Whether management knows the cluster is decided on the proxy rows
|
||||
// themselves, never on how fresh their heartbeats are: a cluster's rows
|
||||
// outlive its proxies' liveness (only the stale-proxy reaper removes them), so
|
||||
// a cluster that exists stays judged as one. Judging on liveness instead would
|
||||
// make the same centralised cluster pass or fail depending on whether its
|
||||
// proxies happened to have heartbeated in the last couple of minutes.
|
||||
//
|
||||
// The single opening left is a cluster management holds no proxy row for at
|
||||
// all: pinning ahead of a proxy's first connection is a legitimate order — the
|
||||
// dedicated path claims an address the same way, before any proxy declares it.
|
||||
func (m *managerImpl) validateGatewayCluster(ctx context.Context, accountID, clusterAddr string) error {
|
||||
declared, err := m.accountClusterSpellings(ctx, accountID, clusterAddr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(declared) == 0 {
|
||||
// No proxy has ever declared this address: an address-first pin.
|
||||
return nil
|
||||
}
|
||||
|
||||
// A cluster management knows has to prove it can serve the gateway, and
|
||||
// only a live proxy reporting the capability proves that. Both an explicit false and an
|
||||
// unreported capability (nothing live in the cluster, or proxies predating
|
||||
// capability reporting) fail here: unusable and unproven are the same
|
||||
// answer for a decision that cannot be revisited later.
|
||||
//
|
||||
// The capability is read per declared spelling and taken as any-true, the
|
||||
// same way it aggregates over a cluster's proxies: the store matches
|
||||
// cluster_address exactly, so a host two proxies spelled differently must
|
||||
// not come back unproven just because it was asked about under one of them.
|
||||
for _, address := range declared {
|
||||
if private := m.store.GetClusterSupportsPrivate(ctx, address); private != nil && *private {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return status.Errorf(status.InvalidArgument,
|
||||
"proxy cluster %s has no private capabilities: the agent network gateway requires a reverse proxy cluster "+
|
||||
"with private capabilities", clusterAddr)
|
||||
}
|
||||
|
||||
// accountClusterSpellings returns every proxy cluster address in the account's
|
||||
// view — its own (BYOP) clusters plus the shared ones — that names the same
|
||||
// host as clusterAddr. Empty means management holds no proxy row for that host
|
||||
// in this account's view.
|
||||
//
|
||||
// A proxy declares its cluster address as the operator spelled it, so identity
|
||||
// is compared on the normalised form rather than byte-equal — an in-memory pass
|
||||
// over the account's clusters, not a query. What comes back is the stored
|
||||
// spelling, because the capability lookup matches cluster_address exactly and
|
||||
// would silently find nothing under a spelling the store never held. The
|
||||
// cluster listing is not gated on heartbeats, so this answer does not change
|
||||
// while a cluster's proxies are merely offline.
|
||||
func (m *managerImpl) accountClusterSpellings(ctx context.Context, accountID, clusterAddr string) ([]string, error) {
|
||||
clusters, err := m.store.GetProxyClusters(ctx, accountID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list proxy clusters: %w", err)
|
||||
}
|
||||
|
||||
var spellings []string
|
||||
for _, cluster := range clusters {
|
||||
normalized, err := types.NormalizeHostname(cluster.Address)
|
||||
if err != nil {
|
||||
// An address declared in a shape we cannot normalise is not one an
|
||||
// endpoint can be allocated beneath.
|
||||
log.WithContext(ctx).Debugf("skipping unusable proxy cluster address %q: %s", cluster.Address, err)
|
||||
continue
|
||||
}
|
||||
if normalized == clusterAddr {
|
||||
spellings = append(spellings, cluster.Address)
|
||||
}
|
||||
}
|
||||
return spellings, nil
|
||||
}
|
||||
|
||||
// bootstrapLabeled allocates a labeled endpoint one label beneath the given
|
||||
// cluster address: Domain = <label>.<proxyAddress>, served by whichever proxy
|
||||
// declares the parent. Labels are adjective-noun tuples; a candidate is
|
||||
@@ -1065,6 +1170,20 @@ func (m *managerImpl) bootstrapLabeled(ctx context.Context, settings *types.Sett
|
||||
if err != nil {
|
||||
return status.Errorf(status.InvalidArgument, "invalid proxy_address: %s", err)
|
||||
}
|
||||
if err := m.requireHostNotForeign(ctx, settings.AccountID, parent); err != nil {
|
||||
return err
|
||||
}
|
||||
// Another account's endpoint at this exact hostname means the proxy that
|
||||
// declares it is theirs, so nothing would serve a label beneath it. Other
|
||||
// accounts' labeled pins under the same cluster are not asked about: a
|
||||
// shared cluster carries many of them by design.
|
||||
if err := m.requireNotClaimedByOtherAccount(ctx, settings.AccountID, parent, m.store.HasGatewayEndpointByOtherAccount); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := m.validateGatewayCluster(ctx, settings.AccountID, parent); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for attempt := 1; attempt <= maxDomainAllocationAttempts; attempt++ {
|
||||
label := labelgen.PickTuple()
|
||||
@@ -1111,6 +1230,41 @@ func (m *managerImpl) bootstrapLabeled(ctx context.Context, settings *types.Sett
|
||||
return fmt.Errorf("allocate agent network endpoint for account %s: %d attempts exhausted", settings.AccountID, maxDomainAllocationAttempts)
|
||||
}
|
||||
|
||||
// requireHostNotForeign refuses to pin the account's gateway onto a host that
|
||||
// another account's proxy declares. The pin's proxy_address is what selects
|
||||
// the proxy that serves the endpoint, and an account-scoped proxy only ever
|
||||
// receives its own account's mappings, so such a pin could never be served —
|
||||
// and the endpoint it assigns is immutable. Shared proxies are not foreign, and
|
||||
// a host no proxy has declared stays pinnable: claiming the address before the
|
||||
// proxy's first connection is the documented order.
|
||||
func (m *managerImpl) requireHostNotForeign(ctx context.Context, accountID, host string) error {
|
||||
foreign, err := m.store.HasForeignAccountProxyAtHost(ctx, host, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("check proxy host ownership: %w", err)
|
||||
}
|
||||
if foreign {
|
||||
return errHostNotAvailable(host)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// requireNotClaimedByOtherAccount refuses the pin when another account's
|
||||
// gateway settings already claim the host in the shape claimed answers for.
|
||||
func (m *managerImpl) requireNotClaimedByOtherAccount(ctx context.Context, accountID, host string, claimed func(context.Context, string, string) (bool, error)) error {
|
||||
taken, err := claimed(ctx, host, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("check agent network gateway claims at host: %w", err)
|
||||
}
|
||||
if taken {
|
||||
return errHostNotAvailable(host)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func errHostNotAvailable(host string) error {
|
||||
return status.Errorf(status.InvalidArgument, "proxy cluster %s is not available to this account", host)
|
||||
}
|
||||
|
||||
// isUniqueConstraintError reports whether err is a database unique-constraint
|
||||
// violation, matched on the driver message because CreateAgentNetworkSettings
|
||||
// deliberately returns the driver error unwrapped.
|
||||
|
||||
@@ -5,12 +5,14 @@ import (
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
@@ -70,6 +72,57 @@ func (f *bootstrapFixture) createSettings(ctx context.Context, accountID, userID
|
||||
return f.manager.CreateSettings(ctx, userID, types.DefaultSettings(accountID), proxyAddress, endpoint)
|
||||
}
|
||||
|
||||
func ptrTo[T any](v T) *T { return &v }
|
||||
|
||||
// seedProxy registers a proxy in clusterAddr, heartbeating now, so the labeled
|
||||
// bootstrap path has a real cluster to validate against. accountID empty makes
|
||||
// it a shared (NetBird-operated) cluster; private mirrors the capability an
|
||||
// proxy with private capabilities reports, nil an unreported one.
|
||||
func (f *bootstrapFixture) seedProxy(t *testing.T, proxyID, accountID, clusterAddr string, private *bool) {
|
||||
t.Helper()
|
||||
f.seedProxyAt(t, proxyID, accountID, clusterAddr, private, time.Now().UTC())
|
||||
}
|
||||
|
||||
// seedProxyAt is seedProxy with an explicit last-seen, for cases that need a
|
||||
// proxy whose heartbeat has aged past the active window while its row (and so
|
||||
// its cluster) is still on record.
|
||||
func (f *bootstrapFixture) seedProxyAt(t *testing.T, proxyID, accountID, clusterAddr string, private *bool, lastSeen time.Time) {
|
||||
t.Helper()
|
||||
p := &proxy.Proxy{
|
||||
ID: proxyID,
|
||||
ClusterAddress: clusterAddr,
|
||||
Status: proxy.StatusConnected,
|
||||
LastSeen: lastSeen,
|
||||
Capabilities: proxy.Capabilities{Private: private},
|
||||
}
|
||||
if accountID != "" {
|
||||
p.AccountID = &accountID
|
||||
}
|
||||
require.NoError(t, f.store.SaveProxy(context.Background(), p), "seeding a proxy must succeed")
|
||||
}
|
||||
|
||||
// seedPrivateCluster is the common case: a shared cluster with a connected
|
||||
// proxy that has private capabilities, which is what a bootstrap requires.
|
||||
func (f *bootstrapFixture) seedPrivateCluster(t *testing.T, clusterAddr string) {
|
||||
t.Helper()
|
||||
f.seedProxy(t, "proxy-"+clusterAddr, "", clusterAddr, ptrTo(true))
|
||||
}
|
||||
|
||||
// requireForeignClusterRefusal asserts the refusal a pin onto another
|
||||
// account's host gets, and that it left no row behind.
|
||||
func (f *bootstrapFixture) requireForeignClusterRefusal(t *testing.T, err error, accountID string) {
|
||||
t.Helper()
|
||||
require.Error(t, err, "another account's host must be refused")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.InvalidArgument, sErr.Type(), "rejection must be a validation error")
|
||||
assert.Contains(t, err.Error(), "not available to this account",
|
||||
"the error must say the host is not the account's to use")
|
||||
|
||||
_, err = f.store.GetAgentNetworkSettings(context.Background(), store.LockingStrengthNone, accountID)
|
||||
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
|
||||
}
|
||||
|
||||
// TestCreateSettingsRequiresPermission pins the gate: bootstrap assigns the
|
||||
// account's immutable endpoint, a settings write requiring the settings
|
||||
// Create permission — and a denial leaves no row behind.
|
||||
@@ -94,6 +147,7 @@ func TestCreateSettingsRequiresPermission(t *testing.T) {
|
||||
func TestCreateSettingsLabeled(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedPrivateCluster(t, "cluster1.example.com")
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "Cluster1.Example.com", "")
|
||||
@@ -167,6 +221,7 @@ func TestCreateSettingsIdentityFieldValidation(t *testing.T) {
|
||||
func TestCreateSettingsConflictsOnSecondBootstrap(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedPrivateCluster(t, "cluster1.example.com")
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
first, err := f.createSettings(ctx, "account1", "user1", "cluster1.example.com", "")
|
||||
@@ -230,3 +285,279 @@ func TestCreateProviderHasNoSettingsSideEffects(t *testing.T) {
|
||||
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
assert.Error(t, err, "provider create must not conjure a settings row")
|
||||
}
|
||||
|
||||
// TestCreateSettingsRejectsOfflineCluster is the guard against deciding on
|
||||
// heartbeat freshness. A centralised cluster is refused while its proxies are
|
||||
// live; the same cluster must stay refused once they stop heartbeating, which
|
||||
// takes only a couple of minutes (proxyActiveThreshold). Judging on liveness
|
||||
// would turn "wait for the proxy to go quiet" into a way to pin the account's
|
||||
// immutable endpoint to a cluster that can never serve it.
|
||||
func TestCreateSettingsRejectsOfflineCluster(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
notPrivate := false
|
||||
|
||||
cases := map[string]*bool{
|
||||
"centralised proxy gone quiet": ¬Private,
|
||||
// A cluster that could serve the gateway still has to have something
|
||||
// live in it to prove so at bootstrap: refusing is the safe direction
|
||||
// (reconnect the proxy and retry) where accepting is permanent.
|
||||
"private proxy gone quiet": ptrTo(true),
|
||||
}
|
||||
for name, private := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxyAt(t, "proxy1", "", "offline.example.com", private,
|
||||
time.Now().UTC().Add(-time.Hour))
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
_, err := f.createSettings(ctx, "account1", "user1", "offline.example.com", "")
|
||||
require.Error(t, err, "a known cluster with nothing live in it must be rejected")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.InvalidArgument, sErr.Type(), "rejection must be a validation error")
|
||||
assert.Contains(t, err.Error(), "private capabilities",
|
||||
"the error must say private capabilities are what is missing")
|
||||
|
||||
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCreateSettingsRequiresPrivateCluster pins the capability gate: the
|
||||
// synthesised gateway service is always private, so a live cluster whose
|
||||
// proxies lack private capabilities cannot serve it and must not
|
||||
// become the account's immutable endpoint.
|
||||
func TestCreateSettingsRequiresPrivateCluster(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
notPrivate := false
|
||||
f.seedProxy(t, "proxy1", "", "central.example.com", ¬Private)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
_, err := f.createSettings(ctx, "account1", "user1", "central.example.com", "")
|
||||
require.Error(t, err, "a cluster without private capabilities must be rejected")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.InvalidArgument, sErr.Type(), "rejection must be a validation error")
|
||||
assert.Contains(t, err.Error(), "private capabilities", "the error must name what the cluster is missing")
|
||||
|
||||
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
|
||||
}
|
||||
|
||||
// TestCreateSettingsAcceptsOwnPrivateCluster pins the BYOP happy path: the
|
||||
// account's own cluster with a connected private-capable proxy is a valid pin.
|
||||
func TestCreateSettingsAcceptsOwnPrivateCluster(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxy(t, "proxy1", "account1", "byop.account1.example.com", ptrTo(true))
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "byop.account1.example.com", "")
|
||||
require.NoError(t, err, "the account's own private cluster must be accepted")
|
||||
assert.Equal(t, "byop.account1.example.com", created.ProxyAddress)
|
||||
}
|
||||
|
||||
// TestCreateSettingsMatchesClusterCasing pins that a cluster spelled with
|
||||
// capitals in the store is still recognised as the same cluster the normalised
|
||||
// proxy_address names, in both directions: a private cluster is accepted and a
|
||||
// centralised one is refused, whatever the casing. The comparison is in memory
|
||||
// over the account's cluster list; the capability lookup is still asked under
|
||||
// the spelling the store actually holds, which is what an exact match needs.
|
||||
func TestCreateSettingsMatchesClusterCasing(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("own private cluster is found", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxy(t, "proxy1", "", "EU.Proxy.Example.com", ptrTo(true))
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "eu.proxy.example.com", "")
|
||||
require.NoError(t, err, "a private cluster declared with capitals must still be accepted")
|
||||
assert.Equal(t, "eu.proxy.example.com", created.ProxyAddress)
|
||||
})
|
||||
|
||||
t.Run("non-private cluster is still refused", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxy(t, "proxy1", "", "Central.Example.com", ptrTo(false))
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
_, err := f.createSettings(ctx, "account1", "user1", "central.example.com", "")
|
||||
require.Error(t, err, "casing must not become a way past the capability check")
|
||||
assert.Contains(t, err.Error(), "private capabilities")
|
||||
})
|
||||
}
|
||||
|
||||
// TestCreateSettingsRejectsForeignCluster pins tenant consistency on the pin:
|
||||
// an account may not pin its gateway onto a host another account's proxy
|
||||
// declares. That proxy only ever receives its own account's mappings, so the
|
||||
// pin could never be served, and the endpoint it assigns is immutable.
|
||||
// Ownership is decided on the proxy rows, not on heartbeat freshness — a
|
||||
// cluster whose proxies are merely offline is still somebody's — and on the
|
||||
// normalised host, since proxies declare their address as the operator
|
||||
// spelled it.
|
||||
func TestCreateSettingsRejectsForeignCluster(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
cases := map[string]struct {
|
||||
spelling string
|
||||
lastSeen time.Time
|
||||
}{
|
||||
"live": {"byop.account2.example.com", time.Now().UTC()},
|
||||
"offline": {"byop.account2.example.com", time.Now().UTC().Add(-time.Hour)},
|
||||
"spelled in caps": {"BYOP.Account2.Example.com", time.Now().UTC()},
|
||||
}
|
||||
for name, tc := range cases {
|
||||
t.Run("labeled "+name, func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxyAt(t, "proxy1", "account2", tc.spelling, ptrTo(true), tc.lastSeen)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
_, err := f.createSettings(ctx, "account1", "user1", "byop.account2.example.com", "")
|
||||
f.requireForeignClusterRefusal(t, err, "account1")
|
||||
})
|
||||
t.Run("self-addressed "+name, func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxyAt(t, "proxy1", "account2", tc.spelling, ptrTo(true), tc.lastSeen)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
_, err := f.createSettings(ctx, "account1", "user1", "", "byop.account2.example.com")
|
||||
f.requireForeignClusterRefusal(t, err, "account1")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCreateSettingsSharedClusterStaysPinnable pins the constraint the
|
||||
// ownership check must respect: a shared (NetBird-operated) cluster is not
|
||||
// anybody's, so any number of accounts pin their gateways to it — including
|
||||
// an account that also runs a proxy of its own elsewhere.
|
||||
func TestCreateSettingsSharedClusterStaysPinnable(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxy(t, "shared", "", "eu.proxy.netbird.io", ptrTo(true))
|
||||
f.seedProxy(t, "own", "account1", "byop.account1.example.com", ptrTo(true))
|
||||
|
||||
for _, account := range []string{"account1", "account2"} {
|
||||
f.expectPermission(account, "user", modules.AgentNetworkSettings, operations.Create, true)
|
||||
created, err := f.createSettings(ctx, account, "user", "eu.proxy.netbird.io", "")
|
||||
require.NoError(t, err, "a shared cluster must stay pinnable by %s", account)
|
||||
assert.Equal(t, "eu.proxy.netbird.io", created.ProxyAddress)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCreateSettingsOwnClusterIsPinnable is the BYOP order in both directions:
|
||||
// the account's own proxy is not a competing claim, whether the pin is labeled
|
||||
// beneath its cluster or self-addressed onto the very host it declares.
|
||||
func TestCreateSettingsOwnClusterIsPinnable(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("labeled", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxy(t, "own", "account1", "byop.account1.example.com", ptrTo(true))
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "byop.account1.example.com", "")
|
||||
require.NoError(t, err, "the account's own cluster must be pinnable")
|
||||
assert.True(t, strings.HasSuffix(created.Domain, ".byop.account1.example.com"))
|
||||
})
|
||||
t.Run("self-addressed", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxy(t, "own", "account1", "gw.account1.example.com", ptrTo(true))
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "", "gw.account1.example.com")
|
||||
require.NoError(t, err, "the host the account's own proxy declares must be pinnable")
|
||||
assert.Equal(t, "gw.account1.example.com", created.ProxyAddress)
|
||||
})
|
||||
}
|
||||
|
||||
// TestCreateSettingsUnknownHostIsPinnable pins the address-first order: a host
|
||||
// no proxy has ever declared is nobody's, so the pin goes through and the
|
||||
// proxy is deployed after.
|
||||
func TestCreateSettingsUnknownHostIsPinnable(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "future.example.com", "")
|
||||
require.NoError(t, err, "a host no proxy has declared must stay pinnable")
|
||||
assert.Equal(t, "future.example.com", created.ProxyAddress)
|
||||
}
|
||||
|
||||
// TestCreateSettingsRejectsHostAnotherAccountPinned covers claims made by pins
|
||||
// rather than proxies, which the proxy-row check cannot see. A labeled pin
|
||||
// beneath a host makes that host the other account's cluster, so a
|
||||
// self-addressed endpoint on it would never be served; a self-addressed
|
||||
// endpoint on a host makes the proxy declaring it theirs, so a label beneath
|
||||
// it would never be served either. Neither is a shared-cluster shape: many
|
||||
// labeled pins under one cluster are asked about in neither direction.
|
||||
func TestCreateSettingsRejectsHostAnotherAccountPinned(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("self-addressed onto another account's cluster", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account2", "user2", modules.AgentNetworkSettings, operations.Create, true)
|
||||
_, err := f.createSettings(ctx, "account2", "user2", "gw.example.com", "")
|
||||
require.NoError(t, err, "account2's labeled pin beneath the host must go through first")
|
||||
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
_, err = f.createSettings(ctx, "account1", "user1", "", "gw.example.com")
|
||||
f.requireForeignClusterRefusal(t, err, "account1")
|
||||
})
|
||||
|
||||
t.Run("labeled beneath another account's endpoint", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account2", "user2", modules.AgentNetworkSettings, operations.Create, true)
|
||||
_, err := f.createSettings(ctx, "account2", "user2", "", "gw.example.com")
|
||||
require.NoError(t, err, "account2's self-addressed endpoint must go through first")
|
||||
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
_, err = f.createSettings(ctx, "account1", "user1", "gw.example.com", "")
|
||||
f.requireForeignClusterRefusal(t, err, "account1")
|
||||
})
|
||||
|
||||
t.Run("labeled beside another account's labeled pin stays allowed", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
for _, account := range []string{"account1", "account2"} {
|
||||
f.expectPermission(account, "user", modules.AgentNetworkSettings, operations.Create, true)
|
||||
_, err := f.createSettings(ctx, account, "user", "eu.proxy.netbird.io", "")
|
||||
require.NoError(t, err, "labeled pins under one cluster are the shared-cluster shape and must not refuse each other")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestCreateSettingsSelfAddressedRequiresPrivateCluster pins that the
|
||||
// capability gate applies to a self-addressed endpoint too: the service behind
|
||||
// it is the same private one, so a proxy that already declares the hostname
|
||||
// must have private capabilities, whether the account's own or a shared cluster's. A
|
||||
// hostname no proxy declares yet stays claimable (TestCreateSettingsSelfAddressed).
|
||||
func TestCreateSettingsSelfAddressedRequiresPrivateCluster(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("centralised proxy at the hostname is refused", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxy(t, "central", "", "gw.example.com", ptrTo(false))
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
_, err := f.createSettings(ctx, "account1", "user1", "", "gw.example.com")
|
||||
require.Error(t, err, "a self-addressed endpoint on a centralised proxy can never be served")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.InvalidArgument, sErr.Type())
|
||||
assert.Contains(t, err.Error(), "private capabilities")
|
||||
|
||||
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
|
||||
})
|
||||
|
||||
t.Run("private proxy at the hostname is accepted", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.seedProxy(t, "private", "", "gw.example.com", ptrTo(true))
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "", "gw.example.com")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "gw.example.com", created.ProxyAddress)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,5 +1,13 @@
|
||||
package domain
|
||||
|
||||
import "time"
|
||||
|
||||
// ValidationTTL is the time available to validate a custom domain registration.
|
||||
const ValidationTTL = 48 * time.Hour
|
||||
|
||||
// ID identifies a custom domain registration.
|
||||
type ID string
|
||||
|
||||
type Type string
|
||||
|
||||
const (
|
||||
@@ -8,12 +16,13 @@ const (
|
||||
)
|
||||
|
||||
type Domain struct {
|
||||
ID string `gorm:"unique;primaryKey;autoIncrement"`
|
||||
Domain string `gorm:"unique"` // Domain records must be unique, this avoids domain reuse across accounts.
|
||||
AccountID string `gorm:"index"`
|
||||
TargetCluster string // The proxy cluster this domain should be validated against
|
||||
Type Type `gorm:"-"`
|
||||
Validated bool
|
||||
ID string `gorm:"unique;primaryKey;autoIncrement"`
|
||||
Domain string `gorm:"unique"` // Domain records must be unique, this avoids domain reuse across accounts.
|
||||
AccountID string `gorm:"index"`
|
||||
TargetCluster string // The proxy cluster this domain should be validated against
|
||||
Type Type `gorm:"-"`
|
||||
Validated bool
|
||||
ValidationExpiresAt *time.Time `gorm:"index"`
|
||||
// SupportsCustomPorts is populated at query time for free domains from the
|
||||
// proxy cluster capabilities. Not persisted.
|
||||
SupportsCustomPorts *bool `gorm:"-"`
|
||||
@@ -36,7 +45,12 @@ func (d *Domain) EventMeta() map[string]any {
|
||||
}
|
||||
}
|
||||
|
||||
// Copy returns a copy with an independent validation deadline.
|
||||
func (d *Domain) Copy() *Domain {
|
||||
dCopy := *d
|
||||
if d.ValidationExpiresAt != nil {
|
||||
expiresAt := *d.ValidationExpiresAt
|
||||
dCopy.ValidationExpiresAt = &expiresAt
|
||||
}
|
||||
return &dCopy
|
||||
}
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
)
|
||||
|
||||
const (
|
||||
validationCleanupInterval = 60 * time.Minute
|
||||
validationCleanupBatch = 100
|
||||
)
|
||||
|
||||
// RunValidationCleanup removes expired registrations on startup and hourly until cancellation.
|
||||
func (m Manager) RunValidationCleanup(ctx context.Context) {
|
||||
ticker := time.NewTicker(validationCleanupInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
m.cleanupExpiredDomains(ctx, time.Now().UTC())
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m Manager) cleanupExpiredDomains(ctx context.Context, now time.Time) {
|
||||
var afterID domain.ID
|
||||
for ctx.Err() == nil {
|
||||
domains, err := m.store.GetExpiredCustomDomains(ctx, now, afterID, validationCleanupBatch)
|
||||
if err != nil {
|
||||
if ctx.Err() == nil {
|
||||
log.WithContext(ctx).WithError(err).Error("list expired custom domain registrations")
|
||||
}
|
||||
return
|
||||
}
|
||||
for _, d := range domains {
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
m.deleteExpiredDomain(ctx, d, now)
|
||||
afterID = domain.ID(d.ID)
|
||||
}
|
||||
if len(domains) < validationCleanupBatch {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m Manager) deleteExpiredDomain(ctx context.Context, d *domain.Domain, now time.Time) {
|
||||
deleted, err := m.store.DeleteExpiredCustomDomain(ctx, d, now)
|
||||
if err != nil {
|
||||
if ctx.Err() == nil {
|
||||
log.WithContext(ctx).WithFields(log.Fields{"accountID": d.AccountID, "domainID": d.ID}).
|
||||
WithError(err).Warn("could not expire custom domain registration")
|
||||
}
|
||||
return
|
||||
}
|
||||
if !deleted {
|
||||
return
|
||||
}
|
||||
meta := d.EventMeta()
|
||||
if d.ValidationExpiresAt != nil {
|
||||
meta["validation_expires_at"] = d.ValidationExpiresAt.UTC().Format(time.RFC3339)
|
||||
}
|
||||
m.accountManager.StoreEvent(ctx, activity.SystemInitiator, d.ID, d.AccountID,
|
||||
activity.CustomDomainValidationExpired, meta)
|
||||
}
|
||||
@@ -0,0 +1,274 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
"github.com/netbirdio/netbird/management/server/mock_server"
|
||||
nbstore "github.com/netbirdio/netbird/management/server/store"
|
||||
)
|
||||
|
||||
func TestValidateDomain_ExpiredRegistration(t *testing.T) {
|
||||
env := setupDomainTest(t)
|
||||
ctx := context.Background()
|
||||
d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "expired.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
expiresAt := time.Now().Add(-time.Second)
|
||||
db := env.store.(*nbstore.SqlStore).GetDB()
|
||||
require.NoError(t, db.Model(&domain.Domain{}).Where("id = ?", d.ID).
|
||||
Update("validation_expires_at", expiresAt).Error)
|
||||
env.resolver.set("validation.expired.example.com", testCluster)
|
||||
|
||||
env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID)
|
||||
|
||||
stored := storedDomain(t, env.store, accountA, d.Domain)
|
||||
require.NotNil(t, stored)
|
||||
assert.False(t, stored.Validated, "an expired registration must not become usable before cleanup runs")
|
||||
}
|
||||
|
||||
func TestCreateDomain_ValidationDeadline(t *testing.T) {
|
||||
env := setupClockDomainTest(t)
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
createdAt := time.Now().UTC()
|
||||
d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "pending.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, d.ValidationExpiresAt)
|
||||
assert.Equal(t, createdAt.Add(48*time.Hour), *d.ValidationExpiresAt, "new registrations get 48 hours")
|
||||
|
||||
time.Sleep(time.Hour)
|
||||
env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID)
|
||||
stored := storedDomain(t, env.store, accountA, d.Domain)
|
||||
require.NotNil(t, stored)
|
||||
require.NotNil(t, stored.ValidationExpiresAt)
|
||||
assert.WithinDuration(t, *d.ValidationExpiresAt, *stored.ValidationExpiresAt, 0, "failed validation must not extend the deadline")
|
||||
})
|
||||
}
|
||||
|
||||
func TestCleanupExpiredDomains_Boundaries(t *testing.T) {
|
||||
env := setupDomainTest(t)
|
||||
events := captureDomainEvents(env)
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC().Truncate(time.Second)
|
||||
tests := []struct {
|
||||
name string
|
||||
expiresAt time.Time
|
||||
validated bool
|
||||
deleted bool
|
||||
}{
|
||||
{"expired", now.Add(-time.Second), false, true},
|
||||
{"deadline", now, false, true},
|
||||
{"pending", now.Add(time.Second), false, false},
|
||||
{"validated", now.Add(-time.Hour), true, false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
d := createExpiringDomain(t, env, tt.name+".example.com", tt.expiresAt)
|
||||
if tt.validated {
|
||||
require.NoError(t, env.store.(*nbstore.SqlStore).GetDB().Model(d).Update("validated", true).Error)
|
||||
}
|
||||
env.manager.cleanupExpiredDomains(ctx, now)
|
||||
stored := storedDomain(t, env.store, accountA, d.Domain)
|
||||
if !tt.deleted {
|
||||
assert.NotNil(t, stored, "pending and validated registrations must survive cleanup")
|
||||
return
|
||||
}
|
||||
assert.Nil(t, stored, "expired unused registrations must be removed")
|
||||
replacement, err := env.manager.CreateDomain(ctx, accountB, accountBUser, d.Domain, testCluster)
|
||||
require.NoError(t, err)
|
||||
assert.NotEqual(t, d.ID, replacement.ID, "the released name must receive a fresh registration")
|
||||
assert.False(t, replacement.Validated, "the new account must validate its own registration")
|
||||
})
|
||||
}
|
||||
got := events.get()
|
||||
require.Len(t, got, 2, "only successful expiration deletions emit events")
|
||||
for _, event := range got {
|
||||
assert.Equal(t, activity.CustomDomainValidationExpired, event.Activity, "use the requested expiration event")
|
||||
assert.Equal(t, activity.SystemInitiator, event.InitiatorID, "cleanup is attributed to the system")
|
||||
assert.Equal(t, accountA, event.AccountID, "expiration belongs to the original account")
|
||||
assert.NotEmpty(t, event.TargetID, "retain the deleted domain ID")
|
||||
assert.NotEmpty(t, event.Meta["domain"], "retain the deleted domain name")
|
||||
assert.NotEmpty(t, event.Meta["validation_expires_at"], "include the validation deadline")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupExpiredDomains_ContinuesPastProtectedBatch(t *testing.T) {
|
||||
env := setupDomainTest(t)
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC()
|
||||
for i := range validationCleanupBatch {
|
||||
d := createExpiringDomain(t, env, fmt.Sprintf("protected-%d.example.com", i), now.Add(-time.Hour))
|
||||
require.NoError(t, env.store.CreateService(ctx, &rpservice.Service{
|
||||
ID: fmt.Sprintf("service-%d", i), AccountID: accountA, Domain: "app." + d.Domain,
|
||||
}))
|
||||
}
|
||||
unprotected := createExpiringDomain(t, env, "unused.example.com", now.Add(-time.Hour))
|
||||
env.manager.cleanupExpiredDomains(ctx, now)
|
||||
assert.Nil(t, storedDomain(t, env.store, accountA, unprotected.Domain), "protected registrations must not starve later batches")
|
||||
remaining, err := env.store.ListCustomDomains(ctx, accountA)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, remaining, validationCleanupBatch, "all registrations with dependent services must survive")
|
||||
}
|
||||
|
||||
func TestCleanupExpiredDomains_ConcurrentWorkers(t *testing.T) {
|
||||
env := setupDomainTest(t)
|
||||
events := captureDomainEvents(env)
|
||||
now := time.Now().UTC()
|
||||
d := createExpiringDomain(t, env, "concurrent.example.com", now.Add(-time.Hour))
|
||||
var workers sync.WaitGroup
|
||||
for range 2 {
|
||||
workers.Go(func() { env.manager.cleanupExpiredDomains(context.Background(), now) })
|
||||
}
|
||||
workers.Wait()
|
||||
assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "one worker must remove the expired registration")
|
||||
assert.Len(t, events.get(), 1, "only the worker that deletes the row may emit the event")
|
||||
}
|
||||
|
||||
func TestRunValidationCleanup_HourlyAndRestart(t *testing.T) {
|
||||
env := setupClockDomainTest(t)
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
events := captureDomainEvents(env)
|
||||
now := time.Now().UTC()
|
||||
startup := createExpiringDomain(t, env, "startup.example.com", now.Add(-time.Hour))
|
||||
hourly := createExpiringDomain(t, env, "hourly.example.com", now.Add(time.Minute))
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
env.manager.RunValidationCleanup(ctx)
|
||||
}()
|
||||
synctest.Wait()
|
||||
assert.Nil(t, storedDomain(t, env.store, accountA, startup.Domain), "startup must collect overdue registrations")
|
||||
time.Sleep(59 * time.Minute)
|
||||
synctest.Wait()
|
||||
assert.NotNil(t, storedDomain(t, env.store, accountA, hourly.Domain), "cleanup must wait for the 60-minute interval")
|
||||
time.Sleep(time.Minute)
|
||||
synctest.Wait()
|
||||
assert.Nil(t, storedDomain(t, env.store, accountA, hourly.Domain), "the hourly scan must collect expired registrations")
|
||||
cancel()
|
||||
<-done
|
||||
|
||||
offline := createExpiringDomain(t, env, "offline.example.com", time.Now().UTC().Add(time.Minute))
|
||||
time.Sleep(2 * time.Hour)
|
||||
assert.NotNil(t, storedDomain(t, env.store, accountA, offline.Domain), "a stopped worker must not continue deleting")
|
||||
ctx, cancel = context.WithCancel(context.Background())
|
||||
done = make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
env.manager.RunValidationCleanup(ctx)
|
||||
}()
|
||||
synctest.Wait()
|
||||
assert.Nil(t, storedDomain(t, env.store, accountA, offline.Domain), "restart must use the persisted deadline")
|
||||
cancel()
|
||||
<-done
|
||||
assert.Len(t, events.get(), 3, "each deletion should emit an expiration event")
|
||||
})
|
||||
}
|
||||
|
||||
type blockingDomainResolver struct {
|
||||
started chan struct{}
|
||||
release chan struct{}
|
||||
}
|
||||
|
||||
func (r blockingDomainResolver) LookupCNAME(context.Context, string) (string, error) {
|
||||
close(r.started)
|
||||
<-r.release
|
||||
return testCluster + ".", nil
|
||||
}
|
||||
|
||||
func TestValidateDomain_DeadlinePassesDuringLookup(t *testing.T) {
|
||||
for _, cleanup := range []bool{false, true} {
|
||||
t.Run(fmt.Sprintf("cleanup=%t", cleanup), func(t *testing.T) {
|
||||
env := setupClockDomainTest(t)
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
events := captureDomainEvents(env)
|
||||
ctx := context.Background()
|
||||
d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "late.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
resolver := blockingDomainResolver{started: make(chan struct{}), release: make(chan struct{})}
|
||||
env.manager.validator.Resolver = resolver
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID)
|
||||
}()
|
||||
<-resolver.started
|
||||
time.Sleep(48 * time.Hour)
|
||||
if cleanup {
|
||||
env.manager.cleanupExpiredDomains(ctx, time.Now().UTC())
|
||||
_, err = env.store.CreateCustomDomain(ctx, accountB, d.Domain, testCluster, false)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
close(resolver.release)
|
||||
<-done
|
||||
owner := accountA
|
||||
if cleanup {
|
||||
assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "late validation must not restore the old claim")
|
||||
owner = accountB
|
||||
}
|
||||
stored := storedDomain(t, env.store, owner, d.Domain)
|
||||
require.NotNil(t, stored)
|
||||
assert.False(t, stored.Validated, "late validation must not validate either claim")
|
||||
for _, event := range events.get() {
|
||||
assert.NotEqual(t, activity.DomainValidated, event.Activity, "a rejected write must not emit a validation event")
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func setupClockDomainTest(t *testing.T) *domainTestEnv {
|
||||
t.Helper()
|
||||
// Network driver watchers cannot share cancellation channels across synctest bubbles.
|
||||
// Store boundary and concurrency tests still exercise the selected database engine.
|
||||
t.Setenv("NETBIRD_STORE_ENGINE", "sqlite")
|
||||
return setupDomainTest(t)
|
||||
}
|
||||
|
||||
func createExpiringDomain(t *testing.T, env *domainTestEnv, name string, expiresAt time.Time) *domain.Domain {
|
||||
t.Helper()
|
||||
d, err := env.store.CreateCustomDomain(context.Background(), accountA, name, testCluster, false)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, env.store.(*nbstore.SqlStore).GetDB().Model(d).Update("validation_expires_at", expiresAt).Error)
|
||||
d.ValidationExpiresAt = &expiresAt
|
||||
return d
|
||||
}
|
||||
|
||||
type domainEvents struct {
|
||||
mu sync.Mutex
|
||||
events []*activity.Event
|
||||
}
|
||||
|
||||
func captureDomainEvents(env *domainTestEnv) *domainEvents {
|
||||
events := &domainEvents{}
|
||||
env.manager.accountManager = &mock_server.MockAccountManager{
|
||||
StoreEventFunc: func(_ context.Context, initiator, target, account string, code activity.ActivityDescriber, meta map[string]any) {
|
||||
if code == activity.DomainAdded {
|
||||
return
|
||||
}
|
||||
events.mu.Lock()
|
||||
defer events.mu.Unlock()
|
||||
events.events = append(events.events, &activity.Event{
|
||||
InitiatorID: initiator, TargetID: target, AccountID: account,
|
||||
Activity: code.(activity.Activity), Meta: meta,
|
||||
})
|
||||
},
|
||||
}
|
||||
return events
|
||||
}
|
||||
|
||||
func (e *domainEvents) get() []*activity.Event {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
return append([]*activity.Event(nil), e.events...)
|
||||
}
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
@@ -18,6 +19,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
nbstore "github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
nbdomain "github.com/netbirdio/netbird/shared/management/domain"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
@@ -32,6 +34,8 @@ type store interface {
|
||||
CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error)
|
||||
UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error)
|
||||
DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error
|
||||
GetExpiredCustomDomains(ctx context.Context, now time.Time, afterID domain.ID, limit int) ([]*domain.Domain, error)
|
||||
DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error)
|
||||
}
|
||||
|
||||
type proxyManager interface {
|
||||
@@ -106,12 +110,13 @@ func (m Manager) GetDomains(ctx context.Context, accountID, userID string) ([]*d
|
||||
// Add custom domains.
|
||||
for _, d := range domains {
|
||||
cd := &domain.Domain{
|
||||
ID: d.ID,
|
||||
Domain: d.Domain,
|
||||
AccountID: accountID,
|
||||
TargetCluster: d.TargetCluster,
|
||||
Type: domain.TypeCustom,
|
||||
Validated: d.Validated,
|
||||
ID: d.ID,
|
||||
Domain: d.Domain,
|
||||
AccountID: accountID,
|
||||
TargetCluster: d.TargetCluster,
|
||||
Type: domain.TypeCustom,
|
||||
Validated: d.Validated,
|
||||
ValidationExpiresAt: d.ValidationExpiresAt,
|
||||
}
|
||||
if d.TargetCluster != "" {
|
||||
cd.SupportsCustomPorts = m.proxyManager.ClusterSupportsCustomPorts(ctx, d.TargetCluster)
|
||||
@@ -126,6 +131,7 @@ func (m Manager) GetDomains(ctx context.Context, accountID, userID string) ([]*d
|
||||
return ret, nil
|
||||
}
|
||||
|
||||
// CreateDomain registers a normalized custom domain and attempts DNS validation.
|
||||
func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName, targetCluster string) (*domain.Domain, error) {
|
||||
ok, ctx, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Services, operations.Create)
|
||||
if err != nil {
|
||||
@@ -135,6 +141,15 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName
|
||||
return nil, status.NewPermissionDeniedError()
|
||||
}
|
||||
|
||||
parsed, err := nbdomain.FromString(strings.TrimSuffix(domainName, "."))
|
||||
if err != nil {
|
||||
return nil, status.Errorf(status.InvalidArgument, "invalid domain: %v", err)
|
||||
}
|
||||
domainName = parsed.PunycodeString()
|
||||
if !nbdomain.IsValidDomainNoWildcard(domainName) {
|
||||
return nil, status.Errorf(status.InvalidArgument, "invalid domain format")
|
||||
}
|
||||
|
||||
// Verify the target cluster is in the available clusters for this account
|
||||
allowList, err := m.getClusterAllowList(ctx, accountID)
|
||||
if err != nil {
|
||||
@@ -243,6 +258,14 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID
|
||||
}).WithError(err).Error("get custom domain from store")
|
||||
return
|
||||
}
|
||||
if d.Validated {
|
||||
return
|
||||
}
|
||||
if d.ValidationExpiresAt == nil || !time.Now().Before(*d.ValidationExpiresAt) {
|
||||
log.WithFields(log.Fields{"accountID": accountID, "domainID": domainID}).
|
||||
Debug("custom domain validation window has expired")
|
||||
return
|
||||
}
|
||||
|
||||
// Validate only against the domain's target cluster
|
||||
targetCluster := d.TargetCluster
|
||||
@@ -263,20 +286,21 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID
|
||||
}).Info("validating domain against target cluster")
|
||||
|
||||
if m.validator.IsValid(context.Background(), d.Domain, []string{targetCluster}) {
|
||||
log.WithFields(log.Fields{
|
||||
"accountID": accountID,
|
||||
"domainID": domainID,
|
||||
"domain": d.Domain,
|
||||
}).Info("domain validated successfully")
|
||||
d.Validated = true
|
||||
if _, err := m.store.UpdateCustomDomain(context.Background(), accountID, d); err != nil {
|
||||
log.WithFields(log.Fields{
|
||||
entry := log.WithFields(log.Fields{
|
||||
"accountID": accountID,
|
||||
"domainID": domainID,
|
||||
"domain": d.Domain,
|
||||
}).WithError(err).Error("update custom domain in store")
|
||||
}).WithError(err)
|
||||
if sErr, ok := status.FromError(err); ok && sErr.Type() == status.PreconditionFailed {
|
||||
entry.Debug("custom domain registration is no longer pending validation")
|
||||
return
|
||||
}
|
||||
entry.Error("update custom domain in store")
|
||||
return
|
||||
}
|
||||
log.WithFields(log.Fields{"accountID": accountID, "domainID": domainID}).
|
||||
Info("custom domain validated successfully")
|
||||
|
||||
m.accountManager.StoreEvent(context.Background(), userID, domainID, accountID, activity.DomainValidated, d.EventMeta())
|
||||
} else {
|
||||
|
||||
@@ -99,7 +99,7 @@ func setupDomainTest(t *testing.T) *domainTestEnv {
|
||||
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", nil, nil)
|
||||
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", "", nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resolver := &stubResolver{cnames: make(map[string]string)}
|
||||
@@ -296,11 +296,8 @@ func TestValidateDomain_PermissionDeniedDoesNotValidate(t *testing.T) {
|
||||
assert.Error(t, err, "the domain must still be unservable")
|
||||
}
|
||||
|
||||
// Validation runs asynchronously, so it can finish after the domain was
|
||||
// deleted and then write a stale row back. gorm's Save falls back to an insert
|
||||
// when an update affects no rows, which would resurrect the domain as
|
||||
// validated; UpdateCustomDomain avoids that by selecting explicit columns.
|
||||
// This pins that behaviour, since dropping the Select would reintroduce it.
|
||||
// A validation finishing after deletion must reject the stale write, without
|
||||
// restoring the registration or reporting successful validation.
|
||||
func TestUpdateCustomDomain_DoesNotResurrectDeletedDomain(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
@@ -315,11 +312,9 @@ func TestUpdateCustomDomain_DoesNotResurrectDeletedDomain(t *testing.T) {
|
||||
require.Nil(t, storedDomain(t, env.store, accountA, "racy.example.com"), "the domain should be gone")
|
||||
|
||||
// What an in-flight validation would write once its CNAME check succeeded.
|
||||
// The write has to succeed for the assertion below to mean anything: a
|
||||
// rejected write would leave the domain absent for the wrong reason.
|
||||
stale.Validated = true
|
||||
_, err = env.store.UpdateCustomDomain(ctx, accountA, stale)
|
||||
require.NoError(t, err, "the update itself must succeed, so absence is not just a failed write")
|
||||
require.Error(t, err, "a deleted registration must reject a late validation")
|
||||
|
||||
assert.Nil(t, storedDomain(t, env.store, accountA, "racy.example.com"),
|
||||
"a late validation write must not recreate a deleted domain")
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -208,6 +209,14 @@ func (s *stubStore) DeleteCustomDomain(context.Context, string, string) error {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) GetExpiredCustomDomains(context.Context, time.Time, domain.ID, int) ([]*domain.Domain, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) DeleteExpiredCustomDomain(context.Context, *domain.Domain, time.Time) (bool, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
// TestGetClusterAllowList_DedicatedGatewayAddressExcluded pins invariant (B)'s
|
||||
// chokepoint: a self-addressed settings pin reserves the account's gateway
|
||||
// address, so it is dropped from the allow list — which, because the
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
func TestCreateDomain_NormalizesName(t *testing.T) {
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
input string
|
||||
canonical string
|
||||
}{
|
||||
{"mixed case", "Apps.Example.COM", "apps.example.com"},
|
||||
{"unicode", "münchen.example.com", "xn--mnchen-3ya.example.com"},
|
||||
{"trailing dot", "apps.example.com.", "apps.example.com"},
|
||||
{"underscore", "My_App.example.com", "my_app.example.com"},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
env.resolver.set("validation."+tt.canonical, testCluster)
|
||||
|
||||
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, tt.input, testCluster)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tt.canonical, created.Domain, "the response must use the normalized name")
|
||||
assert.True(t, created.Validated, "the CNAME lookup must use the normalized name")
|
||||
stored, err := env.store.GetCustomDomain(ctx, accountA, created.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tt.canonical, stored.Domain, "the database must retain the normalized name")
|
||||
|
||||
_, err = env.manager.CreateDomain(ctx, accountB, accountBUser, tt.canonical, testCluster)
|
||||
require.Error(t, err)
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "an equivalent name must return a typed conflict")
|
||||
assert.Equal(t, status.AlreadyExists, sErr.Type(), "normalization must precede the availability check")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateDomain_NormalizedNameCanValidateLater(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "Apps.Example.COM.", testCluster)
|
||||
require.NoError(t, err)
|
||||
require.False(t, created.Validated, "a missing CNAME must leave the normalized registration pending")
|
||||
|
||||
env.resolver.set("validation.apps.example.com", testCluster)
|
||||
env.manager.ValidateDomain(ctx, accountA, accountAUser, created.ID)
|
||||
stored, err := env.store.GetCustomDomain(ctx, accountA, created.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "apps.example.com", stored.Domain, "retrying validation must retain the normalized name")
|
||||
assert.True(t, stored.Validated, "later validation must look up the normalized name")
|
||||
}
|
||||
|
||||
func TestCreateDomain_RejectsInvalidName(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
for _, name := range []string{
|
||||
"", ".", "app..example.com", "app.example.com..", "-app.example.com",
|
||||
"app%.example.com", "app!.example.com", "*.example.com", "app example.com",
|
||||
"https://example.com", strings.Repeat("a", 64) + ".example.com",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
// A matching DNS response must not make a malformed name acceptable.
|
||||
env.resolver.set("validation."+name, testCluster)
|
||||
_, err := env.manager.CreateDomain(ctx, accountA, accountAUser, name, testCluster)
|
||||
require.Error(t, err)
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "invalid names must return a typed client error")
|
||||
assert.Equal(t, status.InvalidArgument, sErr.Type(), "malformed names must be rejected before storage")
|
||||
})
|
||||
}
|
||||
stored, err := env.store.ListCustomDomains(ctx, accountA)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, stored, "invalid registration attempts must not reserve any names")
|
||||
}
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
|
||||
// Manager defines the interface for proxy operations
|
||||
type Manager interface {
|
||||
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error)
|
||||
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error)
|
||||
Disconnect(ctx context.Context, proxyID, sessionID string) error
|
||||
Heartbeat(ctx context.Context, p *Proxy) error
|
||||
GetActiveClusterAddresses(ctx context.Context) ([]string, error)
|
||||
|
||||
@@ -50,7 +50,7 @@ func NewManager(store store, meter metric.Meter) (*Manager, error) {
|
||||
|
||||
// Connect registers a new proxy connection in the database.
|
||||
// capabilities may be nil for old proxies that do not report them.
|
||||
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
|
||||
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
|
||||
now := time.Now()
|
||||
var caps proxy.Capabilities
|
||||
if capabilities != nil {
|
||||
@@ -61,6 +61,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
|
||||
SessionID: sessionID,
|
||||
ClusterAddress: clusterAddress,
|
||||
IPAddress: ipAddress,
|
||||
Version: truncateVersion(version),
|
||||
AccountID: accountID,
|
||||
LastSeen: now,
|
||||
ConnectedAt: &now,
|
||||
@@ -78,6 +79,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
|
||||
"sessionID": sessionID,
|
||||
"clusterAddress": clusterAddress,
|
||||
"ipAddress": ipAddress,
|
||||
"version": p.Version,
|
||||
}).Info("proxy connected")
|
||||
|
||||
return p, nil
|
||||
@@ -184,3 +186,13 @@ func (m *Manager) DeleteAccountCluster(ctx context.Context, clusterAddress, acco
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// truncateVersion cuts a proxy-reported version to the column width so an
|
||||
// oversized value cannot fail the save and block the connect.
|
||||
func truncateVersion(version string) string {
|
||||
runes := []rune(version)
|
||||
if len(runes) <= proxy.MaxVersionLength {
|
||||
return version
|
||||
}
|
||||
return string(runes[:proxy.MaxVersionLength])
|
||||
}
|
||||
|
||||
@@ -4,8 +4,10 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -124,7 +126,7 @@ func TestConnect_WithAccountID(t *testing.T) {
|
||||
}
|
||||
|
||||
mgr := newTestManager(s)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", &accountID, nil)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "0.60.0", &accountID, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotNil(t, savedProxy)
|
||||
@@ -132,6 +134,7 @@ func TestConnect_WithAccountID(t *testing.T) {
|
||||
assert.Equal(t, "session-1", savedProxy.SessionID)
|
||||
assert.Equal(t, "cluster.example.com", savedProxy.ClusterAddress)
|
||||
assert.Equal(t, "10.0.0.1", savedProxy.IPAddress)
|
||||
assert.Equal(t, "0.60.0", savedProxy.Version, "reported proxy version should be stored")
|
||||
assert.Equal(t, &accountID, savedProxy.AccountID)
|
||||
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
|
||||
assert.NotNil(t, savedProxy.ConnectedAt)
|
||||
@@ -147,7 +150,7 @@ func TestConnect_WithoutAccountID(t *testing.T) {
|
||||
}
|
||||
|
||||
mgr := newTestManager(s)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", nil, nil)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", "", nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotNil(t, savedProxy)
|
||||
@@ -155,6 +158,29 @@ func TestConnect_WithoutAccountID(t *testing.T) {
|
||||
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
|
||||
}
|
||||
|
||||
func TestConnect_TruncatesOversizedVersion(t *testing.T) {
|
||||
var savedProxy *proxy.Proxy
|
||||
s := &mockStore{
|
||||
saveProxyFunc: func(_ context.Context, p *proxy.Proxy) error {
|
||||
savedProxy = p
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
// Multi-byte runes make sure the cut counts characters, as varchar does,
|
||||
// and never splits a rune into invalid UTF-8.
|
||||
version := strings.Repeat("ü", proxy.MaxVersionLength+10)
|
||||
|
||||
mgr := newTestManager(s)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", version, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NotNil(t, savedProxy)
|
||||
assert.Equal(t, proxy.MaxVersionLength, utf8.RuneCountInString(savedProxy.Version), "stored version should be cut to the column width")
|
||||
assert.True(t, utf8.ValidString(savedProxy.Version), "stored version should remain valid UTF-8")
|
||||
assert.True(t, strings.HasPrefix(version, savedProxy.Version), "stored version should be a prefix of the reported one")
|
||||
}
|
||||
|
||||
func TestConnect_StoreError(t *testing.T) {
|
||||
s := &mockStore{
|
||||
saveProxyFunc: func(_ context.Context, _ *proxy.Proxy) error {
|
||||
@@ -163,7 +189,7 @@ func TestConnect_StoreError(t *testing.T) {
|
||||
}
|
||||
|
||||
mgr := newTestManager(s)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", nil, nil)
|
||||
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "", nil, nil)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
|
||||
@@ -113,18 +113,18 @@ func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any)
|
||||
}
|
||||
|
||||
// Connect mocks base method.
|
||||
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
|
||||
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
|
||||
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
|
||||
ret0, _ := ret[0].(*Proxy)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// Connect indicates an expected call of Connect.
|
||||
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
|
||||
}
|
||||
|
||||
// CountAccountProxies mocks base method.
|
||||
|
||||
@@ -9,6 +9,9 @@ const (
|
||||
StatusDisconnected = "disconnected"
|
||||
)
|
||||
|
||||
// MaxVersionLength is the width of the Version column, in characters.
|
||||
const MaxVersionLength = 255
|
||||
|
||||
// Capabilities describes what a proxy can handle, as reported via gRPC.
|
||||
// Nil fields mean the proxy never reported this capability.
|
||||
type Capabilities struct {
|
||||
@@ -31,6 +34,7 @@ type Proxy struct {
|
||||
SessionID string `gorm:"type:varchar(36)"`
|
||||
ClusterAddress string `gorm:"type:varchar(255);not null;index:idx_proxy_cluster_status"`
|
||||
IPAddress string `gorm:"type:varchar(45)"`
|
||||
Version string `gorm:"type:varchar(255)"`
|
||||
AccountID *string `gorm:"type:varchar(255);index:idx_proxy_account_id"`
|
||||
LastSeen time.Time `gorm:"not null;index:idx_proxy_last_seen"`
|
||||
ConnectedAt *time.Time
|
||||
|
||||
@@ -30,7 +30,7 @@ func withRealDomainManager(t *testing.T, mgr *Manager, testStore store.Store) {
|
||||
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", nil, nil)
|
||||
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", "", nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
accountMgr := &mock_server.MockAccountManager{
|
||||
|
||||
@@ -31,6 +31,7 @@ import (
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
|
||||
networkmapdbfactory "github.com/netbirdio/netbird/management/internals/network_map_db/factory"
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
|
||||
@@ -111,7 +112,8 @@ func (s *BaseServer) NetworkMapStore() *networkmapdb.NetworkMapDBStoreImpl {
|
||||
s.Config.StoreConfig.Engine,
|
||||
s.Config.Datadir,
|
||||
s.IntegratedValidator(),
|
||||
s.SettingsManager())
|
||||
s.SettingsManager(),
|
||||
)
|
||||
// networkmap db store supports postgres and sqlite backends only
|
||||
// for other backends a fallback is used, so NotSupportedStoreEngineError
|
||||
// is not a fatal error
|
||||
@@ -180,24 +182,7 @@ func (s *BaseServer) RateLimiter() *middleware.APIRateLimiter {
|
||||
|
||||
func (s *BaseServer) GRPCServer() *grpc.Server {
|
||||
return Create(s, func() *grpc.Server {
|
||||
trustedPeers := s.Config.ReverseProxy.TrustedPeers
|
||||
defaultTrustedPeers := []netip.Prefix{netip.MustParsePrefix("0.0.0.0/0"), netip.MustParsePrefix("::/0")}
|
||||
if len(trustedPeers) == 0 || slices.Equal[[]netip.Prefix](trustedPeers, defaultTrustedPeers) {
|
||||
log.WithContext(context.Background()).Warn("TrustedPeers are configured to default value '0.0.0.0/0', '::/0'. This allows connection IP spoofing.")
|
||||
trustedPeers = defaultTrustedPeers
|
||||
}
|
||||
trustedHTTPProxies := s.Config.ReverseProxy.TrustedHTTPProxies
|
||||
trustedProxiesCount := s.Config.ReverseProxy.TrustedHTTPProxiesCount
|
||||
if len(trustedHTTPProxies) > 0 && trustedProxiesCount > 0 {
|
||||
log.WithContext(context.Background()).Warn("TrustedHTTPProxies and TrustedHTTPProxiesCount both are configured. " +
|
||||
"This is not recommended way to extract X-Forwarded-For. Consider using one of these options.")
|
||||
}
|
||||
realipOpts := []realip.Option{
|
||||
realip.WithTrustedPeers(trustedPeers),
|
||||
realip.WithTrustedProxies(trustedHTTPProxies),
|
||||
realip.WithTrustedProxiesCount(trustedProxiesCount),
|
||||
realip.WithHeaders([]string{realip.XForwardedFor, realip.XRealIp}),
|
||||
}
|
||||
realipOpts := realIPOptions(s.Config.ReverseProxy)
|
||||
proxyUnary, proxyStream, proxyAuthClose := nbgrpc.NewProxyAuthInterceptors(s.Store())
|
||||
s.proxyAuthClose = proxyAuthClose
|
||||
gRPCOpts := []grpc.ServerOption{
|
||||
@@ -333,7 +318,7 @@ func (s *BaseServer) AccessLogsManager() accesslogs.Manager {
|
||||
})
|
||||
}
|
||||
|
||||
func loadTLSConfig(certFile string, certKey string) (*tls.Config, error) {
|
||||
func loadTLSConfig(certFile, certKey string) (*tls.Config, error) {
|
||||
// Load server's certificate and private key
|
||||
serverCert, err := tls.LoadX509KeyPair(certFile, certKey)
|
||||
if err != nil {
|
||||
@@ -380,3 +365,37 @@ func streamInterceptor(
|
||||
wrapped.WrappedContext = context.WithValue(ctx, nbContext.RequestIDKey, reqID)
|
||||
return handler(srv, wrapped)
|
||||
}
|
||||
|
||||
// realIPOptions builds the real-IP middleware options.
|
||||
//
|
||||
// Empty TrustedPeers trusts all IPv4 and IPv6 sources. Configure TrustedPeers
|
||||
// with the reverse proxy address or network.
|
||||
//
|
||||
// X-Forwarded-For takes precedence over X-Real-IP.
|
||||
func realIPOptions(cfg nbconfig.ReverseProxy) []realip.Option {
|
||||
trustedPeers := cfg.TrustedPeers
|
||||
if len(trustedPeers) == 0 {
|
||||
trustedPeers = []netip.Prefix{
|
||||
netip.MustParsePrefix("0.0.0.0/0"),
|
||||
netip.MustParsePrefix("::/0"),
|
||||
}
|
||||
}
|
||||
if idx := slices.IndexFunc(trustedPeers, func(p netip.Prefix) bool { return p.Bits() == 0 }); idx >= 0 {
|
||||
log.WithContext(context.Background()).Warnf("TrustedPeers contains the default route %s, which trusts "+
|
||||
"X-Forwarded-For from every client and allows connection IP spoofing. Set TrustedPeers to the address "+
|
||||
"of your reverse proxy.", trustedPeers[idx])
|
||||
}
|
||||
if cfg.TrustedHTTPProxiesCount > 0 {
|
||||
log.WithContext(context.Background()).Warn(
|
||||
"TrustedHTTPProxiesCount skips X-Forwarded-For entries by position before TrustedHTTPProxies filters by address. " +
|
||||
"An incorrect count may skip the real client IP and produce an incorrect source address.",
|
||||
)
|
||||
}
|
||||
|
||||
return []realip.Option{
|
||||
realip.WithTrustedPeers(trustedPeers),
|
||||
realip.WithTrustedProxies(cfg.TrustedHTTPProxies),
|
||||
realip.WithTrustedProxiesCount(cfg.TrustedHTTPProxiesCount),
|
||||
realip.WithHeaders([]string{realip.XForwardedFor, realip.XRealIp}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/realip"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
)
|
||||
|
||||
const (
|
||||
realIPProbeMethod = "/netbird.test.RealIPProbe/Probe"
|
||||
realIPProbeStreamMethod = "/netbird.test.RealIPProbe/ProbeStream"
|
||||
)
|
||||
|
||||
// realIPProbe records the real IP the middleware derived for each call.
|
||||
type realIPProbe struct {
|
||||
got chan string
|
||||
}
|
||||
|
||||
func (p *realIPProbe) record(ctx context.Context) {
|
||||
addr, _ := realip.FromContext(ctx)
|
||||
p.got <- addr.String()
|
||||
}
|
||||
|
||||
func (p *realIPProbe) wait(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
select {
|
||||
case got := <-p.got:
|
||||
return got
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out waiting for probe")
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func startProbeServer(t *testing.T, cfg nbconfig.ReverseProxy) (*grpc.ClientConn, *realIPProbe) {
|
||||
t.Helper()
|
||||
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
probe := &realIPProbe{got: make(chan string, 1)}
|
||||
opts := realIPOptions(cfg)
|
||||
srv := grpc.NewServer(
|
||||
grpc.ChainUnaryInterceptor(realip.UnaryServerInterceptorOpts(opts...)),
|
||||
grpc.ChainStreamInterceptor(realip.StreamServerInterceptorOpts(opts...)),
|
||||
)
|
||||
srv.RegisterService(&grpc.ServiceDesc{
|
||||
ServiceName: "netbird.test.RealIPProbe",
|
||||
HandlerType: (*any)(nil),
|
||||
Methods: []grpc.MethodDesc{{
|
||||
MethodName: "Probe",
|
||||
Handler: func(_ any, ctx context.Context, dec func(any) error, interceptor grpc.UnaryServerInterceptor) (any, error) {
|
||||
req := new(emptypb.Empty)
|
||||
if err := dec(req); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
handler := func(ctx context.Context, _ any) (any, error) {
|
||||
probe.record(ctx)
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
if interceptor == nil {
|
||||
return handler(ctx, req)
|
||||
}
|
||||
return interceptor(ctx, req, &grpc.UnaryServerInfo{FullMethod: realIPProbeMethod}, handler)
|
||||
},
|
||||
}},
|
||||
Streams: []grpc.StreamDesc{{
|
||||
StreamName: "ProbeStream",
|
||||
ServerStreams: true,
|
||||
Handler: func(_ any, stream grpc.ServerStream) error {
|
||||
probe.record(stream.Context())
|
||||
return nil
|
||||
},
|
||||
}},
|
||||
}, probe)
|
||||
|
||||
go func() { _ = srv.Serve(listener) }()
|
||||
t.Cleanup(srv.Stop)
|
||||
|
||||
conn, err := grpc.NewClient(listener.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
|
||||
return conn, probe
|
||||
}
|
||||
|
||||
func callUnary(t *testing.T, conn *grpc.ClientConn, probe *realIPProbe, kv ...string) string {
|
||||
t.Helper()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
ctx = metadata.AppendToOutgoingContext(ctx, kv...)
|
||||
require.NoError(t, conn.Invoke(ctx, realIPProbeMethod, &emptypb.Empty{}, &emptypb.Empty{}))
|
||||
|
||||
return probe.wait(t)
|
||||
}
|
||||
|
||||
func callStream(t *testing.T, conn *grpc.ClientConn, probe *realIPProbe, kv ...string) string {
|
||||
t.Helper()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
ctx = metadata.AppendToOutgoingContext(ctx, kv...)
|
||||
desc := &grpc.StreamDesc{StreamName: "ProbeStream", ServerStreams: true}
|
||||
stream, err := conn.NewStream(ctx, desc, realIPProbeStreamMethod)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, stream.CloseSend())
|
||||
require.ErrorIs(t, stream.RecvMsg(&emptypb.Empty{}), io.EOF)
|
||||
|
||||
return probe.wait(t)
|
||||
}
|
||||
|
||||
func assertRealIP(t *testing.T, cfg nbconfig.ReverseProxy, want string, kv ...string) {
|
||||
t.Helper()
|
||||
|
||||
conn, probe := startProbeServer(t, cfg)
|
||||
t.Run("unary", func(t *testing.T) {
|
||||
assert.Equal(t, want, callUnary(t, conn, probe, kv...))
|
||||
})
|
||||
t.Run("stream", func(t *testing.T) {
|
||||
assert.Equal(t, want, callStream(t, conn, probe, kv...))
|
||||
})
|
||||
}
|
||||
|
||||
func TestRealIPDefaultTrustsForwardedHeaders(t *testing.T) {
|
||||
assertRealIP(t, nbconfig.ReverseProxy{}, "203.0.113.44",
|
||||
realip.XForwardedFor, "203.0.113.44",
|
||||
realip.XRealIp, "203.0.113.44",
|
||||
)
|
||||
}
|
||||
|
||||
func TestRealIPUntrustedPeerIgnoresForwardedHeaders(t *testing.T) {
|
||||
cfg := nbconfig.ReverseProxy{TrustedPeers: []netip.Prefix{netip.MustParsePrefix("10.9.8.7/32")}}
|
||||
|
||||
assertRealIP(t, cfg, "127.0.0.1",
|
||||
realip.XForwardedFor, "203.0.113.44",
|
||||
realip.XRealIp, "203.0.113.44",
|
||||
)
|
||||
}
|
||||
|
||||
func TestRealIPTrustedPeerHonoursForwardedHeaders(t *testing.T) {
|
||||
cfg := nbconfig.ReverseProxy{TrustedPeers: []netip.Prefix{netip.MustParsePrefix("127.0.0.1/32")}}
|
||||
|
||||
assertRealIP(t, cfg, "203.0.113.44",
|
||||
realip.XForwardedFor, "203.0.113.44",
|
||||
realip.XRealIp, "203.0.113.44",
|
||||
)
|
||||
}
|
||||
|
||||
func TestRealIPReadsXRealIPWhenProxyCountSkipsForwardedFor(t *testing.T) {
|
||||
cfg := nbconfig.ReverseProxy{
|
||||
TrustedPeers: []netip.Prefix{netip.MustParsePrefix("127.0.0.1/32")},
|
||||
TrustedHTTPProxiesCount: 1,
|
||||
}
|
||||
|
||||
t.Run("no X-Forwarded-For", func(t *testing.T) {
|
||||
assertRealIP(t, cfg, "203.0.113.44", realip.XRealIp, "203.0.113.44")
|
||||
})
|
||||
t.Run("single-entry X-Forwarded-For", func(t *testing.T) {
|
||||
assertRealIP(t, cfg, "198.51.100.7",
|
||||
realip.XForwardedFor, "203.0.113.44",
|
||||
realip.XRealIp, "198.51.100.7",
|
||||
)
|
||||
})
|
||||
}
|
||||
@@ -23,6 +23,8 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/idp"
|
||||
"github.com/netbirdio/netbird/management/server/metrics"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/lifecycle"
|
||||
"github.com/netbirdio/netbird/shared/profiling"
|
||||
"github.com/netbirdio/netbird/util/wsproxy"
|
||||
wsproxyserver "github.com/netbirdio/netbird/util/wsproxy/server"
|
||||
"github.com/netbirdio/netbird/version"
|
||||
@@ -36,6 +38,8 @@ const (
|
||||
DefaultSelfHostedDomain = "netbird.selfhosted"
|
||||
|
||||
ContainerKeyBaseServer = "baseServer"
|
||||
|
||||
applicationName = "management"
|
||||
)
|
||||
|
||||
type Server interface {
|
||||
@@ -66,7 +70,8 @@ type BaseServer struct {
|
||||
disableLegacyManagementPort bool
|
||||
autoResolveDomains bool
|
||||
|
||||
proxyAuthClose func()
|
||||
proxyAuthClose func()
|
||||
domainCleanupStop func()
|
||||
|
||||
// grpcExtensions holds additional gRPC services, interceptors, and shutdown
|
||||
// hooks registered by external modules via RegisterGRPCExtension. Populated
|
||||
@@ -74,12 +79,15 @@ type BaseServer struct {
|
||||
grpcExtensions []GRPCExtension
|
||||
|
||||
listener net.Listener
|
||||
tlsConfig *tls.Config
|
||||
certManager *autocert.Manager
|
||||
update *version.Update
|
||||
|
||||
errCh chan error
|
||||
wg sync.WaitGroup
|
||||
cancel context.CancelFunc
|
||||
|
||||
lifecycle.StopHandlers
|
||||
}
|
||||
|
||||
// Config holds the configuration parameters for creating a new server
|
||||
@@ -94,6 +102,7 @@ type Config struct {
|
||||
DisableGeoliteUpdate bool
|
||||
UserDeleteFromIDPEnabled bool
|
||||
AutoResolveDomains bool
|
||||
TLSConfig *tls.Config
|
||||
}
|
||||
|
||||
// NewServer initializes and configures a new Server instance
|
||||
@@ -110,9 +119,13 @@ func NewServer(cfg *Config) *BaseServer {
|
||||
disableLegacyManagementPort: cfg.DisableLegacyManagementPort,
|
||||
mgmtMetricsPort: cfg.MgmtMetricsPort,
|
||||
autoResolveDomains: cfg.AutoResolveDomains,
|
||||
tlsConfig: cfg.TLSConfig,
|
||||
}
|
||||
s.container[ContainerKeyBaseServer] = s
|
||||
|
||||
stopProfiling := profiling.Start(applicationName)
|
||||
s.OnStop(stopProfiling)
|
||||
|
||||
return s
|
||||
}
|
||||
|
||||
@@ -122,6 +135,14 @@ func (s *BaseServer) AfterInit(fn func(s *BaseServer)) {
|
||||
|
||||
// Start begins listening for HTTP requests on the configured address
|
||||
func (s *BaseServer) Start(ctx context.Context) error {
|
||||
if err := s.start(ctx); err != nil {
|
||||
s.RunStopHandlers()
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *BaseServer) start(ctx context.Context) error {
|
||||
srvCtx, cancel := context.WithCancel(ctx)
|
||||
s.cancel = cancel
|
||||
s.errCh = make(chan error, 4)
|
||||
@@ -139,21 +160,9 @@ func (s *BaseServer) Start(ctx context.Context) error {
|
||||
}
|
||||
s.EphemeralManager().LoadInitialPeers(srvCtx)
|
||||
|
||||
var tlsConfig *tls.Config
|
||||
tlsEnabled := false
|
||||
if s.Config.HttpConfig.LetsEncryptDomain != "" {
|
||||
s.certManager, err = encryption.CreateCertManager(s.Config.Datadir, s.Config.HttpConfig.LetsEncryptDomain)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed creating LetsEncrypt cert manager: %v", err)
|
||||
}
|
||||
tlsEnabled = true
|
||||
} else if s.Config.HttpConfig.CertFile != "" && s.Config.HttpConfig.CertKey != "" {
|
||||
tlsConfig, err = loadTLSConfig(s.Config.HttpConfig.CertFile, s.Config.HttpConfig.CertKey)
|
||||
if err != nil {
|
||||
log.WithContext(srvCtx).Errorf("cannot load TLS credentials: %v", err)
|
||||
return err
|
||||
}
|
||||
tlsEnabled = true
|
||||
tlsEnabled, err := s.setupTLS(srvCtx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
installationID, err := getInstallationID(srvCtx, s.Store())
|
||||
@@ -215,8 +224,8 @@ func (s *BaseServer) Start(ctx context.Context) error {
|
||||
log.WithContext(ctx).Infof("running HTTP server (LetsEncrypt challenge handler): %s", cml.Addr().String())
|
||||
s.serveHTTP(ctx, cml, s.certManager.HTTPHandler(nil))
|
||||
}
|
||||
case tlsConfig != nil:
|
||||
s.listener, err = tls.Listen("tcp", fmt.Sprintf(":%d", s.mgmtPort), tlsConfig)
|
||||
case s.tlsConfig != nil:
|
||||
s.listener, err = tls.Listen("tcp", fmt.Sprintf(":%d", s.mgmtPort), s.tlsConfig)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed creating TLS listener on port %d: %v", s.mgmtPort, err)
|
||||
}
|
||||
@@ -236,14 +245,60 @@ func (s *BaseServer) Start(ctx context.Context) error {
|
||||
s.update.SetOnUpdateListener(func() {
|
||||
log.WithContext(ctx).Infof("your management version, \"%s\", is outdated, a new management version is available. Learn more here: https://github.com/netbirdio/netbird/releases", version.NetbirdVersion())
|
||||
})
|
||||
s.startDomainCleanup(srvCtx)
|
||||
|
||||
return nil
|
||||
}
|
||||
func (s *BaseServer) startDomainCleanup(ctx context.Context) {
|
||||
if s.domainCleanupStop != nil {
|
||||
return
|
||||
}
|
||||
mgr := s.ReverseProxyDomainManager()
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
done := make(chan struct{})
|
||||
s.domainCleanupStop = func() {
|
||||
cancel()
|
||||
<-done
|
||||
}
|
||||
go func() {
|
||||
defer close(done)
|
||||
mgr.RunValidationCleanup(ctx)
|
||||
}()
|
||||
}
|
||||
|
||||
// setupTLS resolves the listener's TLS source: an injected config wins over the HttpConfig certificate settings
|
||||
func (s *BaseServer) setupTLS(ctx context.Context) (bool, error) {
|
||||
switch {
|
||||
case s.tlsConfig != nil:
|
||||
return true, nil
|
||||
case s.Config.HttpConfig.LetsEncryptDomain != "":
|
||||
certManager, err := encryption.CreateCertManager(s.Config.Datadir, s.Config.HttpConfig.LetsEncryptDomain)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed creating LetsEncrypt cert manager: %v", err)
|
||||
}
|
||||
s.certManager = certManager
|
||||
return true, nil
|
||||
case s.Config.HttpConfig.CertFile != "" && s.Config.HttpConfig.CertKey != "":
|
||||
tlsConfig, err := loadTLSConfig(s.Config.HttpConfig.CertFile, s.Config.HttpConfig.CertKey)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("cannot load TLS credentials: %v", err)
|
||||
return false, err
|
||||
}
|
||||
s.tlsConfig = tlsConfig
|
||||
return true, nil
|
||||
default:
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Stop attempts a graceful shutdown, waiting up to 5 seconds for active connections to finish
|
||||
func (s *BaseServer) Stop() error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
defer s.RunStopHandlers()
|
||||
if s.domainCleanupStop != nil {
|
||||
s.domainCleanupStop()
|
||||
}
|
||||
|
||||
s.IntegratedValidator().Stop(ctx)
|
||||
if s.GeoLocationManager() != nil {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user