mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-06 13:39:07 +02:00
Merge branch 'main' into embedded-vnc
# Conflicts: # client/ui/frontend/src/app.tsx # client/ui/frontend/src/modules/main/MainConnectionStatusSwitch.tsx # client/ui/i18n/locales/uk/common.json # go.sum
This commit is contained in:
@@ -54,6 +54,15 @@ func Execute() error {
|
||||
return rootCmd.Execute()
|
||||
}
|
||||
|
||||
// Customize hands the fully built root command to fn so an embedding binary
|
||||
// can extend or adjust the command tree — most commonly attaching its own
|
||||
// subcommands next to (or under) the built-in ones — before calling Execute.
|
||||
// The root command is constructed in this package's init, so Customize may be
|
||||
// called from the embedding binary's main at any point before Execute.
|
||||
func Customize(fn func(root *cobra.Command)) {
|
||||
fn(rootCmd)
|
||||
}
|
||||
|
||||
func init() {
|
||||
mgmtCmd.Flags().IntVar(&mgmtPort, "port", 80, "server port to listen on (defaults to 443 if TLS is enabled, 80 otherwise")
|
||||
mgmtCmd.Flags().BoolVar(&disableLegacyManagementPort, "disable-legacy-port", false, "disabling the old legacy port (33073)")
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
// TestCustomize verifies an embedding binary can extend the command tree: a
|
||||
// top-level command attached through the hook, and a subcommand attached under
|
||||
// the built-in admin group, are both resolvable exactly as Execute would
|
||||
// resolve them.
|
||||
func TestCustomize(t *testing.T) {
|
||||
topLevel := &cobra.Command{Use: "some-extra", RunE: func(*cobra.Command, []string) error { return nil }}
|
||||
nested := &cobra.Command{Use: "cluster", RunE: func(*cobra.Command, []string) error { return nil }}
|
||||
|
||||
Customize(func(root *cobra.Command) {
|
||||
root.AddCommand(topLevel)
|
||||
for _, c := range root.Commands() {
|
||||
if c.Name() == "admin" {
|
||||
c.AddCommand(nested)
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatal("admin command not found in the root tree")
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
rootCmd.RemoveCommand(topLevel)
|
||||
for _, c := range rootCmd.Commands() {
|
||||
if c.Name() == "admin" {
|
||||
c.RemoveCommand(nested)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
if found, _, err := rootCmd.Find([]string{"some-extra"}); err != nil || found != topLevel {
|
||||
t.Fatalf("top-level command not resolvable: found=%v err=%v", found, err)
|
||||
}
|
||||
if found, _, err := rootCmd.Find([]string{"admin", "cluster"}); err != nil || found != nested {
|
||||
t.Fatalf("nested admin subcommand not resolvable: found=%v err=%v", found, err)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -66,8 +66,8 @@ func TestExtractClusterFromFreeDomain(t *testing.T) {
|
||||
|
||||
func TestExtractClusterFromCustomDomains(t *testing.T) {
|
||||
customDomains := []*domain.Domain{
|
||||
{Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io"},
|
||||
{Domain: "proxy.corp.io", TargetCluster: "us1.proxy.netbird.io"},
|
||||
{Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io", Validated: true},
|
||||
{Domain: "proxy.corp.io", TargetCluster: "us1.proxy.netbird.io", Validated: true},
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
@@ -120,19 +120,49 @@ func TestExtractClusterFromCustomDomains(t *testing.T) {
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cluster, ok := extractClusterFromCustomDomains(tc.domain, customDomains)
|
||||
assert.Equal(t, tc.wantOK, ok)
|
||||
if ok {
|
||||
assert.Equal(t, tc.wantVal, cluster)
|
||||
cluster, match := extractClusterFromCustomDomains(tc.domain, customDomains)
|
||||
if !tc.wantOK {
|
||||
assert.Equal(t, customDomainNoMatch, match, "unrelated domain should not match any custom domain")
|
||||
return
|
||||
}
|
||||
assert.Equal(t, customDomainValidated, match, "validated custom domain should resolve a cluster")
|
||||
assert.Equal(t, tc.wantVal, cluster)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// An unvalidated row must never yield a cluster: the account has not shown it
|
||||
// controls the name, so no service may be bound to it.
|
||||
func TestExtractClusterFromCustomDomains_UnvalidatedDomainRefused(t *testing.T) {
|
||||
customDomains := []*domain.Domain{
|
||||
{Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io", Validated: false},
|
||||
}
|
||||
|
||||
for _, serviceDomain := range []string{"example.com", "app.example.com"} {
|
||||
t.Run(serviceDomain, func(t *testing.T) {
|
||||
cluster, match := extractClusterFromCustomDomains(serviceDomain, customDomains)
|
||||
assert.Equal(t, customDomainUnvalidated, match, "unvalidated row must be reported as such")
|
||||
assert.Empty(t, cluster, "unvalidated row must not resolve a cluster")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A more specific unvalidated row must not shadow a validated parent domain.
|
||||
func TestExtractClusterFromCustomDomains_ValidatedParentWinsOverUnvalidatedChild(t *testing.T) {
|
||||
customDomains := []*domain.Domain{
|
||||
{Domain: "example.com", TargetCluster: "cluster-generic", Validated: true},
|
||||
{Domain: "app.example.com", TargetCluster: "cluster-app", Validated: false},
|
||||
}
|
||||
|
||||
cluster, match := extractClusterFromCustomDomains("app.example.com", customDomains)
|
||||
assert.Equal(t, customDomainValidated, match)
|
||||
assert.Equal(t, "cluster-generic", cluster, "validated parent domain should provide the cluster")
|
||||
}
|
||||
|
||||
func TestExtractClusterFromCustomDomains_OverlappingDomains(t *testing.T) {
|
||||
customDomains := []*domain.Domain{
|
||||
{Domain: "example.com", TargetCluster: "cluster-generic"},
|
||||
{Domain: "app.example.com", TargetCluster: "cluster-app"},
|
||||
{Domain: "example.com", TargetCluster: "cluster-generic", Validated: true},
|
||||
{Domain: "app.example.com", TargetCluster: "cluster-app", Validated: true},
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
@@ -164,8 +194,8 @@ func TestExtractClusterFromCustomDomains_OverlappingDomains(t *testing.T) {
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cluster, ok := extractClusterFromCustomDomains(tc.domain, customDomains)
|
||||
assert.True(t, ok)
|
||||
cluster, match := extractClusterFromCustomDomains(tc.domain, customDomains)
|
||||
assert.Equal(t, customDomainValidated, match)
|
||||
assert.Equal(t, tc.wantVal, cluster)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -26,11 +28,14 @@ type store interface {
|
||||
GetAgentNetworkSettings(ctx context.Context, lockStrength nbstore.LockingStrength, accountID string) (*agentnetworkTypes.Settings, error)
|
||||
|
||||
GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error)
|
||||
GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error)
|
||||
ListFreeDomains(ctx context.Context, accountID string) ([]string, error)
|
||||
ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error)
|
||||
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 {
|
||||
@@ -105,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)
|
||||
@@ -125,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 {
|
||||
@@ -134,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 {
|
||||
@@ -150,6 +166,10 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName
|
||||
return nil, fmt.Errorf("target cluster %s is not available", targetCluster)
|
||||
}
|
||||
|
||||
if err := m.checkDomainAvailable(ctx, domainName); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Attempt an initial validation against the specified cluster only
|
||||
var validated bool
|
||||
if m.validator.IsValid(ctx, domainName, []string{targetCluster}) {
|
||||
@@ -166,6 +186,23 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName
|
||||
return d, nil
|
||||
}
|
||||
|
||||
// checkDomainAvailable reports whether the domain is free to claim. The unique
|
||||
// index on the column is the real guard; this turns the violation into a
|
||||
// conflict the caller can act on instead of a database error, and says nothing
|
||||
// about which account holds the domain.
|
||||
func (m Manager) checkDomainAvailable(ctx context.Context, domainName string) error {
|
||||
_, err := m.store.GetCustomDomainByName(ctx, domainName)
|
||||
if err == nil {
|
||||
return status.Errorf(status.AlreadyExists, "domain %s is already registered", domainName)
|
||||
}
|
||||
|
||||
if sErr, ok := status.FromError(err); ok && sErr.Type() == status.NotFound {
|
||||
return nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("look up domain: %w", err)
|
||||
}
|
||||
|
||||
func (m Manager) DeleteDomain(ctx context.Context, accountID, userID, domainID string) error {
|
||||
ok, ctx, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Services, operations.Delete)
|
||||
if err != nil {
|
||||
@@ -203,7 +240,9 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID
|
||||
log.WithFields(log.Fields{
|
||||
"accountID": accountID,
|
||||
"domainID": domainID,
|
||||
}).WithError(err).Error("validate domain")
|
||||
"userID": userID,
|
||||
}).Error("validate domain: permission denied")
|
||||
return
|
||||
}
|
||||
|
||||
log.WithFields(log.Fields{
|
||||
@@ -219,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
|
||||
@@ -239,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 {
|
||||
@@ -298,9 +346,12 @@ func (m Manager) DeriveClusterFromDomain(ctx context.Context, accountID, domain
|
||||
return "", fmt.Errorf("list custom domains: %w", err)
|
||||
}
|
||||
|
||||
targetCluster, valid := extractClusterFromCustomDomains(domain, customDomains)
|
||||
if valid {
|
||||
targetCluster, match := extractClusterFromCustomDomains(domain, customDomains)
|
||||
switch match {
|
||||
case customDomainValidated:
|
||||
return targetCluster, nil
|
||||
case customDomainUnvalidated:
|
||||
return "", status.Errorf(status.PreconditionFailed, "domain %s is not validated", domain)
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("domain %s does not match any available proxy cluster", domain)
|
||||
@@ -363,19 +414,46 @@ func (m Manager) reservedGatewayAddress(ctx context.Context, accountID string) (
|
||||
return settings.ProxyAddress, nil
|
||||
}
|
||||
|
||||
func extractClusterFromCustomDomains(serviceDomain string, customDomains []*domain.Domain) (string, bool) {
|
||||
// customDomainMatch describes how a service domain relates to the account's
|
||||
// custom domain rows.
|
||||
type customDomainMatch int
|
||||
|
||||
const (
|
||||
customDomainNoMatch customDomainMatch = iota
|
||||
customDomainUnvalidated
|
||||
customDomainValidated
|
||||
)
|
||||
|
||||
// extractClusterFromCustomDomains finds the longest custom domain covering the
|
||||
// service domain and reports its target cluster. Only a validated row yields a
|
||||
// cluster: until the CNAME check has passed the account has not shown it
|
||||
// controls the name, so no traffic may be routed for it.
|
||||
func extractClusterFromCustomDomains(serviceDomain string, customDomains []*domain.Domain) (string, customDomainMatch) {
|
||||
bestCluster := ""
|
||||
bestLen := -1
|
||||
matched := false
|
||||
for _, cd := range customDomains {
|
||||
if serviceDomain != cd.Domain && !strings.HasSuffix(serviceDomain, "."+cd.Domain) {
|
||||
continue
|
||||
}
|
||||
matched = true
|
||||
if !cd.Validated {
|
||||
continue
|
||||
}
|
||||
if l := len(cd.Domain); l > bestLen {
|
||||
bestLen = l
|
||||
bestCluster = cd.TargetCluster
|
||||
}
|
||||
}
|
||||
return bestCluster, bestLen >= 0
|
||||
|
||||
switch {
|
||||
case bestLen >= 0:
|
||||
return bestCluster, customDomainValidated
|
||||
case matched:
|
||||
return "", customDomainUnvalidated
|
||||
default:
|
||||
return "", customDomainNoMatch
|
||||
}
|
||||
}
|
||||
|
||||
// ExtractClusterFromFreeDomain extracts the cluster address from a free domain.
|
||||
|
||||
@@ -0,0 +1,321 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel/metric/noop"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
"github.com/netbirdio/netbird/management/server/mock_server"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
nbstore "github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
const (
|
||||
testCluster = "eu.proxy.test"
|
||||
accountA = "account-a"
|
||||
accountAUser = "account-a-admin"
|
||||
accountB = "account-b"
|
||||
accountBUser = "account-b-admin"
|
||||
accountAMember = "account-a-member"
|
||||
)
|
||||
|
||||
// stubResolver answers CNAME lookups from a table the test controls, so a
|
||||
// domain can point at the cluster or nowhere without touching a real resolver.
|
||||
type stubResolver struct {
|
||||
mu sync.Mutex
|
||||
cnames map[string]string
|
||||
}
|
||||
|
||||
func (r *stubResolver) LookupCNAME(_ context.Context, host string) (string, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
cname, ok := r.cnames[host]
|
||||
if !ok {
|
||||
return "", fmt.Errorf("lookup %s: no such host", host)
|
||||
}
|
||||
return cname + ".", nil
|
||||
}
|
||||
|
||||
func (r *stubResolver) set(host, cname string) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.cnames[host] = cname
|
||||
}
|
||||
|
||||
type domainTestEnv struct {
|
||||
manager Manager
|
||||
store nbstore.Store
|
||||
resolver *stubResolver
|
||||
}
|
||||
|
||||
// setupDomainTest builds the domain manager on a real SQLite store with two
|
||||
// accounts and one active public proxy cluster.
|
||||
func setupDomainTest(t *testing.T) *domainTestEnv {
|
||||
t.Helper()
|
||||
|
||||
ctx := context.Background()
|
||||
testStore, cleanup, err := nbstore.NewTestStoreFromSQL(ctx, "", t.TempDir())
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(cleanup)
|
||||
|
||||
for accountID, userID := range map[string]string{accountA: accountAUser, accountB: accountBUser} {
|
||||
users := map[string]*types.User{
|
||||
userID: {
|
||||
Id: userID,
|
||||
AccountID: accountID,
|
||||
Role: types.UserRoleAdmin,
|
||||
},
|
||||
}
|
||||
if accountID == accountA {
|
||||
// A real member of the account whose role denies Services:Create, so
|
||||
// permission denial is exercised as ok=false rather than as a lookup
|
||||
// error for a user who is not in the account at all.
|
||||
users[accountAMember] = &types.User{
|
||||
Id: accountAMember,
|
||||
AccountID: accountID,
|
||||
Role: types.UserRoleUser,
|
||||
}
|
||||
}
|
||||
|
||||
require.NoError(t, testStore.SaveAccount(ctx, &types.Account{
|
||||
Id: accountID,
|
||||
CreatedBy: userID,
|
||||
Settings: &types.Settings{},
|
||||
Users: users,
|
||||
}))
|
||||
}
|
||||
|
||||
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)
|
||||
require.NoError(t, err)
|
||||
|
||||
resolver := &stubResolver{cnames: make(map[string]string)}
|
||||
|
||||
mgr := Manager{
|
||||
store: testStore,
|
||||
proxyManager: proxyMgr,
|
||||
validator: domain.Validator{Resolver: resolver},
|
||||
permissionsManager: permissions.NewManager(testStore),
|
||||
accountManager: &mock_server.MockAccountManager{
|
||||
StoreEventFunc: func(context.Context, string, string, string, activity.ActivityDescriber, map[string]any) {},
|
||||
},
|
||||
}
|
||||
|
||||
return &domainTestEnv{manager: mgr, store: testStore, resolver: resolver}
|
||||
}
|
||||
|
||||
// storedDomain reads a domain row back through the store so assertions are made
|
||||
// on what was persisted rather than on the value the manager returned.
|
||||
func storedDomain(t *testing.T, s nbstore.Store, accountID, domainName string) *domain.Domain {
|
||||
t.Helper()
|
||||
|
||||
domains, err := s.ListCustomDomains(context.Background(), accountID)
|
||||
require.NoError(t, err)
|
||||
for _, d := range domains {
|
||||
if d.Domain == domainName {
|
||||
return d
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// A domain whose CNAME check fails is stored unvalidated and must not resolve a
|
||||
// cluster, which is what service creation gates on.
|
||||
func TestCreateDomain_FailedLookupIsNotServable(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)
|
||||
assert.False(t, created.Validated, "a domain whose CNAME lookup fails must not be created validated")
|
||||
|
||||
stored := storedDomain(t, env.store, accountA, "apps.example.com")
|
||||
require.NotNil(t, stored, "domain row should exist")
|
||||
assert.False(t, stored.Validated, "persisted row must be unvalidated")
|
||||
|
||||
cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "apps.example.com")
|
||||
require.Error(t, err, "an unvalidated domain must not resolve a cluster")
|
||||
assert.Empty(t, cluster)
|
||||
assert.Contains(t, err.Error(), "not validated", "error should tell the caller what to fix")
|
||||
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "error should be a typed status error")
|
||||
assert.Equal(t, status.PreconditionFailed, sErr.Type())
|
||||
|
||||
_, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "sub.apps.example.com")
|
||||
assert.Error(t, err, "subdomains of an unvalidated custom domain are not servable either")
|
||||
}
|
||||
|
||||
// A second account claiming a registered domain gets a clean conflict, not a
|
||||
// database error surfaced as a 500.
|
||||
func TestCreateDomain_DuplicateIsAConflict(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
|
||||
_, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "shared.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = env.manager.CreateDomain(ctx, accountB, accountBUser, "shared.example.com", testCluster)
|
||||
require.Error(t, err)
|
||||
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "conflict must be a typed status error, not a raw database error")
|
||||
assert.Equal(t, status.AlreadyExists, sErr.Type(), "conflict should map to 409, not 500")
|
||||
assert.NotContains(t, sErr.Message, accountA, "the response must not reveal the holding account")
|
||||
|
||||
assert.Nil(t, storedDomain(t, env.store, accountB, "shared.example.com"), "no row should be written on conflict")
|
||||
}
|
||||
|
||||
// The same account re-adding one of its own domains is a conflict too.
|
||||
func TestCreateDomain_SameAccountDuplicateIsAConflict(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
|
||||
_, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "dup.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = env.manager.CreateDomain(ctx, accountA, accountAUser, "dup.example.com", testCluster)
|
||||
require.Error(t, err)
|
||||
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, status.AlreadyExists, sErr.Type())
|
||||
}
|
||||
|
||||
// The negative control: a validated domain still derives its cluster, for the
|
||||
// bare name and for subdomains, exactly as before.
|
||||
func TestCreateDomain_ValidatedDomainDerivesCluster(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
env.resolver.set("validation.valid.example.com", testCluster)
|
||||
|
||||
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "valid.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
require.True(t, created.Validated, "a matching CNAME should validate on create")
|
||||
|
||||
cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "valid.example.com")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testCluster, cluster)
|
||||
|
||||
cluster, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "app.valid.example.com")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testCluster, cluster, "subdomains of a validated custom domain resolve too")
|
||||
}
|
||||
|
||||
// Validating a domain flips the gate: the same lookup that failed before now
|
||||
// resolves a cluster.
|
||||
func TestValidateDomain_UnlocksClusterDerivation(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
|
||||
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "later.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
require.False(t, created.Validated)
|
||||
|
||||
_, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "later.example.com")
|
||||
require.Error(t, err)
|
||||
|
||||
env.resolver.set("validation.later.example.com", testCluster)
|
||||
env.manager.ValidateDomain(ctx, accountA, accountAUser, created.ID)
|
||||
|
||||
require.True(t, storedDomain(t, env.store, accountA, "later.example.com").Validated)
|
||||
|
||||
cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "later.example.com")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testCluster, cluster)
|
||||
}
|
||||
|
||||
// Free cluster domains are unaffected by the custom domain gate.
|
||||
func TestDeriveClusterFromDomain_FreeDomainUnaffected(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
|
||||
cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "myapp.abc123."+testCluster)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testCluster, cluster)
|
||||
}
|
||||
|
||||
// The manager pre-check exists to turn a conflict into a 409, but the unique
|
||||
// index on the column is what actually guarantees the domain is claimed once.
|
||||
//
|
||||
// Two requests can clear the pre-check concurrently and race to the insert.
|
||||
// Inserting twice through the store reaches the same code path the loser of
|
||||
// that race takes, without the nondeterminism of driving it from goroutines,
|
||||
// and the loser must still see a conflict rather than an internal error.
|
||||
func TestStore_DuplicateDomainRejectedByIndexAsConflict(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
|
||||
_, err := env.store.CreateCustomDomain(ctx, accountA, "indexed.example.com", testCluster, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = env.store.CreateCustomDomain(ctx, accountB, "indexed.example.com", testCluster, false)
|
||||
require.Error(t, err, "the unique index must reject the same domain in a second account")
|
||||
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "the losing insert must return a typed status error")
|
||||
assert.Equal(t, status.AlreadyExists, sErr.Type(), "a lost race is a 409, not a 500")
|
||||
}
|
||||
|
||||
// Validation is what decides whether a domain routes traffic, so a caller
|
||||
// without permission to it must not be able to flip the flag. The check logged
|
||||
// the denial and then carried on, which was inert while nothing read Validated
|
||||
// and is not once cluster derivation gates on it.
|
||||
func TestValidateDomain_PermissionDeniedDoesNotValidate(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
env := setupDomainTest(t)
|
||||
|
||||
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "guarded.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
require.False(t, created.Validated)
|
||||
|
||||
// The CNAME is in place, so the only thing standing between this caller and
|
||||
// a validated domain is the permission check.
|
||||
env.resolver.set("validation.guarded.example.com", testCluster)
|
||||
|
||||
env.manager.ValidateDomain(ctx, accountA, accountAMember, created.ID)
|
||||
|
||||
stored := storedDomain(t, env.store, accountA, "guarded.example.com")
|
||||
require.NotNil(t, stored)
|
||||
assert.False(t, stored.Validated, "a caller without permission must not validate the domain")
|
||||
|
||||
_, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "guarded.example.com")
|
||||
assert.Error(t, err, "the domain must still be unservable")
|
||||
}
|
||||
|
||||
// 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)
|
||||
|
||||
created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "racy.example.com", testCluster)
|
||||
require.NoError(t, err)
|
||||
|
||||
stale := storedDomain(t, env.store, accountA, "racy.example.com")
|
||||
require.NotNil(t, stale)
|
||||
|
||||
require.NoError(t, env.manager.DeleteDomain(ctx, accountA, accountAUser, created.ID))
|
||||
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.
|
||||
stale.Validated = true
|
||||
_, err = env.store.UpdateCustomDomain(ctx, accountA, stale)
|
||||
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"
|
||||
@@ -184,6 +185,10 @@ func (s *stubStore) GetCustomDomain(context.Context, string, string) (*domain.Do
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) GetCustomDomainByName(context.Context, string) (*domain.Domain, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) ListFreeDomains(context.Context, string) ([]string, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
@@ -204,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")
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel/metric/noop"
|
||||
|
||||
domainmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain/manager"
|
||||
proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager"
|
||||
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"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
const validationTestCluster = "eu.proxy.test"
|
||||
|
||||
// withRealDomainManager swaps the stub cluster deriver for the real domain
|
||||
// manager backed by the same store, so service creation is gated by the actual
|
||||
// domain rows rather than by a test double that always agrees.
|
||||
func withRealDomainManager(t *testing.T, mgr *Manager, testStore store.Store) {
|
||||
t.Helper()
|
||||
|
||||
ctx := context.Background()
|
||||
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)
|
||||
require.NoError(t, err)
|
||||
|
||||
accountMgr := &mock_server.MockAccountManager{
|
||||
StoreEventFunc: func(context.Context, string, string, string, activity.ActivityDescriber, map[string]any) {},
|
||||
}
|
||||
mgr.clusterDeriver = domainmanager.NewManager(testStore, proxyMgr, permissions.NewManager(testStore), accountMgr)
|
||||
}
|
||||
|
||||
func newTestService(domain string) *rpservice.Service {
|
||||
return &rpservice.Service{
|
||||
Name: "test-service",
|
||||
Domain: domain,
|
||||
Enabled: true,
|
||||
Mode: rpservice.ModeHTTP,
|
||||
Targets: []*rpservice.Target{{
|
||||
Host: "10.0.0.1",
|
||||
Port: 8080,
|
||||
Protocol: "http",
|
||||
TargetId: testPeerID,
|
||||
TargetType: "peer",
|
||||
Enabled: true,
|
||||
}},
|
||||
}
|
||||
}
|
||||
|
||||
// A service must not bind to a domain the account has not validated, and
|
||||
// nothing may be persisted for the attempt.
|
||||
func TestCreateService_RefusesUnvalidatedDomain(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupIntegrationTest(t)
|
||||
withRealDomainManager(t, mgr, testStore)
|
||||
|
||||
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "unproven.example.com", validationTestCluster, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = mgr.CreateService(ctx, testAccountID, testUserID, newTestService("unproven.example.com"))
|
||||
require.Error(t, err, "an unvalidated domain must not bind a service")
|
||||
assert.Contains(t, err.Error(), "not validated", "the API error should name the actual problem")
|
||||
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "error should be a typed status error")
|
||||
assert.Equal(t, status.PreconditionFailed, sErr.Type())
|
||||
|
||||
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, services, "no service row should be written for a refused domain")
|
||||
}
|
||||
|
||||
// The negative control: a validated domain still binds a service and derives
|
||||
// its cluster exactly as before.
|
||||
func TestCreateService_ValidatedDomainBindsService(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupIntegrationTest(t)
|
||||
withRealDomainManager(t, mgr, testStore)
|
||||
|
||||
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true)
|
||||
require.NoError(t, err)
|
||||
|
||||
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.proven.example.com"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, validationTestCluster, created.ProxyCluster, "service should bind to the domain's target cluster")
|
||||
|
||||
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, services, 1, "the service should be persisted")
|
||||
assert.Equal(t, "app.proven.example.com", services[0].Domain)
|
||||
}
|
||||
|
||||
// An update must not be a way around the creation gate: moving a live service
|
||||
// onto an unvalidated domain has to fail rather than silently keep the old
|
||||
// cluster and start serving the new hostname.
|
||||
func TestUpdateService_RefusesMoveToUnvalidatedDomain(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupIntegrationTest(t)
|
||||
withRealDomainManager(t, mgr, testStore)
|
||||
|
||||
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true)
|
||||
require.NoError(t, err)
|
||||
_, err = testStore.CreateCustomDomain(ctx, testAccountID, "unproven.example.com", validationTestCluster, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.proven.example.com"))
|
||||
require.NoError(t, err)
|
||||
|
||||
moved := *created
|
||||
moved.Domain = "app.unproven.example.com"
|
||||
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, &moved)
|
||||
require.Error(t, err, "moving to an unvalidated domain must fail")
|
||||
assert.Contains(t, err.Error(), "not validated")
|
||||
|
||||
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "app.proven.example.com", stored.Domain, "the service must keep its original domain")
|
||||
}
|
||||
@@ -606,16 +606,19 @@ func (m *Manager) resolveEffectiveCluster(ctx context.Context, accountID string,
|
||||
return existing.ProxyCluster, nil
|
||||
}
|
||||
|
||||
if m.clusterDeriver != nil {
|
||||
derived, err := m.clusterDeriver.DeriveClusterFromDomain(ctx, accountID, svc.Domain)
|
||||
if err != nil {
|
||||
log.WithError(err).Warnf("could not derive cluster from domain %s", svc.Domain)
|
||||
} else {
|
||||
return derived, nil
|
||||
}
|
||||
if m.clusterDeriver == nil {
|
||||
return existing.ProxyCluster, nil
|
||||
}
|
||||
|
||||
return existing.ProxyCluster, nil
|
||||
// Falling back to the old cluster here would let an update move a service
|
||||
// onto a domain the account has not validated, bypassing the check that
|
||||
// creation makes.
|
||||
derived, err := m.clusterDeriver.DeriveClusterFromDomain(ctx, accountID, svc.Domain)
|
||||
if err != nil {
|
||||
return "", status.Errorf(status.PreconditionFailed, "could not derive cluster from domain %s: %v", svc.Domain, err)
|
||||
}
|
||||
|
||||
return derived, nil
|
||||
}
|
||||
|
||||
func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.Store, accountID string, service *service.Service, updateInfo *serviceUpdateInfo, customPorts *bool, effectiveCluster string) error {
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
})
|
||||
}
|
||||
@@ -66,7 +66,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,6 +75,7 @@ type BaseServer struct {
|
||||
grpcExtensions []GRPCExtension
|
||||
|
||||
listener net.Listener
|
||||
tlsConfig *tls.Config
|
||||
certManager *autocert.Manager
|
||||
update *version.Update
|
||||
|
||||
@@ -94,6 +96,7 @@ type Config struct {
|
||||
DisableGeoliteUpdate bool
|
||||
UserDeleteFromIDPEnabled bool
|
||||
AutoResolveDomains bool
|
||||
TLSConfig *tls.Config
|
||||
}
|
||||
|
||||
// NewServer initializes and configures a new Server instance
|
||||
@@ -110,6 +113,7 @@ func NewServer(cfg *Config) *BaseServer {
|
||||
disableLegacyManagementPort: cfg.DisableLegacyManagementPort,
|
||||
mgmtMetricsPort: cfg.MgmtMetricsPort,
|
||||
autoResolveDomains: cfg.AutoResolveDomains,
|
||||
tlsConfig: cfg.TLSConfig,
|
||||
}
|
||||
s.container[ContainerKeyBaseServer] = s
|
||||
|
||||
@@ -139,21 +143,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 +207,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 +228,59 @@ 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()
|
||||
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.
|
||||
@@ -1223,6 +1225,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 +1237,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,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))
|
||||
}
|
||||
@@ -719,8 +719,10 @@ func (am *DefaultAccountManager) schedulePeerLoginExpiration(ctx context.Context
|
||||
log.WithContext(ctx).Tracef("peer login expiration job for account %s is already scheduled", accountID)
|
||||
return
|
||||
}
|
||||
// The job outlives the request that arms it, so it must not inherit the request's cancellation.
|
||||
jobCtx := context.WithoutCancel(ctx)
|
||||
if nextRun, ok := am.getNextPeerExpiration(ctx, accountID); ok {
|
||||
go am.peerLoginExpiry.Schedule(ctx, nextRun, accountID, am.peerLoginExpirationJob(ctx, accountID))
|
||||
go am.peerLoginExpiry.Schedule(jobCtx, nextRun, accountID, am.peerLoginExpirationJob(jobCtx, accountID))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -752,8 +754,9 @@ func (am *DefaultAccountManager) peerInactivityExpirationJob(ctx context.Context
|
||||
// checkAndSchedulePeerInactivityExpiration periodically checks for inactive peers to end their sessions
|
||||
func (am *DefaultAccountManager) checkAndSchedulePeerInactivityExpiration(ctx context.Context, accountID string) {
|
||||
am.peerInactivityExpiry.Cancel(ctx, []string{accountID})
|
||||
jobCtx := context.WithoutCancel(ctx)
|
||||
if nextRun, ok := am.getNextInactivePeerExpiration(ctx, accountID); ok {
|
||||
go am.peerInactivityExpiry.Schedule(ctx, nextRun, accountID, am.peerInactivityExpirationJob(ctx, accountID))
|
||||
go am.peerInactivityExpiry.Schedule(jobCtx, nextRun, accountID, am.peerInactivityExpirationJob(jobCtx, accountID))
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1920,6 +1920,154 @@ func TestDefaultAccountManager_MarkPeerConnected_PeerLoginExpiration(t *testing.
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_SchedulePeerLoginExpiration_IncludesOfflinePeers(t *testing.T) {
|
||||
manager, updateManager, err := createManager(t)
|
||||
require.NoError(t, err, "unable to create account manager")
|
||||
|
||||
accountID, err := manager.GetAccountIDByUserID(context.Background(), auth.UserAuth{UserId: userID})
|
||||
require.NoError(t, err, "unable to create an account")
|
||||
|
||||
connectedKey, offlineKey := addExpiringPeers(t, manager)
|
||||
_, err = manager.UpdateAccountSettings(context.Background(), accountID, userID, &types.Settings{
|
||||
PeerLoginExpiration: time.Hour,
|
||||
PeerLoginExpirationEnabled: true,
|
||||
Extra: &types.ExtraSettings{},
|
||||
})
|
||||
require.NoError(t, err, "expecting to update account settings successfully but got error")
|
||||
manager.peerLoginExpiry.CancelAll(context.Background())
|
||||
|
||||
// The connected peer logged in just now, so a job computed from connected peers alone
|
||||
// would be armed for an hour. The offline peer's login expires in two seconds; a
|
||||
// reconnect of that peer must not have to wait for the connected peer's tick.
|
||||
now := time.Now().UTC()
|
||||
setPeerLogin(t, manager, accountID, connectedKey, true, now)
|
||||
setPeerLogin(t, manager, accountID, offlineKey, false, now.Add(-time.Hour+2*time.Second))
|
||||
|
||||
offlinePeer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, offlineKey)
|
||||
require.NoError(t, err)
|
||||
updateManager.CreateChannel(context.Background(), offlinePeer.ID)
|
||||
|
||||
manager.peerLoginExpiry = NewDefaultScheduler()
|
||||
t.Cleanup(func() { manager.peerLoginExpiry.CancelAll(context.Background()) })
|
||||
manager.schedulePeerLoginExpiration(context.Background(), accountID)
|
||||
|
||||
// The flag is committed per peer before the disconnect fans out, so wait for both.
|
||||
require.Eventually(t, func() bool {
|
||||
peer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, offlineKey)
|
||||
return err == nil && peer.Status.LoginExpired && !updateManager.HasChannel(offlinePeer.ID)
|
||||
}, 10*time.Second, 100*time.Millisecond, "offline peer should be expired and disconnected at its own deadline")
|
||||
|
||||
connectedPeer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, connectedKey)
|
||||
require.NoError(t, err)
|
||||
assert.False(t, connectedPeer.Status.LoginExpired, "connected peer with a fresh login must not expire")
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_SchedulePeerLoginExpiration_DetachesRequestContext(t *testing.T) {
|
||||
manager, _, err := createManager(t)
|
||||
require.NoError(t, err, "unable to create account manager")
|
||||
|
||||
accountID, err := manager.GetAccountIDByUserID(context.Background(), auth.UserAuth{UserId: userID})
|
||||
require.NoError(t, err, "unable to create an account")
|
||||
connectedKey, _ := addExpiringPeers(t, manager)
|
||||
setPeerLogin(t, manager, accountID, connectedKey, true, time.Now().UTC())
|
||||
|
||||
scheduled := make(chan context.Context, 1)
|
||||
manager.peerLoginExpiry = &MockScheduler{
|
||||
IsSchedulerRunningFunc: func(string) bool { return false },
|
||||
ScheduleFunc: func(ctx context.Context, _ time.Duration, _ string, _ func() (time.Duration, bool)) {
|
||||
scheduled <- ctx
|
||||
},
|
||||
}
|
||||
|
||||
requestCtx, cancel := context.WithCancel(context.Background())
|
||||
manager.schedulePeerLoginExpiration(requestCtx, accountID)
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case jobCtx := <-scheduled:
|
||||
assert.NoError(t, jobCtx.Err(), "the expiration job must outlive the request that armed it")
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("timeout while waiting for the job to be scheduled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_ExpireAndUpdatePeers_SkipsPeerThatLoggedInAgain(t *testing.T) {
|
||||
manager, updateManager, err := createManager(t)
|
||||
require.NoError(t, err, "unable to create account manager")
|
||||
|
||||
accountID, err := manager.GetAccountIDByUserID(context.Background(), auth.UserAuth{UserId: userID})
|
||||
require.NoError(t, err, "unable to create an account")
|
||||
|
||||
reloggedKey, staleKey := addExpiringPeers(t, manager)
|
||||
_, err = manager.UpdateAccountSettings(context.Background(), accountID, userID, &types.Settings{
|
||||
PeerLoginExpiration: time.Hour,
|
||||
PeerLoginExpirationEnabled: true,
|
||||
Extra: &types.ExtraSettings{},
|
||||
})
|
||||
require.NoError(t, err, "expecting to update account settings successfully but got error")
|
||||
manager.peerLoginExpiry.CancelAll(context.Background())
|
||||
|
||||
expiredLogin := time.Now().UTC().Add(-2 * time.Hour)
|
||||
setPeerLogin(t, manager, accountID, reloggedKey, true, expiredLogin)
|
||||
setPeerLogin(t, manager, accountID, staleKey, true, expiredLogin)
|
||||
|
||||
expiredPeers, err := manager.getExpiredPeers(context.Background(), accountID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, expiredPeers, 2, "both peers should be due for expiration")
|
||||
|
||||
// The job holds the candidate list while one peer completes a fresh login, which
|
||||
// moves its deadline into the future and must win over the stale candidate entry.
|
||||
setPeerLogin(t, manager, accountID, reloggedKey, true, time.Now().UTC())
|
||||
|
||||
reloggedPeer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, reloggedKey)
|
||||
require.NoError(t, err)
|
||||
stalePeer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, staleKey)
|
||||
require.NoError(t, err)
|
||||
updateManager.CreateChannel(context.Background(), reloggedPeer.ID)
|
||||
updateManager.CreateChannel(context.Background(), stalePeer.ID)
|
||||
|
||||
err = manager.expireAndUpdatePeers(context.Background(), accountID, expiredPeers, peerExpirationSessionExpired)
|
||||
require.NoError(t, err)
|
||||
|
||||
reloggedPeer, err = manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, reloggedKey)
|
||||
require.NoError(t, err)
|
||||
assert.False(t, reloggedPeer.Status.LoginExpired, "a peer that logged in again must not be flagged from the stale candidate list")
|
||||
assert.True(t, reloggedPeer.Status.Connected, "the re-logged peer must keep its connected status")
|
||||
assert.True(t, updateManager.HasChannel(reloggedPeer.ID), "the re-logged peer's update channel must stay open")
|
||||
|
||||
stalePeer, err = manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, staleKey)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, stalePeer.Status.LoginExpired, "a peer that is still due must be flagged")
|
||||
assert.False(t, updateManager.HasChannel(stalePeer.ID), "the expired peer's update channel must be closed")
|
||||
}
|
||||
|
||||
// addExpiringPeers registers two SSO peers with login expiration enabled and returns their public keys.
|
||||
func addExpiringPeers(t *testing.T, manager *DefaultAccountManager) (string, string) {
|
||||
t.Helper()
|
||||
keys := make([]string, 0, 2)
|
||||
for _, hostname := range []string{"connected-peer", "offline-peer"} {
|
||||
key, err := wgtypes.GenerateKey()
|
||||
require.NoError(t, err, "unable to generate WireGuard key")
|
||||
_, _, _, _, err = manager.AddPeer(context.Background(), "", "", userID, &nbpeer.Peer{
|
||||
Key: key.PublicKey().String(),
|
||||
Meta: nbpeer.PeerSystemMeta{Hostname: hostname},
|
||||
LoginExpirationEnabled: true,
|
||||
}, false)
|
||||
require.NoError(t, err, "unable to add peer")
|
||||
keys = append(keys, key.PublicKey().String())
|
||||
}
|
||||
return keys[0], keys[1]
|
||||
}
|
||||
|
||||
func setPeerLogin(t *testing.T, manager *DefaultAccountManager, accountID, peerKey string, connected bool, lastLogin time.Time) {
|
||||
t.Helper()
|
||||
peer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, peerKey)
|
||||
require.NoError(t, err)
|
||||
peer.Status.Connected = connected
|
||||
peer.LastLogin = &lastLogin
|
||||
require.NoError(t, manager.Store.SavePeer(context.Background(), accountID, peer))
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_MarkPeerDisconnected_SchedulesInactivityExpiration(t *testing.T) {
|
||||
manager, _, err := createManager(t)
|
||||
require.NoError(t, err, "unable to create account manager")
|
||||
@@ -2702,7 +2850,7 @@ func TestAccount_GetNextPeerExpiration(t *testing.T) {
|
||||
expectedNextExpiration: time.Duration(0),
|
||||
},
|
||||
{
|
||||
name: "No connected peers, no expiration",
|
||||
name: "Offline peer with expiration, return expiration",
|
||||
peers: map[string]*nbpeer.Peer{
|
||||
"peer-1": {
|
||||
Status: &nbpeer.PeerStatus{
|
||||
@@ -2721,8 +2869,33 @@ func TestAccount_GetNextPeerExpiration(t *testing.T) {
|
||||
},
|
||||
expiration: time.Second,
|
||||
expirationEnabled: false,
|
||||
expectedNextRun: false,
|
||||
expectedNextExpiration: time.Duration(0),
|
||||
expectedNextRun: true,
|
||||
expectedNextExpiration: time.Second,
|
||||
},
|
||||
{
|
||||
name: "Offline peer with the earliest deadline defines the next run",
|
||||
peers: map[string]*nbpeer.Peer{
|
||||
"peer-1": {
|
||||
Status: &nbpeer.PeerStatus{
|
||||
Connected: true,
|
||||
},
|
||||
LoginExpirationEnabled: true,
|
||||
LastLogin: util.ToPtr(time.Now().UTC()),
|
||||
UserID: userID,
|
||||
},
|
||||
"peer-2": {
|
||||
Status: &nbpeer.PeerStatus{
|
||||
Connected: false,
|
||||
},
|
||||
LoginExpirationEnabled: true,
|
||||
LastLogin: util.ToPtr(time.Now().UTC().Add(-50 * time.Minute)),
|
||||
UserID: userID,
|
||||
},
|
||||
},
|
||||
expiration: time.Hour,
|
||||
expirationEnabled: true,
|
||||
expectedNextRun: true,
|
||||
expectedNextExpiration: 10 * time.Minute,
|
||||
},
|
||||
{
|
||||
name: "Connected peers with disabled expiration, no expiration",
|
||||
@@ -3343,7 +3516,13 @@ func buildTestManager(t testing.TB, store store.Store, nmdataStore *networkmapdb
|
||||
|
||||
eventStore := &activity.InMemoryEventStore{}
|
||||
|
||||
metrics, err := telemetry.NewDefaultAppMetrics(context.Background())
|
||||
// Everything built here watches this context; cancelling it on cleanup stops
|
||||
// the metrics flushers, caches and controllers instead of leaking them for
|
||||
// the rest of the package run.
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
|
||||
metrics, err := telemetry.NewDefaultAppMetrics(ctx)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
@@ -3362,8 +3541,6 @@ func buildTestManager(t testing.TB, store store.Store, nmdataStore *networkmapdb
|
||||
Return(nil).
|
||||
AnyTimes()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
cacheStore, err := cache.NewStore(ctx, 100*time.Millisecond, 300*time.Millisecond, 100)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
|
||||
@@ -284,6 +284,9 @@ const (
|
||||
// AgentNetworkSettingsDeleted indicates that a user deleted the Agent Network account settings, releasing the endpoint
|
||||
AgentNetworkSettingsDeleted Activity = 142
|
||||
|
||||
// CustomDomainValidationExpired indicates that an unvalidated domain registration expired.
|
||||
CustomDomainValidationExpired Activity = 143
|
||||
|
||||
AccountDeleted Activity = 99999
|
||||
)
|
||||
|
||||
@@ -461,9 +464,10 @@ var activityMap = map[Activity]Code{
|
||||
AccountMetricsPushEnabled: {"Account metrics push enabled", "account.setting.metrics.push.enable"},
|
||||
AccountMetricsPushDisabled: {"Account metrics push disabled", "account.setting.metrics.push.disable"},
|
||||
|
||||
DomainAdded: {"Domain added", "domain.add"},
|
||||
DomainDeleted: {"Domain deleted", "domain.delete"},
|
||||
DomainValidated: {"Domain validated", "domain.validate"},
|
||||
DomainAdded: {"Domain added", "domain.add"},
|
||||
DomainDeleted: {"Domain deleted", "domain.delete"},
|
||||
DomainValidated: {"Domain validated", "domain.validate"},
|
||||
CustomDomainValidationExpired: {"Unvalidated domain registration expired", "domain.validation.expire"},
|
||||
}
|
||||
|
||||
// StringCode returns a string code of the activity
|
||||
|
||||
@@ -165,16 +165,16 @@ func (store *Store) Get(ctx context.Context, accountID string, offset, limit int
|
||||
return store.processResult(ctx, events)
|
||||
}
|
||||
|
||||
// Save an event in the SQLite events table end encrypt the "email" element in meta map
|
||||
func (store *Store) Save(_ context.Context, event *activity.Event) (*activity.Event, error) {
|
||||
// Save persists an activity event and encrypts deleted user details using the caller's context.
|
||||
func (store *Store) Save(ctx context.Context, event *activity.Event) (*activity.Event, error) {
|
||||
eventCopy := event.Copy()
|
||||
meta, err := store.saveDeletedUserEmailAndNameInEncrypted(eventCopy)
|
||||
meta, err := store.saveDeletedUserEmailAndNameInEncrypted(ctx, eventCopy)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
eventCopy.Meta = meta
|
||||
|
||||
if err = store.db.Create(eventCopy).Error; err != nil {
|
||||
if err = store.db.WithContext(ctx).Create(eventCopy).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -183,7 +183,7 @@ func (store *Store) Save(_ context.Context, event *activity.Event) (*activity.Ev
|
||||
|
||||
// saveDeletedUserEmailAndNameInEncrypted if the meta contains email and name then store it in encrypted way and delete
|
||||
// this item from meta map
|
||||
func (store *Store) saveDeletedUserEmailAndNameInEncrypted(event *activity.Event) (map[string]any, error) {
|
||||
func (store *Store) saveDeletedUserEmailAndNameInEncrypted(ctx context.Context, event *activity.Event) (map[string]any, error) {
|
||||
email, ok := event.Meta["email"]
|
||||
if !ok {
|
||||
return event.Meta, nil
|
||||
@@ -211,7 +211,7 @@ func (store *Store) saveDeletedUserEmailAndNameInEncrypted(event *activity.Event
|
||||
}
|
||||
deletedUser.Name = encryptedName
|
||||
|
||||
err = store.db.Clauses(clause.OnConflict{
|
||||
err = store.db.WithContext(ctx).Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{"email", "name"}),
|
||||
}).Create(deletedUser).Error
|
||||
|
||||
@@ -7,11 +7,49 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
"github.com/netbirdio/netbird/util/crypt"
|
||||
)
|
||||
|
||||
func TestSave_CancellationWhileWaitingForConnection(t *testing.T) {
|
||||
t.Setenv(storeEngineEnv, "sqlite")
|
||||
key, err := crypt.GenerateKey()
|
||||
require.NoError(t, err)
|
||||
store, err := NewSqlStore(context.Background(), t.TempDir(), key)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { assert.NoError(t, store.Close(context.Background())) })
|
||||
db, err := store.db.DB()
|
||||
require.NoError(t, err)
|
||||
conn, err := db.Conn(context.Background())
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
|
||||
defer cancel()
|
||||
result := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := store.Save(ctx, &activity.Event{
|
||||
Timestamp: time.Now().UTC(), Activity: activity.CustomDomainValidationExpired,
|
||||
AccountID: "account-id", TargetID: "domain-id", InitiatorID: activity.SystemInitiator,
|
||||
})
|
||||
result <- err
|
||||
}()
|
||||
select {
|
||||
case err := <-result:
|
||||
assert.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
require.NoError(t, conn.Close())
|
||||
case <-time.After(time.Second):
|
||||
// Release the connection so a regression cannot leave the writer running.
|
||||
require.NoError(t, conn.Close())
|
||||
assert.ErrorIs(t, <-result, context.DeadlineExceeded)
|
||||
t.Error("activity writes must stop waiting when their deadline expires")
|
||||
}
|
||||
events, err := store.Get(context.Background(), "account-id", 0, 10, true)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, events, "a timed-out write must not persist after the connection is released")
|
||||
}
|
||||
|
||||
func TestNewSqlStore(t *testing.T) {
|
||||
dataDir := t.TempDir()
|
||||
key, _ := crypt.GenerateKey()
|
||||
|
||||
@@ -89,6 +89,7 @@ func TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection(t *t
|
||||
|
||||
// Bootstrap is an explicit settings create; providers have no settings
|
||||
// side effects anymore.
|
||||
seedPrivateProxyCluster(t, am.Store, clusterAddr)
|
||||
before, err := mgr.CreateSettings(ctx, adminUserID, agenttypes.DefaultSettings(accountID), clusterAddr, "")
|
||||
require.NoError(t, err, "CreateSettings must bootstrap the row")
|
||||
require.Equal(t, clusterAddr, before.ProxyAddress, "proxy address pinned at bootstrap")
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
networkmap "github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
agenttypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
rpproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
@@ -92,6 +93,7 @@ func TestAgentNetwork_ProviderCRUD_FansOutToProxyAndClientPeers(t *testing.T) {
|
||||
// UpdateAccountPeers, which is the path under test.
|
||||
agentMgr := agentnetwork.NewManager(am.Store, permissions.NewManager(am.Store), am, nil)
|
||||
|
||||
seedPrivateProxyCluster(t, am.Store, clusterAddr)
|
||||
_, err = agentMgr.CreateSettings(ctx, adminUserID, agenttypes.DefaultSettings(accountID), clusterAddr, "")
|
||||
require.NoError(t, err, "CreateSettings must bootstrap the endpoint")
|
||||
// The bootstrap itself reconciles and queues updates on both channels;
|
||||
@@ -222,3 +224,22 @@ func synthZoneRData(sync *nbproto.SyncResponse, clusterAddr, fqdn string) string
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// seedPrivateProxyCluster registers a connected proxy with private capabilities in a
|
||||
// netbird client for clusterAddr, matching what a real deployment looks like
|
||||
// when the account bootstraps: the agent-network gateway service is always
|
||||
// private, so its cluster has to be one that can serve private services.
|
||||
func seedPrivateProxyCluster(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: "agent-net-proxy-" + clusterAddr,
|
||||
SessionID: "agent-net-session",
|
||||
ClusterAddress: clusterAddr,
|
||||
LastSeen: now,
|
||||
ConnectedAt: &now,
|
||||
Status: rpproxy.StatusConnected,
|
||||
Capabilities: rpproxy.Capabilities{Private: &private},
|
||||
}), "seeding the proxy cluster must succeed")
|
||||
}
|
||||
|
||||
+26
-15
@@ -61,23 +61,34 @@ func (am *DefaultAccountManager) GetEvents(ctx context.Context, accountID, userI
|
||||
return filtered, nil
|
||||
}
|
||||
|
||||
// StoreEvent records an activity, waiting for expiration events before cleanup can stop.
|
||||
func (am *DefaultAccountManager) StoreEvent(ctx context.Context, initiatorID, targetID, accountID string, activityID activity.ActivityDescriber, meta map[string]any) {
|
||||
if isEnabled() {
|
||||
go func() {
|
||||
_, err := am.eventStore.Save(ctx, &activity.Event{
|
||||
Timestamp: time.Now().UTC(),
|
||||
Activity: activityID.(activity.Activity),
|
||||
InitiatorID: initiatorID,
|
||||
TargetID: targetID,
|
||||
AccountID: accountID,
|
||||
Meta: meta,
|
||||
})
|
||||
if err != nil {
|
||||
// todo add metric
|
||||
log.WithContext(ctx).Errorf("received an error while storing an activity event, error: %s", err)
|
||||
}
|
||||
}()
|
||||
if !isEnabled() {
|
||||
return
|
||||
}
|
||||
eventStore := am.eventStore
|
||||
save := func(ctx context.Context) {
|
||||
_, err := eventStore.Save(ctx, &activity.Event{
|
||||
Timestamp: time.Now().UTC(),
|
||||
Activity: activityID.(activity.Activity),
|
||||
InitiatorID: initiatorID,
|
||||
TargetID: targetID,
|
||||
AccountID: accountID,
|
||||
Meta: meta,
|
||||
})
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("received an error while storing an activity event, error: %s", err)
|
||||
}
|
||||
}
|
||||
if activityID == activity.CustomDomainValidationExpired {
|
||||
// The domain is already deleted; shutdown must allow its audit write to finish.
|
||||
ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second)
|
||||
defer cancel()
|
||||
save(ctx)
|
||||
return
|
||||
}
|
||||
// Request cancellation must not discard the audit record of a completed operation.
|
||||
go save(context.WithoutCancel(ctx))
|
||||
}
|
||||
|
||||
type eventUserInfo struct {
|
||||
|
||||
@@ -6,10 +6,52 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
|
||||
"github.com/netbirdio/netbird/util/crypt"
|
||||
)
|
||||
|
||||
func TestStoreEvent_CanceledContext(t *testing.T) {
|
||||
t.Setenv("NB_EVENT_ACTIVITY_LOG_ENABLED", "true")
|
||||
t.Setenv("NB_ACTIVITY_EVENT_STORE_ENGINE", "sqlite")
|
||||
for _, code := range []activity.Activity{activity.CustomDomainValidationExpired, activity.DomainAdded} {
|
||||
t.Run(code.StringCode(), func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
key, err := crypt.GenerateKey()
|
||||
require.NoError(t, err)
|
||||
eventStore, err := activitystore.NewSqlStore(context.Background(), dir, key)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { assert.NoError(t, eventStore.Close(context.Background())) })
|
||||
manager := &DefaultAccountManager{eventStore: eventStore}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
// The operation already succeeded when shutdown or the request cancels its context.
|
||||
manager.StoreEvent(ctx, activity.SystemInitiator, "domain-id", "account-id",
|
||||
code, map[string]any{"domain": "expired.example.com"})
|
||||
if code != activity.CustomDomainValidationExpired {
|
||||
require.Eventually(t, func() bool {
|
||||
events, err := eventStore.Get(context.Background(), "account-id", 0, 10, true)
|
||||
return err == nil && len(events) == 1
|
||||
}, time.Second, time.Millisecond, "asynchronous events must survive request cancellation")
|
||||
}
|
||||
require.NoError(t, eventStore.Close(context.Background()))
|
||||
|
||||
reopened, err := activitystore.NewSqlStore(context.Background(), dir, key)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { assert.NoError(t, reopened.Close(context.Background())) })
|
||||
events, err := reopened.Get(context.Background(), "account-id", 0, 10, true)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, events, 1, "the event must be persisted before shutdown closes the store")
|
||||
assert.Equal(t, code, events[0].Activity, "persist the requested activity")
|
||||
assert.Equal(t, "domain-id", events[0].TargetID, "retain the registration ID")
|
||||
assert.Equal(t, "expired.example.com", events[0].Meta["domain"], "retain the domain name")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func generateAndStoreEvents(t *testing.T, manager *DefaultAccountManager, typ activity.Activity, initiatorID, targetID,
|
||||
accountID string, count int) {
|
||||
t.Helper()
|
||||
|
||||
@@ -101,10 +101,8 @@ func (am *DefaultAccountManager) CreateGroup(ctx context.Context, accountID, use
|
||||
return status.Errorf(status.Internal, "failed to create group: %v", err)
|
||||
}
|
||||
|
||||
for _, peerID := range newGroup.Peers {
|
||||
if err := transaction.AddPeerToGroup(ctx, accountID, peerID, newGroup.ID); err != nil {
|
||||
return status.Errorf(status.Internal, "failed to add peer %s to group %s: %v", peerID, newGroup.ID, err)
|
||||
}
|
||||
if err = syncGroupMembership(ctx, transaction, accountID, newGroup.ID, newGroup.Peers, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
snap, err = affectedpeers.Load(ctx, transaction, accountID, change)
|
||||
@@ -200,6 +198,9 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use
|
||||
|
||||
// syncGroupMembership applies the peer membership delta for a group within a transaction.
|
||||
func syncGroupMembership(ctx context.Context, transaction store.Store, accountID, groupID string, peersToAdd, peersToRemove []string) error {
|
||||
if err := validateGroupPeers(ctx, transaction, accountID, peersToAdd); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, peerID := range peersToAdd {
|
||||
if err := transaction.AddPeerToGroup(ctx, accountID, peerID, groupID); err != nil {
|
||||
return status.Errorf(status.Internal, "failed to add peer %s to group %s: %v", peerID, groupID, err)
|
||||
@@ -213,6 +214,25 @@ func syncGroupMembership(ctx context.Context, transaction store.Store, accountID
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateGroupPeers(ctx context.Context, transaction store.Store, accountID string, peerIDs []string) error {
|
||||
if len(peerIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
peers, err := transaction.GetPeersByIDs(ctx, store.LockingStrengthNone, accountID, peerIDs)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, peerID := range peerIDs {
|
||||
if _, ok := peers[peerID]; !ok {
|
||||
return status.Errorf(status.InvalidArgument, "peer with ID %s not found", peerID)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// CreateGroups adds new groups to the account.
|
||||
// Note: This function does not acquire the global lock.
|
||||
// It is the caller's responsibility to ensure proper locking is in place before invoking this method.
|
||||
@@ -540,7 +560,7 @@ func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, gr
|
||||
change := affectedpeers.Change{OutputPeerIDs: []string{peerID}, LinkGroups: []string{groupID}}
|
||||
|
||||
err := am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
if err := transaction.AddPeerToGroup(ctx, accountID, peerID, groupID); err != nil {
|
||||
if err := syncGroupMembership(ctx, transaction, accountID, groupID, []string{peerID}, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -754,6 +774,14 @@ func validateDeleteGroup(ctx context.Context, transaction store.Store, group *ty
|
||||
return &GroupLinkError{"agent network policy", linkedPolicy.Name}
|
||||
}
|
||||
|
||||
isLinked, linkedRule, err := isGroupLinkedToAgentNetworkBudgetRule(ctx, transaction, group.AccountID, group.ID)
|
||||
if err != nil {
|
||||
return status.Errorf(status.Internal, "failed to check agent network budget rules")
|
||||
}
|
||||
if isLinked {
|
||||
return &GroupLinkError{"agent network budget rule", linkedRule.Name}
|
||||
}
|
||||
|
||||
return checkGroupLinkedToSettings(ctx, transaction, group)
|
||||
}
|
||||
|
||||
@@ -925,6 +953,26 @@ func isGroupLinkedToAgentNetworkPolicy(ctx context.Context, transaction store.St
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// isGroupLinkedToAgentNetworkBudgetRule checks if a group is a target of any
|
||||
// account-level agent network budget rule.
|
||||
func isGroupLinkedToAgentNetworkBudgetRule(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *agentNetworkTypes.AccountBudgetRule, error) {
|
||||
rules, err := transaction.GetAccountAgentNetworkBudgetRules(ctx, store.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("error retrieving agent network budget rules while checking group linkage: %v", err)
|
||||
return false, nil, err
|
||||
}
|
||||
|
||||
for _, rule := range rules {
|
||||
if rule == nil {
|
||||
continue
|
||||
}
|
||||
if slices.Contains(rule.TargetGroups, groupID) {
|
||||
return true, rule, nil
|
||||
}
|
||||
}
|
||||
return false, nil, nil
|
||||
}
|
||||
|
||||
// areGroupChangesAffectPeers checks if any changes to the specified groups will affect peers.
|
||||
// It fetches each collection once and checks all groupIDs against them in memory.
|
||||
func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) {
|
||||
|
||||
@@ -11,10 +11,10 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.org/x/exp/maps"
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
@@ -132,6 +132,11 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) {
|
||||
"grp-for-agent-network-policy",
|
||||
"agent network policy",
|
||||
},
|
||||
{
|
||||
"agent network budget rule",
|
||||
"grp-for-agent-network-budget-rule",
|
||||
"agent network budget rule",
|
||||
},
|
||||
{
|
||||
"reverse proxy private service access group",
|
||||
"grp-for-rp-private",
|
||||
@@ -152,6 +157,16 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) {
|
||||
return
|
||||
}
|
||||
|
||||
group, getErr := am.GetGroup(context.Background(), account.Id, testCase.groupID, groupAdminUserID)
|
||||
if getErr != nil {
|
||||
t.Errorf("group %s should still exist after failed deletion: %s", testCase.groupID, getErr)
|
||||
return
|
||||
}
|
||||
if group == nil {
|
||||
t.Errorf("group %s was deleted despite the failed deletion", testCase.groupID)
|
||||
return
|
||||
}
|
||||
|
||||
var sErr *status.Error
|
||||
if errors.As(err, &sErr) {
|
||||
if sErr.Message != testCase.expectedReason {
|
||||
@@ -240,6 +255,12 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) {
|
||||
groupIDs: []string{"grp-for-agent-network-policy"},
|
||||
expectedReasons: []string{"agent network policy"},
|
||||
},
|
||||
{
|
||||
name: "agent network budget rule",
|
||||
groupIDs: []string{"grp-for-agent-network-budget-rule"},
|
||||
expectedReasons: []string{"agent network budget rule"},
|
||||
expectedNotDeleted: []string{"grp-for-agent-network-budget-rule"},
|
||||
},
|
||||
{
|
||||
name: "reverse proxy services",
|
||||
groupIDs: []string{"grp-for-rp-private", "grp-for-rp-bearer"},
|
||||
@@ -501,6 +522,14 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t
|
||||
Peers: make([]string, 0),
|
||||
}
|
||||
|
||||
groupForAgentNetworkBudgetRule := &types.Group{
|
||||
ID: "grp-for-agent-network-budget-rule",
|
||||
AccountID: "account-id",
|
||||
Name: "Group for agent network budget rules",
|
||||
Issued: types.GroupIssuedAPI,
|
||||
Peers: make([]string, 0),
|
||||
}
|
||||
|
||||
groupForRPPrivate := &types.Group{
|
||||
ID: "grp-for-rp-private",
|
||||
AccountID: "account-id",
|
||||
@@ -573,6 +602,7 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForUsers)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForIntegration)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForAgentNetworkPolicy)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForAgentNetworkBudgetRule)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPPrivate)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPBearer)
|
||||
|
||||
@@ -587,6 +617,20 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
budgetRuleDecoy := agentNetworkTypes.NewAccountBudgetRule(accountID)
|
||||
budgetRuleDecoy.Name = "Unrelated agent network budget rule"
|
||||
budgetRuleDecoy.TargetGroups = []string{"unrelated-group"}
|
||||
if err := am.Store.SaveAgentNetworkBudgetRule(context.Background(), budgetRuleDecoy); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
budgetRule := agentNetworkTypes.NewAccountBudgetRule(accountID)
|
||||
budgetRule.Name = "Example agent network budget rule"
|
||||
budgetRule.TargetGroups = []string{groupForAgentNetworkBudgetRule.ID}
|
||||
if err := am.Store.SaveAgentNetworkBudgetRule(context.Background(), budgetRule); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
// The decoy services are created first so the linkage check has to scan
|
||||
// past services that do not reference the groups under test.
|
||||
rpServices := []*rpservice.Service{
|
||||
@@ -1236,3 +1280,82 @@ func Test_IncrementNetworkSerial(t *testing.T) {
|
||||
|
||||
assert.Equal(t, totalPeers, int(account.Network.Serial), "Expected %d serial increases in account %s, got %d", totalPeers, accountID, account.Network.Serial)
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_GroupPeersMustBelongToAccount(t *testing.T) {
|
||||
manager, _, account, peer1, _, _ := setupNetworkMapTest(t)
|
||||
|
||||
otherAccount, err := createAccount(manager, "other_account", "other_user", "")
|
||||
require.NoError(t, err)
|
||||
|
||||
foreignPeer := &peer2.Peer{
|
||||
ID: "foreign-peer",
|
||||
AccountID: otherAccount.Id,
|
||||
Key: "foreign-key",
|
||||
DNSLabel: "foreign-peer",
|
||||
IP: uint32ToIP(1),
|
||||
}
|
||||
require.NoError(t, manager.Store.AddPeerToAccount(context.Background(), foreignPeer))
|
||||
|
||||
assertRejected := func(t *testing.T, err error) {
|
||||
t.Helper()
|
||||
require.Error(t, err)
|
||||
s, ok := status.FromError(err)
|
||||
require.True(t, ok, "expected status error, got %v", err)
|
||||
assert.Equal(t, status.InvalidArgument, s.Type(), "peer outside the account should be rejected as invalid argument")
|
||||
}
|
||||
|
||||
t.Run("create rejects foreign peer", func(t *testing.T) {
|
||||
err := manager.CreateGroup(context.Background(), account.Id, userID, &types.Group{
|
||||
Name: "foreign",
|
||||
Issued: types.GroupIssuedAPI,
|
||||
Peers: []string{peer1.ID, foreignPeer.ID},
|
||||
})
|
||||
assertRejected(t, err)
|
||||
|
||||
_, err = manager.Store.GetGroupByName(context.Background(), store.LockingStrengthNone, account.Id, "foreign")
|
||||
assert.Error(t, err, "rejected create must not persist the group")
|
||||
})
|
||||
|
||||
t.Run("update rejects foreign and unknown peers", func(t *testing.T) {
|
||||
group := &types.Group{ID: "own", Name: "own", Issued: types.GroupIssuedAPI, Peers: []string{peer1.ID}}
|
||||
require.NoError(t, manager.CreateGroup(context.Background(), account.Id, userID, group))
|
||||
|
||||
group.Peers = []string{peer1.ID, foreignPeer.ID}
|
||||
assertRejected(t, manager.UpdateGroup(context.Background(), account.Id, userID, group))
|
||||
|
||||
group.Peers = []string{peer1.ID, "does-not-exist"}
|
||||
assertRejected(t, manager.UpdateGroup(context.Background(), account.Id, userID, group))
|
||||
|
||||
stored, err := manager.Store.GetGroupByID(context.Background(), store.LockingStrengthNone, account.Id, group.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{peer1.ID}, stored.Peers, "rejected updates must not change membership")
|
||||
})
|
||||
|
||||
t.Run("update tolerates and drops pre-existing dangling members", func(t *testing.T) {
|
||||
group := &types.Group{ID: "polluted", Name: "polluted", Issued: types.GroupIssuedAPI, Peers: []string{peer1.ID}}
|
||||
require.NoError(t, manager.CreateGroup(context.Background(), account.Id, userID, group))
|
||||
require.NoError(t, manager.Store.AddPeerToGroup(context.Background(), account.Id, foreignPeer.ID, group.ID))
|
||||
|
||||
group.Peers = []string{peer1.ID, foreignPeer.ID}
|
||||
assert.NoError(t, manager.UpdateGroup(context.Background(), account.Id, userID, group), "keeping an existing member must not be rejected")
|
||||
|
||||
group.Peers = []string{peer1.ID}
|
||||
require.NoError(t, manager.UpdateGroup(context.Background(), account.Id, userID, group))
|
||||
|
||||
stored, err := manager.Store.GetGroupByID(context.Background(), store.LockingStrengthNone, account.Id, group.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{peer1.ID}, stored.Peers, "dangling member should be removed once omitted")
|
||||
})
|
||||
|
||||
t.Run("direct add rejects foreign and unknown peers", func(t *testing.T) {
|
||||
group := &types.Group{ID: "direct", Name: "direct", Issued: types.GroupIssuedAPI, Peers: []string{peer1.ID}}
|
||||
require.NoError(t, manager.CreateGroup(context.Background(), account.Id, userID, group))
|
||||
|
||||
assertRejected(t, manager.GroupAddPeer(context.Background(), account.Id, group.ID, foreignPeer.ID))
|
||||
assertRejected(t, manager.GroupAddPeer(context.Background(), account.Id, group.ID, "does-not-exist"))
|
||||
|
||||
stored, err := manager.Store.GetGroupByID(context.Background(), store.LockingStrengthNone, account.Id, group.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{peer1.ID}, stored.Peers, "rejected direct adds must not change membership")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -148,13 +148,10 @@ func (h *handler) updateGroup(w http.ResponseWriter, r *http.Request) {
|
||||
peers = *req.Peers
|
||||
}
|
||||
|
||||
resources := make([]types.Resource, 0)
|
||||
if req.Resources != nil {
|
||||
for _, res := range *req.Resources {
|
||||
resource := types.Resource{}
|
||||
resource.FromAPIRequest(&res)
|
||||
resources = append(resources, resource)
|
||||
}
|
||||
resources, err := resourcesFromAPIRequest(req.Resources)
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
|
||||
group := types.Group{
|
||||
@@ -210,13 +207,10 @@ func (h *handler) createGroup(w http.ResponseWriter, r *http.Request) {
|
||||
peers = *req.Peers
|
||||
}
|
||||
|
||||
resources := make([]types.Resource, 0)
|
||||
if req.Resources != nil {
|
||||
for _, res := range *req.Resources {
|
||||
resource := types.Resource{}
|
||||
resource.FromAPIRequest(&res)
|
||||
resources = append(resources, resource)
|
||||
}
|
||||
resources, err := resourcesFromAPIRequest(req.Resources)
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
|
||||
group := types.Group{
|
||||
@@ -335,11 +329,30 @@ func toGroupResponse(peers []*nbpeer.Peer, group *types.Group) *api.Group {
|
||||
gr.PeersCount = len(gr.Peers)
|
||||
|
||||
for _, res := range group.Resources {
|
||||
resResp := res.ToAPIResponse()
|
||||
gr.Resources = append(gr.Resources, *resResp)
|
||||
if resResp := res.ToAPIResponse(); resResp != nil {
|
||||
gr.Resources = append(gr.Resources, *resResp)
|
||||
}
|
||||
}
|
||||
|
||||
gr.ResourcesCount = len(gr.Resources)
|
||||
|
||||
return &gr
|
||||
}
|
||||
|
||||
func resourcesFromAPIRequest(req *[]api.Resource) ([]types.Resource, error) {
|
||||
resources := make([]types.Resource, 0)
|
||||
if req == nil {
|
||||
return resources, nil
|
||||
}
|
||||
|
||||
for _, res := range *req {
|
||||
if res.Id == "" || !types.ResourceType(res.Type).Valid() {
|
||||
return nil, status.Errorf(status.InvalidArgument, "resource id shouldn't be empty and type must be one of: peer, domain, host, subnet")
|
||||
}
|
||||
resource := types.Resource{}
|
||||
resource.FromAPIRequest(&res)
|
||||
resources = append(resources, resource)
|
||||
}
|
||||
|
||||
return resources, nil
|
||||
}
|
||||
|
||||
@@ -8,8 +8,8 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -208,6 +208,33 @@ func TestWriteGroup(t *testing.T) {
|
||||
expectedStatus: http.StatusUnprocessableEntity,
|
||||
expectedBody: false,
|
||||
},
|
||||
{
|
||||
name: "Write Group POST Empty Resource",
|
||||
requestType: http.MethodPost,
|
||||
requestPath: "/api/groups",
|
||||
requestBody: bytes.NewBuffer(
|
||||
[]byte(`{"name":"With Resource","resources":[{}]}`)),
|
||||
expectedStatus: http.StatusUnprocessableEntity,
|
||||
expectedBody: false,
|
||||
},
|
||||
{
|
||||
name: "Write Group PUT Empty Resource",
|
||||
requestType: http.MethodPut,
|
||||
requestPath: "/api/groups/id-existed",
|
||||
requestBody: bytes.NewBuffer(
|
||||
[]byte(`{"name":"With Resource","resources":[{"id":"","type":"host"}]}`)),
|
||||
expectedStatus: http.StatusUnprocessableEntity,
|
||||
expectedBody: false,
|
||||
},
|
||||
{
|
||||
name: "Write Group POST Unknown Resource Type",
|
||||
requestType: http.MethodPost,
|
||||
requestPath: "/api/groups",
|
||||
requestBody: bytes.NewBuffer(
|
||||
[]byte(`{"name":"With Resource","resources":[{"id":"res-1","type":"banana"}]}`)),
|
||||
expectedStatus: http.StatusUnprocessableEntity,
|
||||
expectedBody: false,
|
||||
},
|
||||
{
|
||||
name: "Write Group PUT OK",
|
||||
requestType: http.MethodPut,
|
||||
@@ -376,6 +403,20 @@ func TestGetAllGroups(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestToGroupResponseSkipsEmptyResource(t *testing.T) {
|
||||
group := &types.Group{
|
||||
ID: "id-resources",
|
||||
Name: "Resources",
|
||||
Issued: types.GroupIssuedAPI,
|
||||
Resources: []types.Resource{{}, {ID: "res-1", Type: types.ResourceTypeHost}},
|
||||
}
|
||||
|
||||
got := toGroupResponse(nil, group)
|
||||
|
||||
assert.Equal(t, 1, got.ResourcesCount)
|
||||
assert.Equal(t, []api.Resource{{Id: "res-1", Type: api.ResourceType(types.ResourceTypeHost)}}, got.Resources)
|
||||
}
|
||||
|
||||
func TestDeleteGroup(t *testing.T) {
|
||||
tt := []struct {
|
||||
name string
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -10,6 +11,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/http/testing/testing_tools"
|
||||
"github.com/netbirdio/netbird/management/server/http/testing/testing_tools/channel"
|
||||
@@ -34,7 +36,7 @@ func Test_Events_GetAll(t *testing.T) {
|
||||
|
||||
for _, user := range users {
|
||||
t.Run(user.name+" - Get all events", func(t *testing.T) {
|
||||
apiHandler, _, _ := channel.BuildApiBlackBoxWithDBState(t, "../testdata/events.sql", nil, false)
|
||||
apiHandler, accountManager, _ := channel.BuildApiBlackBoxWithDBState(t, "../testdata/events.sql", nil, false)
|
||||
|
||||
// First, perform a mutation to generate an event (create a group as admin)
|
||||
groupBody, err := json.Marshal(&api.GroupRequest{Name: "eventTestGroup"})
|
||||
@@ -44,7 +46,14 @@ func Test_Events_GetAll(t *testing.T) {
|
||||
createReq := testing_tools.BuildRequest(t, groupBody, http.MethodPost, "/api/groups", testing_tools.TestAdminId)
|
||||
createRecorder := httptest.NewRecorder()
|
||||
apiHandler.ServeHTTP(createRecorder, createReq)
|
||||
assert.Equal(t, http.StatusOK, createRecorder.Code, "Failed to create group to generate event")
|
||||
require.Equal(t, http.StatusOK, createRecorder.Code, "Failed to create group to generate event")
|
||||
|
||||
// Group creation returns before its asynchronous audit write finishes.
|
||||
require.EventuallyWithT(t, func(c *assert.CollectT) {
|
||||
events, err := accountManager.GetEvents(context.Background(), testing_tools.TestAccountId, testing_tools.TestAdminId)
|
||||
assert.NoError(c, err)
|
||||
assert.NotEmpty(c, events, "wait for the group creation event before checking permissions")
|
||||
}, time.Second, 10*time.Millisecond)
|
||||
|
||||
// Now query events
|
||||
req := testing_tools.BuildRequest(t, []byte{}, http.MethodGet, "/api/events", user.userId)
|
||||
|
||||
@@ -23,6 +23,9 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
// maxDiscoveryDocumentSize caps the discovery document read at 1 MiB. Providers serve a few kilobytes.
|
||||
const maxDiscoveryDocumentSize = 1 << 20
|
||||
|
||||
// oidcProviderJSON represents the OpenID Connect discovery document
|
||||
type oidcProviderJSON struct {
|
||||
Issuer string `json:"issuer"`
|
||||
@@ -35,6 +38,10 @@ func validateOIDCIssuer(ctx context.Context, issuer string) error {
|
||||
|
||||
httpClient := &http.Client{
|
||||
Timeout: 10 * time.Second,
|
||||
// An issuer that redirects its own discovery document is misconfigured.
|
||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, wellKnown, nil)
|
||||
@@ -48,22 +55,22 @@ func validateOIDCIssuer(ctx context.Context, issuer string) error {
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: unable to read response body: %v", types.ErrIdentityProviderIssuerUnreachable, err)
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("%w: %s", types.ErrIdentityProviderIssuerUnreachable, resp.Status)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("%w: %s: %s", types.ErrIdentityProviderIssuerUnreachable, resp.Status, body)
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxDiscoveryDocumentSize+1))
|
||||
if err != nil || len(body) > maxDiscoveryDocumentSize {
|
||||
return fmt.Errorf("%w: failed to decode provider discovery object", types.ErrIdentityProviderIssuerUnreachable)
|
||||
}
|
||||
|
||||
var p oidcProviderJSON
|
||||
if err := json.Unmarshal(body, &p); err != nil {
|
||||
return fmt.Errorf("%w: failed to decode provider discovery object: %v", types.ErrIdentityProviderIssuerUnreachable, err)
|
||||
return fmt.Errorf("%w: failed to decode provider discovery object", types.ErrIdentityProviderIssuerUnreachable)
|
||||
}
|
||||
|
||||
if p.Issuer != issuer {
|
||||
return fmt.Errorf("%w: expected %q got %q", types.ErrIdentityProviderIssuerMismatch, issuer, p.Issuer)
|
||||
return fmt.Errorf("%w: %q", types.ErrIdentityProviderIssuerMismatch, issuer)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -151,15 +158,15 @@ func (am *DefaultAccountManager) CreateIdentityProvider(ctx context.Context, acc
|
||||
return nil, status.NewPermissionDeniedError()
|
||||
}
|
||||
|
||||
if err := validateIdentityProviderConfig(ctx, idpConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
embeddedManager, ok := am.idpManager.(*idp.EmbeddedIdPManager)
|
||||
if !ok {
|
||||
return nil, status.Errorf(status.Internal, "identity provider management requires embedded IdP")
|
||||
}
|
||||
|
||||
if err := validateIdentityProviderConfig(ctx, idpConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Generate ID if not provided
|
||||
if idpConfig.ID == "" {
|
||||
idpConfig.ID = generateIdentityProviderID(idpConfig.Type)
|
||||
@@ -188,15 +195,15 @@ func (am *DefaultAccountManager) UpdateIdentityProvider(ctx context.Context, acc
|
||||
return nil, status.NewPermissionDeniedError()
|
||||
}
|
||||
|
||||
if err := validateIdentityProviderConfig(ctx, idpConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
embeddedManager, ok := am.idpManager.(*idp.EmbeddedIdPManager)
|
||||
if !ok {
|
||||
return nil, status.Errorf(status.Internal, "identity provider management requires embedded IdP")
|
||||
}
|
||||
|
||||
if err := validateIdentityProviderConfig(ctx, idpConfig); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
idpConfig.ID = idpID
|
||||
idpConfig.AccountID = accountID
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -121,7 +122,7 @@ func createManagerWithEmbeddedIdPModeAndSetup(
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_CreateIdentityProvider_Validation(t *testing.T) {
|
||||
manager, _, err := createManager(t)
|
||||
manager, _, err := createManagerWithEmbeddedIdP(t)
|
||||
require.NoError(t, err)
|
||||
|
||||
userID := "testingUser"
|
||||
@@ -233,7 +234,7 @@ func TestUpdateUserAuthWithSingleModeKeepsConfiguredDomain(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_UpdateIdentityProvider_Validation(t *testing.T) {
|
||||
manager, _, err := createManager(t)
|
||||
manager, _, err := createManagerWithEmbeddedIdP(t)
|
||||
require.NoError(t, err)
|
||||
|
||||
userID := "testingUser"
|
||||
@@ -355,3 +356,45 @@ func TestValidateOIDCIssuer_TrailingSlash(t *testing.T) {
|
||||
require.Error(t, err)
|
||||
assert.True(t, errors.Is(err, types.ErrIdentityProviderIssuerMismatch))
|
||||
}
|
||||
|
||||
func TestValidateOIDCIssuer_DoesNotFollowRedirects(t *testing.T) {
|
||||
var reached bool
|
||||
target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
reached = true
|
||||
w.WriteHeader(http.StatusForbidden)
|
||||
}))
|
||||
t.Cleanup(target.Close)
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, target.URL+"/redirect-target", http.StatusFound)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
err := validateOIDCIssuer(context.Background(), srv.URL)
|
||||
require.Error(t, err)
|
||||
assert.False(t, reached, "Redirects are not followed")
|
||||
}
|
||||
|
||||
func TestValidateOIDCIssuer_BoundsResponseSize(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"issuer":"` + strings.Repeat("a", maxDiscoveryDocumentSize) + `"}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
err := validateOIDCIssuer(context.Background(), srv.URL)
|
||||
require.ErrorIs(t, err, types.ErrIdentityProviderIssuerUnreachable)
|
||||
assert.NotErrorIs(t, err, types.ErrIdentityProviderIssuerMismatch)
|
||||
}
|
||||
|
||||
func TestValidateOIDCIssuer_RejectsTrailingContent(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"issuer":"http://` + r.Host + `"} {"issuer":"second"}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
err := validateOIDCIssuer(context.Background(), srv.URL)
|
||||
require.ErrorIs(t, err, types.ErrIdentityProviderIssuerUnreachable,
|
||||
"Content after the first object is not a valid discovery document")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
package migration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
)
|
||||
|
||||
// MigrateCustomDomainValidationExpiry gives existing pending registrations a validation window.
|
||||
func MigrateCustomDomainValidationExpiry(ctx context.Context, db *gorm.DB) error {
|
||||
result := db.WithContext(ctx).Model(&domain.Domain{}).
|
||||
Where("validated = ? AND validation_expires_at IS NULL", false).
|
||||
Update("validation_expires_at", time.Now().UTC().Add(domain.ValidationTTL))
|
||||
if result.Error != nil {
|
||||
return fmt.Errorf("backfill custom domain validation expiry: %w", result.Error)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package migration_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
"github.com/netbirdio/netbird/management/server/migration"
|
||||
)
|
||||
|
||||
func TestMigrateCustomDomainValidationExpiry(t *testing.T) {
|
||||
db := setupDatabase(t)
|
||||
require.NoError(t, db.AutoMigrate(&domain.Domain{}))
|
||||
t.Cleanup(func() { require.NoError(t, db.Migrator().DropTable(&domain.Domain{})) })
|
||||
ctx := context.Background()
|
||||
existingDeadline := time.Now().UTC().Add(time.Hour).Truncate(time.Second)
|
||||
rows := []domain.Domain{
|
||||
{ID: "legacy", Domain: "legacy.example.com"},
|
||||
{ID: "validated", Domain: "validated.example.com", Validated: true},
|
||||
{ID: "pending", Domain: "pending.example.com", ValidationExpiresAt: &existingDeadline},
|
||||
}
|
||||
require.NoError(t, db.Create(&rows).Error)
|
||||
before := time.Now().UTC()
|
||||
require.NoError(t, migration.MigrateCustomDomainValidationExpiry(ctx, db))
|
||||
after := time.Now().UTC()
|
||||
var migrated domain.Domain
|
||||
require.NoError(t, db.First(&migrated, "id = ?", "legacy").Error)
|
||||
require.NotNil(t, migrated.ValidationExpiresAt)
|
||||
assert.WithinRange(t, *migrated.ValidationExpiresAt, before.Truncate(time.Millisecond).Add(48*time.Hour), after.Add(48*time.Hour+time.Millisecond), "legacy pending registrations get a full window")
|
||||
deadline := *migrated.ValidationExpiresAt
|
||||
require.NoError(t, migration.MigrateCustomDomainValidationExpiry(ctx, db))
|
||||
require.NoError(t, db.First(&migrated, "id = ?", "legacy").Error)
|
||||
assert.Equal(t, deadline, *migrated.ValidationExpiresAt, "repeated migration must not extend the deadline")
|
||||
var validated, pending domain.Domain
|
||||
require.NoError(t, db.First(&validated, "id = ?", "validated").Error)
|
||||
require.NoError(t, db.First(&pending, "id = ?", "pending").Error)
|
||||
assert.Nil(t, validated.ValidationExpiresAt, "validated domains do not acquire an expiry")
|
||||
require.NotNil(t, pending.ValidationExpiresAt)
|
||||
assert.WithinDuration(t, existingDeadline, *pending.ValidationExpiresAt, 0, "existing deadlines must be preserved")
|
||||
}
|
||||
@@ -1494,9 +1494,12 @@ func checkAuth(ctx context.Context, loginUserID string, peer *nbpeer.Peer) error
|
||||
|
||||
func peerLoginExpired(ctx context.Context, peer *nbpeer.Peer, settings *types.Settings) bool {
|
||||
expired, expiresIn := peer.LoginExpired(settings.PeerLoginExpiration)
|
||||
expired = settings.PeerLoginExpirationEnabled && expired
|
||||
if expired || peer.Status.LoginExpired {
|
||||
log.WithContext(ctx).Debugf("peer's %s login expired %v ago", peer.ID, expiresIn)
|
||||
if settings.PeerLoginExpirationEnabled && expired {
|
||||
log.WithContext(ctx).Debugf("peer's %s login expired %v ago", peer.ID, -expiresIn)
|
||||
return true
|
||||
}
|
||||
if peer.Status.LoginExpired {
|
||||
log.WithContext(ctx).Debugf("peer's %s login is marked as expired", peer.ID)
|
||||
return true
|
||||
}
|
||||
return false
|
||||
@@ -1643,7 +1646,9 @@ func (am *DefaultAccountManager) UpdateAccountPeer(ctx context.Context, accountI
|
||||
|
||||
// getNextPeerExpiration returns the minimum duration in which the next peer of the account will expire if it was found.
|
||||
// If there is no peer that expires this function returns false and a duration of 0.
|
||||
// This function only considers peers that haven't been expired yet and that are connected.
|
||||
// This function only considers peers that haven't been expired yet. Offline peers count too:
|
||||
// a running job is never re-armed on connect, so a peer that reconnects with an old login
|
||||
// must already be part of the scheduled run.
|
||||
func (am *DefaultAccountManager) getNextPeerExpiration(ctx context.Context, accountID string) (time.Duration, bool) {
|
||||
peersWithExpiry, err := am.Store.GetAccountPeersWithExpiration(ctx, store.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
@@ -1663,8 +1668,7 @@ func (am *DefaultAccountManager) getNextPeerExpiration(ctx context.Context, acco
|
||||
|
||||
var nextExpiry *time.Duration
|
||||
for _, peer := range peersWithExpiry {
|
||||
// consider only connected peers because others will require login on connecting to the management server
|
||||
if peer.Status.LoginExpired || !peer.Status.Connected {
|
||||
if peer.Status.LoginExpired {
|
||||
continue
|
||||
}
|
||||
_, duration := peer.LoginExpired(settings.PeerLoginExpiration)
|
||||
|
||||
@@ -117,6 +117,7 @@ func (wm *DefaultScheduler) Schedule(ctx context.Context, in time.Duration, ID s
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(in)
|
||||
period := in
|
||||
|
||||
wm.jobs[ID] = cancel
|
||||
log.WithContext(ctx).Debugf("scheduled a job %s to run in %s. There are %d total jobs scheduled.", ID, in.String(), len(wm.jobs))
|
||||
@@ -136,14 +137,18 @@ func (wm *DefaultScheduler) Schedule(ctx context.Context, in time.Duration, ID s
|
||||
if !reschedule {
|
||||
wm.mu.Lock()
|
||||
defer wm.mu.Unlock()
|
||||
delete(wm.jobs, ID)
|
||||
// A Cancel during job() may have registered a replacement under this ID.
|
||||
if current, ok := wm.jobs[ID]; ok && current == cancel {
|
||||
delete(wm.jobs, ID)
|
||||
}
|
||||
log.WithContext(ctx).Debugf("job %s is not scheduled to run again", ID)
|
||||
ticker.Stop()
|
||||
return
|
||||
}
|
||||
// we need this comparison to avoid resetting the ticker with the same duration and missing the current elapsesed time
|
||||
if runIn != in {
|
||||
if runIn != period {
|
||||
ticker.Reset(runIn)
|
||||
period = runIn
|
||||
}
|
||||
case <-cancel:
|
||||
log.WithContext(ctx).Debugf("job %s was canceled, stopping timer", ID)
|
||||
|
||||
@@ -6,10 +6,12 @@ import (
|
||||
"math/rand"
|
||||
"runtime"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestScheduler_Performance(t *testing.T) {
|
||||
@@ -150,3 +152,90 @@ func TestScheduler_Schedule(t *testing.T) {
|
||||
scheduler.cancel(context.Background(), jobID)
|
||||
|
||||
}
|
||||
|
||||
func TestScheduler_Schedule_ResetsTickerAfterReturningInitialInterval(t *testing.T) {
|
||||
jobID := "test-scheduler-job-2"
|
||||
scheduler := NewDefaultScheduler()
|
||||
defer scheduler.Cancel(context.Background(), []string{jobID})
|
||||
|
||||
initial := 30 * time.Millisecond
|
||||
stretched := 400 * time.Millisecond
|
||||
runs := make(chan time.Time, 3)
|
||||
count := 0
|
||||
// The first run stretches the period; the second returns the initial interval again,
|
||||
// which must shrink the period back instead of keeping the stretched one.
|
||||
job := func() (nextRunIn time.Duration, reschedule bool) {
|
||||
count++
|
||||
runs <- time.Now()
|
||||
switch count {
|
||||
case 1:
|
||||
return stretched, true
|
||||
case 2:
|
||||
return initial, true
|
||||
default:
|
||||
return 0, false
|
||||
}
|
||||
}
|
||||
scheduler.Schedule(context.Background(), initial, jobID, job)
|
||||
|
||||
var stamps []time.Time
|
||||
for len(stamps) < 3 {
|
||||
select {
|
||||
case ts := <-runs:
|
||||
stamps = append(stamps, ts)
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatalf("timed out after %d runs", len(stamps))
|
||||
}
|
||||
}
|
||||
assert.Less(t, stamps[2].Sub(stamps[1]), stretched/2, "returning the initial interval must reset the stretched ticker")
|
||||
}
|
||||
|
||||
func TestScheduler_Schedule_StaleCompletionKeepsReplacement(t *testing.T) {
|
||||
jobID := "test-scheduler-job-3"
|
||||
scheduler := NewDefaultScheduler()
|
||||
defer scheduler.Cancel(context.Background(), []string{jobID})
|
||||
|
||||
started := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
staleJob := func() (nextRunIn time.Duration, reschedule bool) {
|
||||
close(started)
|
||||
<-release
|
||||
return 0, false
|
||||
}
|
||||
scheduler.Schedule(context.Background(), 10*time.Millisecond, jobID, staleJob)
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("timed out waiting for the first job to start")
|
||||
}
|
||||
|
||||
// Cancel the job while it is still executing and register a replacement under the
|
||||
// same ID, as the expiration paths do on a settings change.
|
||||
scheduler.Cancel(context.Background(), []string{jobID})
|
||||
var replacementRuns atomic.Int32
|
||||
scheduler.Schedule(context.Background(), 20*time.Millisecond, jobID, func() (nextRunIn time.Duration, reschedule bool) {
|
||||
replacementRuns.Add(1)
|
||||
return 20 * time.Millisecond, true
|
||||
})
|
||||
require.True(t, scheduler.IsSchedulerRunning(jobID), "replacement must be registered")
|
||||
|
||||
// The stale job now completes without rescheduling; its cleanup must leave the
|
||||
// replacement's entry in place.
|
||||
close(release)
|
||||
assert.Never(t, func() bool { return !scheduler.IsSchedulerRunning(jobID) }, 200*time.Millisecond, 10*time.Millisecond,
|
||||
"stale completion must not drop the replacement job")
|
||||
|
||||
var duplicateRuns atomic.Int32
|
||||
scheduler.Schedule(context.Background(), 10*time.Millisecond, jobID, func() (nextRunIn time.Duration, reschedule bool) {
|
||||
duplicateRuns.Add(1)
|
||||
return 10 * time.Millisecond, true
|
||||
})
|
||||
assert.Never(t, func() bool { return duplicateRuns.Load() > 0 }, 100*time.Millisecond, 10*time.Millisecond,
|
||||
"a duplicate schedule must be refused while the replacement is registered")
|
||||
|
||||
scheduler.Cancel(context.Background(), []string{jobID})
|
||||
assert.False(t, scheduler.IsSchedulerRunning(jobID), "cancel must find and remove the replacement")
|
||||
runsAfterCancel := replacementRuns.Load()
|
||||
assert.Never(t, func() bool { return replacementRuns.Load() > runsAfterCancel+1 }, 150*time.Millisecond, 10*time.Millisecond,
|
||||
"the replacement must stop after cancel")
|
||||
}
|
||||
|
||||
@@ -3160,7 +3160,12 @@ func NewMysqlStore(ctx context.Context, dsn string, metrics telemetry.AppMetrics
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return NewSqlStore(ctx, db, types.MysqlStoreEngine, metrics, skipMigration)
|
||||
store, err := NewSqlStore(ctx, db, types.MysqlStoreEngine, metrics, skipMigration)
|
||||
if err != nil {
|
||||
closeGormDB(db)
|
||||
return nil, err
|
||||
}
|
||||
return store, nil
|
||||
}
|
||||
|
||||
func getGormConfig() *gorm.Config {
|
||||
@@ -3219,23 +3224,20 @@ func NewSqliteStoreFromFileStore(ctx context.Context, fileStore *FileStore, data
|
||||
|
||||
// NewPostgresqlStoreFromSqlStore restores a store from SqlStore and stores Postgres DB.
|
||||
func NewPostgresqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, dsn string, metrics telemetry.AppMetrics) (*SqlStore, error) {
|
||||
store, err := NewPostgresqlStoreForTests(ctx, dsn, metrics, false)
|
||||
return newPostgresqlStoreFromSqlStore(ctx, sqliteStore, dsn, metrics, false)
|
||||
}
|
||||
|
||||
func newPostgresqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) {
|
||||
store, err := NewPostgresqlStoreForTests(ctx, dsn, metrics, skipMigration)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
err = store.SaveInstallationID(ctx, sqliteStore.GetInstallationID())
|
||||
if err != nil {
|
||||
if err := seedFromSqliteStore(ctx, store, sqliteStore); err != nil {
|
||||
closeStore(ctx, store)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, account := range sqliteStore.GetAllAccounts(ctx) {
|
||||
err := store.SaveAccount(ctx, account)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return store, nil
|
||||
}
|
||||
|
||||
@@ -3247,11 +3249,14 @@ func NewPostgresqlStoreForTests(ctx context.Context, dsn string, metrics telemet
|
||||
}
|
||||
pool, err := connectToPgDbForTests(context.Background(), dsn)
|
||||
if err != nil {
|
||||
closeGormDB(db)
|
||||
return nil, err
|
||||
}
|
||||
store, err := NewSqlStore(ctx, db, types.PostgresStoreEngine, metrics, skipMigration)
|
||||
if err != nil {
|
||||
// Release the sessions, or the caller cannot drop the database.
|
||||
pool.Close()
|
||||
closeGormDB(db)
|
||||
return nil, err
|
||||
}
|
||||
store.pool = pool
|
||||
@@ -3285,22 +3290,42 @@ func connectToPgDbForTests(ctx context.Context, dsn string) (*pgxpool.Pool, erro
|
||||
|
||||
// NewMysqlStoreFromSqlStore restores a store from SqlStore and stores MySQL DB.
|
||||
func NewMysqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, dsn string, metrics telemetry.AppMetrics) (*SqlStore, error) {
|
||||
store, err := NewMysqlStore(ctx, dsn, metrics, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return newMysqlStoreFromSqlStore(ctx, sqliteStore, dsn, metrics, false)
|
||||
}
|
||||
|
||||
err = store.SaveInstallationID(ctx, sqliteStore.GetInstallationID())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
// seedFromSqliteStore copies the installation ID and the accounts of the
|
||||
// sqlite seed store into a freshly created engine store.
|
||||
func seedFromSqliteStore(ctx context.Context, store, sqliteStore *SqlStore) error {
|
||||
if err := store.SaveInstallationID(ctx, sqliteStore.GetInstallationID()); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, account := range sqliteStore.GetAllAccounts(ctx) {
|
||||
err := store.SaveAccount(ctx, account)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
if err := store.SaveAccount(ctx, account); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// closeStore releases a store that is not handed to the caller, so a failed
|
||||
// seed does not leak its connection and pool.
|
||||
func closeStore(ctx context.Context, store *SqlStore) {
|
||||
store.Close(ctx)
|
||||
if store.pool != nil {
|
||||
store.pool.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func newMysqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) {
|
||||
store, err := NewMysqlStore(ctx, dsn, metrics, skipMigration)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := seedFromSqliteStore(ctx, store, sqliteStore); err != nil {
|
||||
closeStore(ctx, store)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return store, nil
|
||||
}
|
||||
@@ -3479,7 +3504,7 @@ func (s *SqlStore) GetPeerGroups(ctx context.Context, lockStrength LockingStreng
|
||||
var groups []*types.Group
|
||||
query := tx.
|
||||
Joins("JOIN group_peers ON group_peers.group_id = groups.id").
|
||||
Where("group_peers.peer_id = ?", peerId).
|
||||
Where("groups.account_id = ? AND group_peers.peer_id = ?", accountId, peerId).
|
||||
Preload(clause.Associations).
|
||||
Find(&groups)
|
||||
|
||||
@@ -5059,7 +5084,7 @@ func (s *SqlStore) GetPeersByGroupIDs(ctx context.Context, accountID string, gro
|
||||
Select("DISTINCT peer_id").
|
||||
Where("account_id = ? AND group_id IN ?", accountID, groupIDs)
|
||||
|
||||
result := s.db.Where("id IN (?)", peerIDsSubquery).Find(&peers)
|
||||
result := s.db.Where("account_id = ? AND id IN (?)", accountID, peerIDsSubquery).Find(&peers)
|
||||
if result.Error != nil {
|
||||
log.WithContext(ctx).Errorf("failed to get peers by group IDs: %s", result.Error)
|
||||
return nil, status.Errorf(status.Internal, "failed to get peers by group IDs")
|
||||
@@ -5692,6 +5717,23 @@ func (s *SqlStore) ListCustomDomains(ctx context.Context, accountID string) ([]*
|
||||
return domains, nil
|
||||
}
|
||||
|
||||
// GetCustomDomainByName returns the custom domain row holding the given name,
|
||||
// regardless of which account owns it.
|
||||
func (s *SqlStore) GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) {
|
||||
customDomain := &domain.Domain{}
|
||||
result := s.db.Take(customDomain, "domain = ?", domainName)
|
||||
if result.Error != nil {
|
||||
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
||||
return nil, status.Errorf(status.NotFound, "custom domain %s not found", domainName)
|
||||
}
|
||||
|
||||
log.WithContext(ctx).Errorf("failed to get custom domain by name from store: %v", result.Error)
|
||||
return nil, status.Errorf(status.Internal, "failed to get custom domain from store")
|
||||
}
|
||||
|
||||
return customDomain, nil
|
||||
}
|
||||
|
||||
func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error) {
|
||||
newDomain := &domain.Domain{
|
||||
ID: xid.New().String(), // Generate our own ID because gorm doesn't always configure the database to handle this for us.
|
||||
@@ -5701,8 +5743,24 @@ func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, dom
|
||||
Type: domain.TypeCustom,
|
||||
Validated: validated,
|
||||
}
|
||||
if !validated {
|
||||
expiresAt := time.Now().UTC().Add(domain.ValidationTTL)
|
||||
newDomain.ValidationExpiresAt = &expiresAt
|
||||
}
|
||||
result := s.db.Create(newDomain)
|
||||
if result.Error != nil {
|
||||
// The unique index is the last guard when two requests clear the
|
||||
// manager's availability check at the same time. The one that loses the
|
||||
// insert is a conflict, not an internal failure.
|
||||
var count int64
|
||||
if err := s.db.Model(&domain.Domain{}).Where("domain = ?", domainName).Count(&count).Error; err == nil && count > 0 {
|
||||
// The insert error is logged even on this path: the name being taken
|
||||
// is what the caller has to act on, but if the insert also failed for
|
||||
// an unrelated reason the operator still needs to see it.
|
||||
log.WithContext(ctx).Warnf("create reverse proxy custom domain %s rejected, name already registered: %v", domainName, result.Error)
|
||||
return nil, status.Errorf(status.AlreadyExists, "domain %s is already registered", domainName)
|
||||
}
|
||||
|
||||
log.WithContext(ctx).Errorf("failed to create reverse proxy custom domain to store: %v", result.Error)
|
||||
return nil, status.Errorf(status.Internal, "failed to create reverse proxy custom domain to store")
|
||||
}
|
||||
@@ -5710,12 +5768,21 @@ func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, dom
|
||||
return newDomain, nil
|
||||
}
|
||||
|
||||
// UpdateCustomDomain completes validation only while the original registration is pending.
|
||||
func (s *SqlStore) UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error) {
|
||||
d.AccountID = accountID
|
||||
result := s.db.Select("*").Save(d)
|
||||
if !d.Validated {
|
||||
return nil, status.Errorf(status.InvalidArgument, "custom domain update must complete validation")
|
||||
}
|
||||
result := s.db.WithContext(ctx).Model(&domain.Domain{}).
|
||||
Where(accountAndIDQueryCondition, accountID, d.ID).
|
||||
Where("domain = ? AND target_cluster = ?", d.Domain, d.TargetCluster).
|
||||
Where("validated = ? AND validation_expires_at > ?", false, time.Now().UTC()).
|
||||
Update("validated", true)
|
||||
if result.Error != nil {
|
||||
log.WithContext(ctx).Errorf("failed to update reverse proxy custom domain to store: %v", result.Error)
|
||||
return nil, status.Errorf(status.Internal, "failed to update reverse proxy custom domain to store")
|
||||
return nil, fmt.Errorf("validate custom domain in store: %w", result.Error)
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "custom domain registration is no longer pending validation")
|
||||
}
|
||||
|
||||
return d, nil
|
||||
@@ -6410,6 +6477,25 @@ func (s *SqlStore) IsClusterAddressConflicting(ctx context.Context, clusterAddre
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// HasForeignAccountProxyAtHost reports whether a proxy owned by a different
|
||||
// account declares this host. Shared proxies (account_id IS NULL) are not
|
||||
// foreign: a shared cluster is what most accounts pin their agent network
|
||||
// gateway to. The match folds case because proxies declare their address as
|
||||
// the operator spelled it while the caller's host is normalised; that costs a
|
||||
// scan of the proxies table, taken once per account when its gateway is
|
||||
// bootstrapped, not on the per-connect path IsClusterAddressConflicting serves.
|
||||
func (s *SqlStore) HasForeignAccountProxyAtHost(ctx context.Context, host, accountID string) (bool, error) {
|
||||
var count int64
|
||||
result := s.db.
|
||||
Model(&proxy.Proxy{}).
|
||||
Where("LOWER(cluster_address) = LOWER(?) AND account_id IS NOT NULL AND account_id != ?", host, accountID).
|
||||
Count(&count)
|
||||
if result.Error != nil {
|
||||
return false, status.Errorf(status.Internal, "check proxy host ownership: %v", result.Error)
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (s *SqlStore) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error {
|
||||
result := s.db.
|
||||
Where("cluster_address = ? AND account_id = ?", clusterAddress, accountID).
|
||||
|
||||
@@ -315,6 +315,36 @@ func (s *SqlStore) GetAllAgentNetworkSettings(ctx context.Context, lockStrength
|
||||
return settings, nil
|
||||
}
|
||||
|
||||
// HasGatewayClusterPinnedByOtherAccount reports whether another account has a
|
||||
// labeled agent network gateway pinned beneath host, making host its cluster.
|
||||
// A self-addressed endpoint on the very same hostname is not counted: that
|
||||
// collision is the domain unique index's to refuse, as a conflict. Case-folded,
|
||||
// since a settings row written before hostnames were normalised may carry
|
||||
// capitals; one row per account, so the scan is cheap.
|
||||
func (s *SqlStore) HasGatewayClusterPinnedByOtherAccount(ctx context.Context, host, accountID string) (bool, error) {
|
||||
return s.countGatewayRowsByOtherAccount(ctx, "LOWER(proxy_address) = LOWER(?) AND LOWER(domain) <> LOWER(proxy_address)", host, accountID)
|
||||
}
|
||||
|
||||
// HasGatewayEndpointByOtherAccount reports whether host is another account's
|
||||
// agent network endpoint hostname (domain). Case-folded for the same reason as
|
||||
// HasGatewayClusterPinnedByOtherAccount.
|
||||
func (s *SqlStore) HasGatewayEndpointByOtherAccount(ctx context.Context, host, accountID string) (bool, error) {
|
||||
return s.countGatewayRowsByOtherAccount(ctx, "LOWER(domain) = LOWER(?)", host, accountID)
|
||||
}
|
||||
|
||||
func (s *SqlStore) countGatewayRowsByOtherAccount(ctx context.Context, predicate, host, accountID string) (bool, error) {
|
||||
var count int64
|
||||
result := s.db.
|
||||
Model(&agentNetworkTypes.Settings{}).
|
||||
Where(predicate+" AND account_id != ?", host, accountID).
|
||||
Count(&count)
|
||||
if result.Error != nil {
|
||||
log.WithContext(ctx).Errorf("failed to check agent network gateway claims at host: %v", result.Error)
|
||||
return false, status.Errorf(status.Internal, "check agent network gateway claims at host")
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// GetAgentNetworkSettingsByProxyAddress returns every Settings row whose
|
||||
// gateway is served by the proxy declaring the given cluster address. Used by
|
||||
// cluster-scoped synthesis to find the accounts a shared proxy serves.
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
// GetExpiredCustomDomains lists pending registrations in stable batches across accounts.
|
||||
func (s *SqlStore) GetExpiredCustomDomains(ctx context.Context, now time.Time, afterID domain.ID, limit int) ([]*domain.Domain, error) {
|
||||
var domains []*domain.Domain
|
||||
result := s.db.WithContext(ctx).
|
||||
Where("validated = ? AND validation_expires_at <= ? AND id > ?", false, now, string(afterID)).
|
||||
Order("id").Limit(limit).Find(&domains)
|
||||
if result.Error != nil {
|
||||
return nil, fmt.Errorf("list expired custom domains: %w", result.Error)
|
||||
}
|
||||
return domains, nil
|
||||
}
|
||||
|
||||
// DeleteExpiredCustomDomain deletes an expired registration only if no service uses its namespace.
|
||||
func (s *SqlStore) DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error) {
|
||||
db := s.db.WithContext(ctx)
|
||||
services := customDomainServices(db, d)
|
||||
result := db.Where(accountAndIDQueryCondition, d.AccountID, d.ID).
|
||||
Where("domain = ? AND validated = ? AND validation_expires_at <= ?", d.Domain, false, now).
|
||||
Where("NOT EXISTS (?)", services.Select("1")).Delete(&domain.Domain{})
|
||||
if result.Error != nil {
|
||||
return false, fmt.Errorf("delete expired custom domain: %w", result.Error)
|
||||
}
|
||||
if result.RowsAffected > 0 {
|
||||
return true, nil
|
||||
}
|
||||
var count int64
|
||||
if err := customDomainServices(db, d).Count(&count).Error; err != nil {
|
||||
return false, fmt.Errorf("check expired custom domain services: %w", err)
|
||||
}
|
||||
if count > 0 {
|
||||
return false, status.Errorf(status.PreconditionFailed, "expired custom domain still has dependent services")
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func customDomainServices(db *gorm.DB, d *domain.Domain) *gorm.DB {
|
||||
name := strings.ToLower(strings.TrimSuffix(d.Domain, "."))
|
||||
// Shared domain validation permits underscores, and older rows may contain
|
||||
// other LIKE metacharacters.
|
||||
escaped := strings.NewReplacer("!", "!!", "%", "!%", "_", "!_").Replace(name)
|
||||
return db.Model(&rpservice.Service{}).Where(
|
||||
"LOWER(domain) IN ? OR LOWER(domain) LIKE ? ESCAPE '!' OR LOWER(domain) LIKE ? ESCAPE '!'",
|
||||
[]string{name, name + "."}, "%."+escaped, "%."+escaped+".",
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"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"
|
||||
)
|
||||
|
||||
func TestDeleteExpiredCustomDomain_ServiceDependencies(t *testing.T) {
|
||||
runTestForAllEngines(t, "", func(t *testing.T, store Store) {
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC()
|
||||
db := store.(*SqlStore).db
|
||||
require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", "")))
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
domainName string
|
||||
serviceHost string
|
||||
protected bool
|
||||
}{
|
||||
{"exact", "example.com", "example.com", true},
|
||||
{"subdomain", "example.com", "deep.app.example.com", true},
|
||||
{"case", "example.com", "APP.EXAMPLE.COM.", true},
|
||||
{"suffix-boundary", "example.com", "notexample.com", false},
|
||||
{"literal underscore", "a_b.example.com", "app.a_b.example.com", true},
|
||||
{"underscore wildcard", "a_b.example.com", "app.axb.example.com", false},
|
||||
{"legacy percent wildcard", "a%b.example.com", "app.axxb.example.com", false},
|
||||
{"legacy escape character", "a!b.example.com", "app.ab.example.com", false},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
d, err := store.CreateCustomDomain(ctx, "owner", tt.domainName, "cluster", false)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, db.Model(d).Update("validation_expires_at", now.Add(-time.Hour)).Error)
|
||||
svc := &rpservice.Service{ID: "legacy", AccountID: "owner", Domain: tt.serviceHost}
|
||||
require.NoError(t, store.CreateService(ctx, svc))
|
||||
deleted, err := store.DeleteExpiredCustomDomain(ctx, d, now)
|
||||
if tt.protected {
|
||||
require.Error(t, err)
|
||||
assert.False(t, deleted, "service namespaces must remain reserved")
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
assert.True(t, deleted, "a hostname outside the namespace must not prevent cleanup")
|
||||
}
|
||||
require.NoError(t, db.Delete(svc).Error)
|
||||
require.NoError(t, db.Delete(d).Error)
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestDeleteExpiredCustomDomain_RechecksValidation(t *testing.T) {
|
||||
runTestForAllEngines(t, "", func(t *testing.T, store Store) {
|
||||
ctx := context.Background()
|
||||
require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", "")))
|
||||
d, err := store.CreateCustomDomain(ctx, "owner", "validated.example.com", "cluster", false)
|
||||
require.NoError(t, err)
|
||||
d, err = store.GetCustomDomain(ctx, "owner", d.ID)
|
||||
require.NoError(t, err)
|
||||
stale := d.Copy()
|
||||
d.Validated = true
|
||||
_, err = store.UpdateCustomDomain(ctx, "owner", d)
|
||||
require.NoError(t, err)
|
||||
deleted, err := store.DeleteExpiredCustomDomain(ctx, stale, time.Now().Add(domain.ValidationTTL))
|
||||
require.NoError(t, err)
|
||||
assert.False(t, deleted, "a stale cleanup candidate must not delete a validated registration")
|
||||
stored, err := store.GetCustomDomain(ctx, "owner", d.ID)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, stored.Validated, "the validated registration must remain usable")
|
||||
require.NotNil(t, stored.ValidationExpiresAt)
|
||||
assert.Equal(t, stale.ValidationExpiresAt, stored.ValidationExpiresAt, "validation must preserve the original deadline")
|
||||
})
|
||||
}
|
||||
@@ -64,7 +64,7 @@ func assertGetAccountLoadsCustomDomains(t *testing.T, store Store) {
|
||||
|
||||
_, err := store.CreateCustomDomain(ctx, accountID, "example.com", "eu.proxy.netbird.io", true)
|
||||
require.NoError(t, err, "creating the first custom domain must succeed")
|
||||
_, err = store.CreateCustomDomain(ctx, accountID, "apps.acme.io", "us.proxy.netbird.io", false)
|
||||
pending, err := store.CreateCustomDomain(ctx, accountID, "apps.acme.io", "us.proxy.netbird.io", false)
|
||||
require.NoError(t, err, "creating the second custom domain must succeed")
|
||||
|
||||
account, err := store.GetAccount(ctx, accountID)
|
||||
@@ -75,6 +75,10 @@ func assertGetAccountLoadsCustomDomains(t *testing.T, store Store) {
|
||||
for _, d := range account.Domains {
|
||||
require.NotNil(t, d)
|
||||
byDomain[d.Domain] = d.TargetCluster
|
||||
if d.ID == pending.ID {
|
||||
require.NotNil(t, d.ValidationExpiresAt)
|
||||
assert.WithinDuration(t, *pending.ValidationExpiresAt, *d.ValidationExpiresAt, time.Millisecond, "both account loaders must preserve the validation deadline")
|
||||
}
|
||||
}
|
||||
assert.Equal(t, "eu.proxy.netbird.io", byDomain["example.com"], "custom domain must carry its target cluster")
|
||||
assert.Equal(t, "us.proxy.netbird.io", byDomain["apps.acme.io"], "custom domain must carry its target cluster")
|
||||
|
||||
@@ -2847,6 +2847,14 @@ func TestSqlStore_GetPeerGroups(t *testing.T) {
|
||||
groups, err = store.GetPeerGroups(context.Background(), LockingStrengthNone, accountID, peerID)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, groups, 2)
|
||||
|
||||
foreignPeerID := "foreign-peer"
|
||||
err = store.AddPeerToGroup(context.Background(), accountID, foreignPeerID, "cfefqs706sqkneg59g4h")
|
||||
require.NoError(t, err)
|
||||
|
||||
groups, err = store.GetPeerGroups(context.Background(), LockingStrengthNone, "other-account", foreignPeerID)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, groups, "groups of another account must not be returned")
|
||||
}
|
||||
|
||||
func TestSqlStore_GetAccountPeers(t *testing.T) {
|
||||
@@ -4042,9 +4050,15 @@ func TestSqlStore_GetPeersByGroupIDs(t *testing.T) {
|
||||
}
|
||||
require.NoError(t, store.CreateGroups(ctx, accountID, groups))
|
||||
|
||||
otherAccount := newAccountWithId(ctx, "other-account", "other-user", "")
|
||||
require.NoError(t, store.SaveAccount(ctx, otherAccount))
|
||||
foreignPeer := &nbpeer.Peer{ID: "foreign-peer", AccountID: otherAccount.Id}
|
||||
require.NoError(t, store.AddPeerToAccount(ctx, foreignPeer))
|
||||
|
||||
require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer1, group1ID))
|
||||
require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer2, group1ID))
|
||||
require.NoError(t, store.AddPeerToGroup(ctx, accountID, peer1, group2ID))
|
||||
require.NoError(t, store.AddPeerToGroup(ctx, accountID, foreignPeer.ID, group1ID))
|
||||
|
||||
peers, err := store.GetPeersByGroupIDs(ctx, accountID, tt.groupIDs)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -4,6 +4,7 @@ package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
@@ -15,6 +16,7 @@ import (
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
@@ -302,8 +304,11 @@ type Store interface {
|
||||
GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error)
|
||||
ListFreeDomains(ctx context.Context, accountID string) ([]string, error)
|
||||
ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error)
|
||||
GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error)
|
||||
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)
|
||||
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)
|
||||
DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error
|
||||
|
||||
CreateAccessLog(ctx context.Context, log *accesslogs.AccessLogEntry) error
|
||||
@@ -337,6 +342,9 @@ type Store interface {
|
||||
CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error)
|
||||
IsClusterAddressConflicting(ctx context.Context, clusterAddress, accountID string) (bool, error)
|
||||
HasActiveProxyAtClusterAddress(ctx context.Context, clusterAddress string) (bool, error)
|
||||
HasForeignAccountProxyAtHost(ctx context.Context, host, accountID string) (bool, error)
|
||||
HasGatewayClusterPinnedByOtherAccount(ctx context.Context, host, accountID string) (bool, error)
|
||||
HasGatewayEndpointByOtherAccount(ctx context.Context, host, accountID string) (bool, error)
|
||||
DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error
|
||||
|
||||
GetCustomDomainsCounts(ctx context.Context) (total int64, validated int64, err error)
|
||||
@@ -641,6 +649,9 @@ func migratePostAuto(ctx context.Context, db *gorm.DB) error {
|
||||
|
||||
func getMigrationsPostAuto(ctx context.Context) []migrationFunc {
|
||||
return []migrationFunc{
|
||||
func(db *gorm.DB) error {
|
||||
return migration.MigrateCustomDomainValidationExpiry(ctx, db)
|
||||
},
|
||||
func(db *gorm.DB) error {
|
||||
return migration.CreateIndexIfNotExists[nbpeer.Peer](ctx, db, "idx_account_ip", "account_id", "ip")
|
||||
},
|
||||
@@ -726,6 +737,7 @@ func NewTestStoreFromSQL(ctx context.Context, filename string, dataDir string) (
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
store.Close(ctx)
|
||||
return nil, nil, fmt.Errorf("failed to create test store after %d attempts: %v", maxRetries, err)
|
||||
}
|
||||
|
||||
@@ -752,14 +764,15 @@ func addAllGroupToAccount(ctx context.Context, store Store) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func getSqlStoreEngine(ctx context.Context, store *SqlStore, kind types.Engine) (Store, func(), error) {
|
||||
func getSqlStoreEngine(ctx context.Context, sqliteStore *SqlStore, kind types.Engine) (Store, func(), error) {
|
||||
store := sqliteStore
|
||||
var cleanup func()
|
||||
var err error
|
||||
switch kind {
|
||||
case types.PostgresStoreEngine:
|
||||
store, cleanup, err = newReusedPostgresStore(ctx, store, kind)
|
||||
store, cleanup, err = newReusedPostgresStore(ctx, sqliteStore, kind)
|
||||
case types.MysqlStoreEngine:
|
||||
store, cleanup, err = newReusedMysqlStore(ctx, store, kind)
|
||||
store, cleanup, err = newReusedMysqlStore(ctx, sqliteStore, kind)
|
||||
default:
|
||||
cleanup = func() {
|
||||
// sqlite doesn't need to be cleaned up
|
||||
@@ -775,6 +788,11 @@ func getSqlStoreEngine(ctx context.Context, store *SqlStore, kind types.Engine)
|
||||
if store.pool != nil {
|
||||
store.pool.Close()
|
||||
}
|
||||
if store != sqliteStore {
|
||||
// The sqlite store only seeded the engine under test; without this
|
||||
// every test leaks its connection and the opener goroutines.
|
||||
sqliteStore.Close(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
return store, closeConnection, nil
|
||||
@@ -799,19 +817,23 @@ func newReusedPostgresStore(ctx context.Context, store *SqlStore, kind types.Eng
|
||||
return nil, nil, fmt.Errorf("failed to open postgres connection: %v", err)
|
||||
}
|
||||
|
||||
dsn, cleanup, err := createRandomDB(dsn, db, kind)
|
||||
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
sqlDB.Close()
|
||||
template, err := postgresSchemaTemplate(ctx, dsn, db)
|
||||
if err != nil {
|
||||
closeGormDB(db)
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
dsn, cleanup, err := createRandomDB(dsn, db, kind, template)
|
||||
|
||||
closeGormDB(db)
|
||||
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
store, err = NewPostgresqlStoreFromSqlStore(ctx, store, dsn, nil)
|
||||
store, err = newPostgresqlStoreFromSqlStore(ctx, store, dsn, nil, true)
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
@@ -844,7 +866,13 @@ func newReusedMysqlStore(ctx context.Context, store *SqlStore, kind types.Engine
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
sqlDB.SetMaxIdleConns(1)
|
||||
|
||||
dsn, cleanup, err := createRandomDB(dsn, db, kind)
|
||||
tableDDL, err := mysqlSchemaTemplate(ctx, dsn, db)
|
||||
if err != nil {
|
||||
sqlDB.Close()
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
dsn, cleanup, err := createRandomDB(dsn, db, kind, "")
|
||||
|
||||
sqlDB.Close()
|
||||
|
||||
@@ -852,14 +880,200 @@ func newReusedMysqlStore(ctx context.Context, store *SqlStore, kind types.Engine
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
store, err = NewMysqlStoreFromSqlStore(ctx, store, dsn, nil)
|
||||
if err := cloneMysqlSchema(ctx, dsn, tableDDL); err != nil {
|
||||
cleanup()
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
store, err = newMysqlStoreFromSqlStore(ctx, store, dsn, nil, true)
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
return store, cleanup, nil
|
||||
}
|
||||
|
||||
// schemaTemplates remembers, per engine and server, a database that went
|
||||
// through the full migration once in this process. Every later test database
|
||||
// is cloned from it, so a test pays for CREATE DATABASE and a schema copy
|
||||
// instead of the 40-table AutoMigrate plus every pre and post migration, which
|
||||
// is what made each MySQL test store cost well over a second in CI.
|
||||
var (
|
||||
schemaTemplatesMu sync.Mutex
|
||||
schemaTemplates = map[string]*schemaTemplate{}
|
||||
)
|
||||
|
||||
type schemaTemplate struct {
|
||||
dbName string
|
||||
// tableDDL holds the CREATE TABLE statements of the template. MySQL has no
|
||||
// server-side database template, so the schema is replayed statement by
|
||||
// statement into each test database.
|
||||
tableDDL []string
|
||||
}
|
||||
|
||||
func schemaTemplateKey(engine types.Engine, dsn string) string {
|
||||
return string(engine) + "|" + dsn
|
||||
}
|
||||
|
||||
func newTestDBName(prefix string) string {
|
||||
return fmt.Sprintf("%s_%s", prefix, strings.ReplaceAll(uuid.New().String(), "-", "_"))
|
||||
}
|
||||
|
||||
// postgresSchemaTemplate returns the name of a fully migrated database that
|
||||
// CREATE DATABASE ... TEMPLATE can copy, creating it on first use.
|
||||
func postgresSchemaTemplate(ctx context.Context, baseDSN string, admin *gorm.DB) (string, error) {
|
||||
schemaTemplatesMu.Lock()
|
||||
defer schemaTemplatesMu.Unlock()
|
||||
|
||||
key := schemaTemplateKey(types.PostgresStoreEngine, baseDSN)
|
||||
if tpl, ok := schemaTemplates[key]; ok {
|
||||
return tpl.dbName, nil
|
||||
}
|
||||
|
||||
name := newTestDBName("test_template")
|
||||
if err := admin.Exec(fmt.Sprintf("CREATE DATABASE %s", name)).Error; err != nil {
|
||||
return "", fmt.Errorf("create postgres template database: %w", err)
|
||||
}
|
||||
|
||||
tplStore, err := NewPostgresqlStoreForTests(ctx, replaceDBName(baseDSN, name), nil, false)
|
||||
if err != nil {
|
||||
dropDatabase(admin, name)
|
||||
return "", fmt.Errorf("migrate postgres template database: %w", err)
|
||||
}
|
||||
// TEMPLATE refuses a source that still has sessions, so release both handles
|
||||
// before the first clone.
|
||||
tplStore.Close(ctx)
|
||||
if tplStore.pool != nil {
|
||||
tplStore.pool.Close()
|
||||
}
|
||||
|
||||
schemaTemplates[key] = &schemaTemplate{dbName: name}
|
||||
return name, nil
|
||||
}
|
||||
|
||||
// mysqlSchemaTemplate returns the CREATE TABLE statements of a fully migrated
|
||||
// database, migrating one on first use.
|
||||
func mysqlSchemaTemplate(ctx context.Context, baseDSN string, admin *gorm.DB) ([]string, error) {
|
||||
schemaTemplatesMu.Lock()
|
||||
defer schemaTemplatesMu.Unlock()
|
||||
|
||||
key := schemaTemplateKey(types.MysqlStoreEngine, baseDSN)
|
||||
if tpl, ok := schemaTemplates[key]; ok {
|
||||
return tpl.tableDDL, nil
|
||||
}
|
||||
|
||||
name := newTestDBName("test_template")
|
||||
if err := admin.Exec(fmt.Sprintf("CREATE DATABASE %s", name)).Error; err != nil {
|
||||
return nil, fmt.Errorf("create mysql template database: %w", err)
|
||||
}
|
||||
|
||||
tplStore, err := NewMysqlStore(ctx, replaceDBName(baseDSN, name), nil, false)
|
||||
if err != nil {
|
||||
dropDatabase(admin, name)
|
||||
return nil, fmt.Errorf("migrate mysql template database: %w", err)
|
||||
}
|
||||
tableDDL, err := mysqlTableDDL(ctx, tplStore.db, name)
|
||||
tplStore.Close(ctx)
|
||||
if err != nil {
|
||||
dropDatabase(admin, name)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
schemaTemplates[key] = &schemaTemplate{dbName: name, tableDDL: tableDDL}
|
||||
return tableDDL, nil
|
||||
}
|
||||
|
||||
func mysqlTableDDL(ctx context.Context, db *gorm.DB, dbName string) ([]string, error) {
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
tables, err := mysqlTableNames(ctx, sqlDB, dbName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
tableDDL := make([]string, 0, len(tables))
|
||||
for _, table := range tables {
|
||||
var name, createStmt string
|
||||
row := sqlDB.QueryRowContext(ctx, fmt.Sprintf("SHOW CREATE TABLE %s.%s", dbName, table))
|
||||
if err := row.Scan(&name, &createStmt); err != nil {
|
||||
return nil, fmt.Errorf("read create statement of %s: %w", table, err)
|
||||
}
|
||||
tableDDL = append(tableDDL, createStmt)
|
||||
}
|
||||
return tableDDL, nil
|
||||
}
|
||||
|
||||
func mysqlTableNames(ctx context.Context, sqlDB *sql.DB, dbName string) ([]string, error) {
|
||||
rows, err := sqlDB.QueryContext(ctx, fmt.Sprintf("SHOW TABLES FROM %s", dbName))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list template tables: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var tables []string
|
||||
for rows.Next() {
|
||||
var table string
|
||||
if err := rows.Scan(&table); err != nil {
|
||||
return nil, fmt.Errorf("scan template table name: %w", err)
|
||||
}
|
||||
tables = append(tables, table)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("list template tables: %w", err)
|
||||
}
|
||||
return tables, nil
|
||||
}
|
||||
|
||||
// cloneMysqlSchema replays the template's CREATE TABLE statements into the
|
||||
// database the DSN points at.
|
||||
func cloneMysqlSchema(ctx context.Context, dsn string, tableDDL []string) error {
|
||||
db, err := gorm.Open(mysql.Open(dsn+"?charset=utf8&parseTime=True&loc=Local"), getGormConfig())
|
||||
if err != nil {
|
||||
return fmt.Errorf("connect to test database: %w", err)
|
||||
}
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer sqlDB.Close()
|
||||
|
||||
// The statements come out of SHOW TABLES in name order, not dependency
|
||||
// order, and their foreign keys reference tables of the session's default
|
||||
// database. Pin a single connection so the session setting below covers
|
||||
// every statement, and connect straight to the new database so unqualified
|
||||
// references land there.
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
if _, err := sqlDB.ExecContext(ctx, "SET FOREIGN_KEY_CHECKS = 0"); err != nil {
|
||||
return fmt.Errorf("disable foreign key checks: %w", err)
|
||||
}
|
||||
for _, stmt := range tableDDL {
|
||||
if _, err := sqlDB.ExecContext(ctx, stmt); err != nil {
|
||||
return fmt.Errorf("replay table definition: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// dropDatabase removes a template that never became usable, so a failed setup
|
||||
// does not leave it behind on a shared server. The server may still be tearing
|
||||
// down the sessions the failed migration held, so the drop retries while
|
||||
// Postgres reports the database as in use.
|
||||
func dropDatabase(admin *gorm.DB, name string) {
|
||||
if err := execWithTemplateRetry(admin, fmt.Sprintf("DROP DATABASE IF EXISTS %s", name)); err != nil {
|
||||
log.Warnf("failed to drop template database %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func closeGormDB(db *gorm.DB) {
|
||||
if sqlDB, _ := db.DB(); sqlDB != nil {
|
||||
sqlDB.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func openDBWithRetry(dsn string, engine types.Engine, maxRetries int) (*gorm.DB, error) {
|
||||
var db *gorm.DB
|
||||
var err error
|
||||
@@ -885,10 +1099,16 @@ func openDBWithRetry(dsn string, engine types.Engine, maxRetries int) (*gorm.DB,
|
||||
return nil, err
|
||||
}
|
||||
|
||||
func createRandomDB(dsn string, db *gorm.DB, engine types.Engine) (string, func(), error) {
|
||||
dbName := fmt.Sprintf("test_db_%s", strings.ReplaceAll(uuid.New().String(), "-", "_"))
|
||||
// createRandomDB creates a uniquely named database for one test. On postgres a
|
||||
// non-empty template is copied server-side with CREATE DATABASE ... TEMPLATE.
|
||||
func createRandomDB(dsn string, db *gorm.DB, engine types.Engine, template string) (string, func(), error) {
|
||||
dbName := newTestDBName("test_db")
|
||||
|
||||
if err := db.Exec(fmt.Sprintf("CREATE DATABASE %s", dbName)).Error; err != nil {
|
||||
createStmt := fmt.Sprintf("CREATE DATABASE %s", dbName)
|
||||
if template != "" && engine == types.PostgresStoreEngine {
|
||||
createStmt = fmt.Sprintf("CREATE DATABASE %s TEMPLATE %s", dbName, template)
|
||||
}
|
||||
if err := execWithTemplateRetry(db, createStmt); err != nil {
|
||||
return "", nil, fmt.Errorf("failed to create database: %v", err)
|
||||
}
|
||||
|
||||
@@ -954,6 +1174,20 @@ func createRandomDB(dsn string, db *gorm.DB, engine types.Engine) (string, func(
|
||||
return replaceDBName(dsn, dbName), cleanup, nil
|
||||
}
|
||||
|
||||
// execWithTemplateRetry runs a statement, retrying briefly when postgres still
|
||||
// sees the template's just-closed sessions and refuses to copy it.
|
||||
func execWithTemplateRetry(db *gorm.DB, stmt string) error {
|
||||
var err error
|
||||
for attempt := 0; attempt < 20; attempt++ {
|
||||
err = db.Exec(stmt).Error
|
||||
if err == nil || !strings.Contains(err.Error(), "is being accessed by other users") {
|
||||
return err
|
||||
}
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func replaceDBName(dsn, newDBName string) string {
|
||||
re := regexp.MustCompile(`(?P<pre>[:/@])(?P<dbname>[^/?]+)(?P<post>\?|$)`)
|
||||
return re.ReplaceAllString(dsn, `${pre}`+newDBName+`${post}`)
|
||||
|
||||
@@ -555,6 +555,21 @@ func (mr *MockStoreMockRecorder) DeleteDNSRecord(ctx, accountID, zoneID, recordI
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteDNSRecord", reflect.TypeOf((*MockStore)(nil).DeleteDNSRecord), ctx, accountID, zoneID, recordID)
|
||||
}
|
||||
|
||||
// DeleteExpiredCustomDomain mocks base method.
|
||||
func (m *MockStore) DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "DeleteExpiredCustomDomain", ctx, d, now)
|
||||
ret0, _ := ret[0].(bool)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// DeleteExpiredCustomDomain indicates an expected call of DeleteExpiredCustomDomain.
|
||||
func (mr *MockStoreMockRecorder) DeleteExpiredCustomDomain(ctx, d, now any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteExpiredCustomDomain", reflect.TypeOf((*MockStore)(nil).DeleteExpiredCustomDomain), ctx, d, now)
|
||||
}
|
||||
|
||||
// DeleteGroup mocks base method.
|
||||
func (m *MockStore) DeleteGroup(ctx context.Context, accountID, groupID string) error {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -1941,6 +1956,21 @@ func (mr *MockStoreMockRecorder) GetCustomDomain(ctx, accountID, domainID any) *
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetCustomDomain", reflect.TypeOf((*MockStore)(nil).GetCustomDomain), ctx, accountID, domainID)
|
||||
}
|
||||
|
||||
// GetCustomDomainByName mocks base method.
|
||||
func (m *MockStore) GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetCustomDomainByName", ctx, domainName)
|
||||
ret0, _ := ret[0].(*domain.Domain)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetCustomDomainByName indicates an expected call of GetCustomDomainByName.
|
||||
func (mr *MockStoreMockRecorder) GetCustomDomainByName(ctx, domainName any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetCustomDomainByName", reflect.TypeOf((*MockStore)(nil).GetCustomDomainByName), ctx, domainName)
|
||||
}
|
||||
|
||||
// GetCustomDomainsCounts mocks base method.
|
||||
func (m *MockStore) GetCustomDomainsCounts(ctx context.Context) (int64, int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -1987,6 +2017,21 @@ func (mr *MockStoreMockRecorder) GetEmbeddedProxyPeerIDsByCluster(ctx, accountID
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEmbeddedProxyPeerIDsByCluster", reflect.TypeOf((*MockStore)(nil).GetEmbeddedProxyPeerIDsByCluster), ctx, accountID)
|
||||
}
|
||||
|
||||
// GetExpiredCustomDomains mocks base method.
|
||||
func (m *MockStore) GetExpiredCustomDomains(ctx context.Context, now time.Time, afterID domain.ID, limit int) ([]*domain.Domain, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetExpiredCustomDomains", ctx, now, afterID, limit)
|
||||
ret0, _ := ret[0].([]*domain.Domain)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetExpiredCustomDomains indicates an expected call of GetExpiredCustomDomains.
|
||||
func (mr *MockStoreMockRecorder) GetExpiredCustomDomains(ctx, now, afterID, limit any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetExpiredCustomDomains", reflect.TypeOf((*MockStore)(nil).GetExpiredCustomDomains), ctx, now, afterID, limit)
|
||||
}
|
||||
|
||||
// GetExpiredEphemeralServices mocks base method.
|
||||
func (m *MockStore) GetExpiredEphemeralServices(ctx context.Context, ttl time.Duration, limit int) ([]*service.Service, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -3020,6 +3065,51 @@ func (mr *MockStoreMockRecorder) HasActiveProxyAtClusterAddress(ctx, clusterAddr
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasActiveProxyAtClusterAddress", reflect.TypeOf((*MockStore)(nil).HasActiveProxyAtClusterAddress), ctx, clusterAddress)
|
||||
}
|
||||
|
||||
// HasForeignAccountProxyAtHost mocks base method.
|
||||
func (m *MockStore) HasForeignAccountProxyAtHost(ctx context.Context, host, accountID string) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "HasForeignAccountProxyAtHost", ctx, host, accountID)
|
||||
ret0, _ := ret[0].(bool)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// HasForeignAccountProxyAtHost indicates an expected call of HasForeignAccountProxyAtHost.
|
||||
func (mr *MockStoreMockRecorder) HasForeignAccountProxyAtHost(ctx, host, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasForeignAccountProxyAtHost", reflect.TypeOf((*MockStore)(nil).HasForeignAccountProxyAtHost), ctx, host, accountID)
|
||||
}
|
||||
|
||||
// HasGatewayClusterPinnedByOtherAccount mocks base method.
|
||||
func (m *MockStore) HasGatewayClusterPinnedByOtherAccount(ctx context.Context, host, accountID string) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "HasGatewayClusterPinnedByOtherAccount", ctx, host, accountID)
|
||||
ret0, _ := ret[0].(bool)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// HasGatewayClusterPinnedByOtherAccount indicates an expected call of HasGatewayClusterPinnedByOtherAccount.
|
||||
func (mr *MockStoreMockRecorder) HasGatewayClusterPinnedByOtherAccount(ctx, host, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasGatewayClusterPinnedByOtherAccount", reflect.TypeOf((*MockStore)(nil).HasGatewayClusterPinnedByOtherAccount), ctx, host, accountID)
|
||||
}
|
||||
|
||||
// HasGatewayEndpointByOtherAccount mocks base method.
|
||||
func (m *MockStore) HasGatewayEndpointByOtherAccount(ctx context.Context, host, accountID string) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "HasGatewayEndpointByOtherAccount", ctx, host, accountID)
|
||||
ret0, _ := ret[0].(bool)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// HasGatewayEndpointByOtherAccount indicates an expected call of HasGatewayEndpointByOtherAccount.
|
||||
func (mr *MockStoreMockRecorder) HasGatewayEndpointByOtherAccount(ctx, host, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasGatewayEndpointByOtherAccount", reflect.TypeOf((*MockStore)(nil).HasGatewayEndpointByOtherAccount), ctx, host, accountID)
|
||||
}
|
||||
|
||||
// IncrementAgentNetworkConsumption mocks base method.
|
||||
func (m *MockStore) IncrementAgentNetworkConsumption(ctx context.Context, accountID string, kind types.ConsumptionDimension, dimID string, windowSeconds int64, windowStart time.Time, tokensIn, tokensOut int64, costUSD float64) error {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -37,6 +37,18 @@ func CreateMysqlTestContainer() (func(), string, error) {
|
||||
mysql.WithDatabase("testing"),
|
||||
mysql.WithUsername("root"),
|
||||
mysql.WithPassword("testing"),
|
||||
// Every test creates and drops a database with about 40 tables, so with
|
||||
// the server defaults the run is dominated by durability work: each
|
||||
// CREATE TABLE fsyncs the redo log, the binary log and the doublewrite
|
||||
// buffer. None of it protects anything in a container that is discarded
|
||||
// after the run. Tables stay in per-table files on purpose: in the shared
|
||||
// system tablespace the cost of every CREATE and DROP grew with the number
|
||||
// of databases the run had already created.
|
||||
testcontainers.WithCmd("mysqld",
|
||||
"--innodb-flush-log-at-trx-commit=0",
|
||||
"--innodb-doublewrite=OFF",
|
||||
"--skip-log-bin",
|
||||
),
|
||||
testcontainers.WithWaitStrategy(
|
||||
wait.ForLog("/usr/sbin/mysqld: ready for connections").
|
||||
WithOccurrence(1).WithStartupTimeout(15*time.Second).WithPollInterval(100*time.Millisecond),
|
||||
|
||||
@@ -404,7 +404,7 @@ func (a *Account) GetExpiredPeers() []*nbpeer.Peer {
|
||||
|
||||
// GetNextPeerExpiration returns the minimum duration in which the next peer of the account will expire if it was found.
|
||||
// If there is no peer that expires this function returns false and a duration of 0.
|
||||
// This function only considers peers that haven't been expired yet and that are connected.
|
||||
// This function only considers peers that haven't been expired yet, whether connected or not.
|
||||
func (a *Account) GetNextPeerExpiration() (time.Duration, bool) {
|
||||
peersWithExpiry := a.GetPeersWithExpiration()
|
||||
if len(peersWithExpiry) == 0 {
|
||||
@@ -412,8 +412,7 @@ func (a *Account) GetNextPeerExpiration() (time.Duration, bool) {
|
||||
}
|
||||
var nextExpiry *time.Duration
|
||||
for _, peer := range peersWithExpiry {
|
||||
// consider only connected peers because others will require login on connecting to the management server
|
||||
if peer.Status.LoginExpired || !peer.Status.Connected {
|
||||
if peer.Status.LoginExpired {
|
||||
continue
|
||||
}
|
||||
_, duration := peer.LoginExpired(a.Settings.PeerLoginExpiration)
|
||||
|
||||
@@ -3,6 +3,7 @@ package types
|
||||
import (
|
||||
"errors"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Identity provider validation errors
|
||||
@@ -99,7 +100,16 @@ func (idp *IdentityProvider) Validate() error {
|
||||
}
|
||||
if idp.Issuer != "" {
|
||||
parsedURL, err := url.Parse(idp.Issuer)
|
||||
if err != nil || parsedURL.Scheme == "" || parsedURL.Host == "" {
|
||||
if err != nil || parsedURL.Host == "" {
|
||||
return ErrIdentityProviderIssuerInvalid
|
||||
}
|
||||
if parsedURL.Scheme != "https" {
|
||||
return ErrIdentityProviderIssuerInvalid
|
||||
}
|
||||
if parsedURL.User != nil {
|
||||
return ErrIdentityProviderIssuerInvalid
|
||||
}
|
||||
if strings.ContainsAny(idp.Issuer, "?#") {
|
||||
return ErrIdentityProviderIssuerInvalid
|
||||
}
|
||||
}
|
||||
|
||||
@@ -135,3 +135,54 @@ func TestIdentityProvider_Validate(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIdentityProvider_ValidateRejectsNonOriginIssuers(t *testing.T) {
|
||||
issuers := []string{
|
||||
"https://idp.example.com/realms/nb?foo=bar",
|
||||
"https://idp.example.com/realms/nb#section",
|
||||
"https://user:pass@idp.example.com",
|
||||
"ftp://idp.example.com",
|
||||
"ldap://idp.example.com",
|
||||
"http://idp.example.com",
|
||||
}
|
||||
|
||||
for _, issuer := range issuers {
|
||||
t.Run(issuer, func(t *testing.T) {
|
||||
idp := &IdentityProvider{
|
||||
Name: "test",
|
||||
Type: IdentityProviderTypeOIDC,
|
||||
Issuer: issuer,
|
||||
ClientID: "client-id",
|
||||
}
|
||||
assert.ErrorIs(t, idp.Validate(), ErrIdentityProviderIssuerInvalid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIdentityProvider_ValidateAcceptsOriginAndPath(t *testing.T) {
|
||||
for _, issuer := range []string{"https://idp.example.com", "https://idp.example.com/realms/nb", "https://127.0.0.1:5556/dex"} {
|
||||
t.Run(issuer, func(t *testing.T) {
|
||||
idp := &IdentityProvider{
|
||||
Name: "test",
|
||||
Type: IdentityProviderTypeOIDC,
|
||||
Issuer: issuer,
|
||||
ClientID: "client-id",
|
||||
}
|
||||
assert.NoError(t, idp.Validate())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIdentityProviderValidateRejectsBareDelimiters(t *testing.T) {
|
||||
for _, issuer := range []string{"https://idp.example.com/realms/nb?", "https://idp.example.com/realms/nb#"} {
|
||||
t.Run(issuer, func(t *testing.T) {
|
||||
idp := &IdentityProvider{
|
||||
Name: "test",
|
||||
Type: IdentityProviderTypeOIDC,
|
||||
Issuer: issuer,
|
||||
ClientID: "client-id",
|
||||
}
|
||||
assert.ErrorIs(t, idp.Validate(), ErrIdentityProviderIssuerInvalid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -285,6 +285,25 @@ func (u *User) EncryptSensitiveData(enc *crypt.FieldEncrypt) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func MaskEmail(email string) string {
|
||||
local, domain, found := strings.Cut(email, "@")
|
||||
if !found || local == "" || domain == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
// Runes, not bytes, so a non-ASCII local part is not cut mid-character.
|
||||
runes := []rune(local)
|
||||
|
||||
// Keeping the first two and the last needs a local part of at least four to
|
||||
// hide anything at all: at three or fewer those are the whole of it, and the
|
||||
// address would be recoverable in full from what is meant to conceal it.
|
||||
if len(runes) < 4 {
|
||||
return "****@" + domain
|
||||
}
|
||||
|
||||
return string(runes[:2]) + "****" + string(runes[len(runes)-1]) + "@" + domain
|
||||
}
|
||||
|
||||
// DecryptSensitiveData decrypts the user's sensitive fields (Email and Name) in place.
|
||||
func (u *User) DecryptSensitiveData(enc *crypt.FieldEncrypt) error {
|
||||
if enc == nil {
|
||||
|
||||
@@ -296,3 +296,144 @@ func TestUser_EncryptDecryptRoundTrip(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaskEmail(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
email string
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "ordinary address keeps the first two, the last, and the domain",
|
||||
email: "admin@example.com",
|
||||
expected: "ad****n@example.com",
|
||||
},
|
||||
{
|
||||
name: "four characters is the shortest local part that reveals anything",
|
||||
email: "abcd@example.com",
|
||||
expected: "ab****d@example.com",
|
||||
},
|
||||
{
|
||||
name: "three character local part is masked whole, since a lead and tail would be all of it",
|
||||
email: "abc@example.com",
|
||||
expected: "****@example.com",
|
||||
},
|
||||
{
|
||||
name: "two character local part is masked whole",
|
||||
email: "ab@example.com",
|
||||
expected: "****@example.com",
|
||||
},
|
||||
{
|
||||
name: "single character local part is masked whole",
|
||||
email: "a@b.co",
|
||||
expected: "****@b.co",
|
||||
},
|
||||
{
|
||||
name: "mask width does not report the length it stands in for",
|
||||
email: "a.very.long.local.part@example.com",
|
||||
expected: "a.****t@example.com",
|
||||
},
|
||||
{
|
||||
name: "a local part far longer than the mask is still reduced to three characters",
|
||||
email: "finance.department.notifications.owner.account@example.com",
|
||||
expected: "fi****t@example.com",
|
||||
},
|
||||
{
|
||||
name: "plus addressing is masked along with the rest of the local part",
|
||||
email: "admin+netbird@example.com",
|
||||
expected: "ad****d@example.com",
|
||||
},
|
||||
{
|
||||
name: "separators inside the local part are not treated specially",
|
||||
email: "first.last-name_x@example.com",
|
||||
expected: "fi****x@example.com",
|
||||
},
|
||||
{
|
||||
name: "case is preserved rather than normalised",
|
||||
email: "Admin@Example.COM",
|
||||
expected: "Ad****n@Example.COM",
|
||||
},
|
||||
{
|
||||
name: "subdomains stay intact",
|
||||
email: "owner@mail.corp.example.com",
|
||||
expected: "ow****r@mail.corp.example.com",
|
||||
},
|
||||
{
|
||||
name: "german umlauts count as single characters",
|
||||
email: "müller@example.de",
|
||||
expected: "mü****r@example.de",
|
||||
},
|
||||
{
|
||||
name: "cyrillic local part is cut on runes",
|
||||
email: "иванов@example.ru",
|
||||
expected: "ив****в@example.ru",
|
||||
},
|
||||
{
|
||||
name: "cjk local part of three runes is masked whole, counted in runes not bytes",
|
||||
email: "用户名@example.cn",
|
||||
expected: "****@example.cn",
|
||||
},
|
||||
{
|
||||
name: "cjk local part of four runes reveals the first two and the last",
|
||||
email: "用户名字@example.cn",
|
||||
expected: "用户****字@example.cn",
|
||||
},
|
||||
{
|
||||
name: "arabic local part is cut on runes",
|
||||
email: "مستخدم@example.sa",
|
||||
expected: "مس****م@example.sa",
|
||||
},
|
||||
{
|
||||
name: "two rune non-ascii local part is masked whole",
|
||||
email: "ää@example.de",
|
||||
expected: "****@example.de",
|
||||
},
|
||||
{
|
||||
name: "astral plane runes are not split into surrogates",
|
||||
email: "a🎉bc@example.com",
|
||||
expected: "a🎉****c@example.com",
|
||||
},
|
||||
{
|
||||
name: "a non-ascii domain is left alone",
|
||||
email: "admin@münchen.example",
|
||||
expected: "ad****n@münchen.example",
|
||||
},
|
||||
{
|
||||
name: "only the first separator splits, so a second stays in the domain",
|
||||
email: "a@b@example.com",
|
||||
expected: "****@b@example.com",
|
||||
},
|
||||
{
|
||||
name: "empty email has nothing to mask",
|
||||
email: "",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "value without a separator is not an address",
|
||||
email: "not-an-email",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "missing local part is not an address",
|
||||
email: "@example.com",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "missing domain is not an address",
|
||||
email: "admin@",
|
||||
expected: "",
|
||||
},
|
||||
{
|
||||
name: "a bare separator is not an address",
|
||||
email: "@",
|
||||
expected: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
assert.Equal(t, tc.expected, MaskEmail(tc.email))
|
||||
})
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
+88
-16
@@ -1177,28 +1177,35 @@ func (am *DefaultAccountManager) expireAndUpdatePeers(ctx context.Context, accou
|
||||
dnsDomain := am.networkMapController.GetDNSDomain(settings)
|
||||
|
||||
var peerIDs []string
|
||||
for _, peer := range peers {
|
||||
defer func() {
|
||||
if len(peerIDs) == 0 {
|
||||
return
|
||||
}
|
||||
// this will trigger peer disconnect from the management service
|
||||
log.Debugf("Expiring %d peers for account %s", len(peerIDs), accountID)
|
||||
am.networkMapController.DisconnectPeers(ctx, accountID, peerIDs)
|
||||
}()
|
||||
for _, candidate := range peers {
|
||||
// nolint:staticcheck
|
||||
ctx = context.WithValue(ctx, nbcontext.PeerIDKey, peer.Key)
|
||||
peerCtx := context.WithValue(ctx, nbcontext.PeerIDKey, candidate.Key)
|
||||
|
||||
if peer.UserID == "" {
|
||||
if candidate.UserID == "" {
|
||||
// we do not want to expire peers that are added via setup key
|
||||
continue
|
||||
}
|
||||
|
||||
if peer.Status.LoginExpired {
|
||||
peer, err := am.expirePeerIfStillDue(peerCtx, accountID, candidate.ID, settings, reason)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if peer == nil {
|
||||
continue
|
||||
}
|
||||
peerIDs = append(peerIDs, peer.ID)
|
||||
peer.MarkLoginExpired(true)
|
||||
|
||||
if err := am.Store.SavePeerStatus(ctx, accountID, peer.ID, *peer.Status); err != nil {
|
||||
return err
|
||||
}
|
||||
meta := peer.EventMeta(dnsDomain)
|
||||
meta["reason"] = string(reason)
|
||||
am.StoreEvent(
|
||||
ctx,
|
||||
peerCtx,
|
||||
peer.UserID, peer.ID, accountID,
|
||||
activity.PeerLoginExpired, meta,
|
||||
)
|
||||
@@ -1215,15 +1222,53 @@ func (am *DefaultAccountManager) expireAndUpdatePeers(ctx context.Context, accou
|
||||
if err != nil {
|
||||
return fmt.Errorf("notify network map controller of peer update: %w", err)
|
||||
}
|
||||
|
||||
if len(peerIDs) != 0 {
|
||||
// this will trigger peer disconnect from the management service
|
||||
log.Debugf("Expiring %d peers for account %s", len(peerIDs), accountID)
|
||||
am.networkMapController.DisconnectPeers(ctx, accountID, peerIDs)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// expirePeerIfStillDue flags the peer as login-expired and returns its fresh copy, or nil
|
||||
// when it no longer qualifies. The candidate list is read without a lock, so a login that
|
||||
// landed in between would otherwise be overwritten with a stale expired status.
|
||||
func (am *DefaultAccountManager) expirePeerIfStillDue(ctx context.Context, accountID, peerID string, settings *types.Settings, reason peerExpirationReason) (*nbpeer.Peer, error) {
|
||||
var expired *nbpeer.Peer
|
||||
err := am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
peer, err := transaction.GetPeerByID(ctx, store.LockingStrengthUpdate, accountID, peerID)
|
||||
if err != nil {
|
||||
if s, ok := status.FromError(err); ok && s.Type() == status.NotFound {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
if peer.Status.LoginExpired || !peerExpirationDue(peer, settings, reason) {
|
||||
return nil
|
||||
}
|
||||
peer.MarkLoginExpired(true)
|
||||
if err := transaction.SavePeerStatus(ctx, accountID, peer.ID, *peer.Status); err != nil {
|
||||
return err
|
||||
}
|
||||
expired = peer
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return expired, nil
|
||||
}
|
||||
|
||||
// peerExpirationDue re-evaluates a time-based expiry against the peer's current state.
|
||||
// Administrative reasons expire the peer unconditionally.
|
||||
func peerExpirationDue(peer *nbpeer.Peer, settings *types.Settings, reason peerExpirationReason) bool {
|
||||
switch reason {
|
||||
case peerExpirationSessionExpired:
|
||||
expired, _ := peer.LoginExpired(settings.PeerLoginExpiration)
|
||||
return settings.PeerLoginExpirationEnabled && expired
|
||||
case peerExpirationInactivity:
|
||||
expired, _ := peer.SessionExpired(settings.PeerInactivityExpiration)
|
||||
return settings.PeerInactivityExpirationEnabled && expired
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
func (am *DefaultAccountManager) deleteUserFromIDP(ctx context.Context, targetUserID, accountID string) error {
|
||||
if am.userDeleteFromIDPEnabled {
|
||||
log.WithContext(ctx).Debugf("user %s deleted from IdP", targetUserID)
|
||||
@@ -1403,6 +1448,25 @@ func (am *DefaultAccountManager) deleteRegularUser(ctx context.Context, accountI
|
||||
return updateAccountPeers, nil
|
||||
}
|
||||
|
||||
// pendingApprovalError refuses a user awaiting approval, naming the owner who
|
||||
// can approve them when their address resolves. Failing to resolve one is not a
|
||||
// reason to withhold the refusal, so the lookup is best effort.
|
||||
func (am *DefaultAccountManager) pendingApprovalError(ctx context.Context, accountID string) error {
|
||||
owner, err := am.GetOwnerInfo(ctx, accountID)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Debugf("pending approval refusal: owner of account %s did not resolve: %v", accountID, err)
|
||||
return status.NewUserPendingApprovalError()
|
||||
}
|
||||
|
||||
masked := types.MaskEmail(owner.Email)
|
||||
if masked == "" {
|
||||
log.WithContext(ctx).Debugf("pending approval refusal: no address found for the owner of account %s", accountID)
|
||||
return status.NewUserPendingApprovalError()
|
||||
}
|
||||
|
||||
return status.NewUserPendingApprovalByOwnerError(masked)
|
||||
}
|
||||
|
||||
// GetOwnerInfo retrieves the owner information for a given account ID.
|
||||
func (am *DefaultAccountManager) GetOwnerInfo(ctx context.Context, accountID string) (*types.UserInfo, error) {
|
||||
owner, err := am.Store.GetAccountOwner(ctx, store.LockingStrengthNone, accountID)
|
||||
@@ -1460,6 +1524,14 @@ func (am *DefaultAccountManager) GetCurrentUserInfo(ctx context.Context, userAut
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// A user pending approval is blocked too, and the dashboard needs to tell
|
||||
// the two apart: one is a dead end, the other resolves by itself once the
|
||||
// owner acts. Naming that owner needs the address the IdP holds, which is
|
||||
// why this is answered here rather than in the permission gate.
|
||||
if user.IsBlocked() && user.PendingApproval {
|
||||
return nil, am.pendingApprovalError(ctx, user.AccountID)
|
||||
}
|
||||
|
||||
if user.IsBlocked() {
|
||||
return nil, status.NewUserBlockedError()
|
||||
}
|
||||
|
||||
@@ -1779,6 +1779,42 @@ func TestDefaultAccountManager_GetCurrentUserInfo(t *testing.T) {
|
||||
}
|
||||
require.NoError(t, store.SaveAccount(context.Background(), account2))
|
||||
|
||||
account3 := newAccountWithId(context.Background(), "account3", "account3Owner", "", "owner@example.com", "", false)
|
||||
account3.Users["pending-user"] = &types.User{
|
||||
Id: "pending-user",
|
||||
AccountID: account3.Id,
|
||||
Role: types.UserRoleUser,
|
||||
Blocked: true,
|
||||
PendingApproval: true,
|
||||
}
|
||||
require.NoError(t, store.SaveAccount(context.Background(), account3))
|
||||
|
||||
// The owner has no address to name, so the refusal falls back to the generic one.
|
||||
account4 := newAccountWithId(context.Background(), "account4", "account4Owner", "", "", "", false)
|
||||
account4.Users["pending-user-without-owner-email"] = &types.User{
|
||||
Id: "pending-user-without-owner-email",
|
||||
AccountID: account4.Id,
|
||||
Role: types.UserRoleUser,
|
||||
Blocked: true,
|
||||
PendingApproval: true,
|
||||
}
|
||||
require.NoError(t, store.SaveAccount(context.Background(), account4))
|
||||
|
||||
// No user holds the owner role, so the owner lookup itself fails.
|
||||
account5 := newAccountWithId(context.Background(), "account5", "account5Admin", "", "", "", false)
|
||||
account5.Users["account5Admin"].Role = types.UserRoleAdmin
|
||||
account5.Users["pending-user-without-owner"] = &types.User{
|
||||
Id: "pending-user-without-owner",
|
||||
AccountID: account5.Id,
|
||||
Role: types.UserRoleUser,
|
||||
Blocked: true,
|
||||
PendingApproval: true,
|
||||
}
|
||||
require.NoError(t, store.SaveAccount(context.Background(), account5))
|
||||
|
||||
account6 := newAccountWithId(context.Background(), "account6", "account6Owner", "", "stranger@example.com", "", false)
|
||||
require.NoError(t, store.SaveAccount(context.Background(), account6))
|
||||
|
||||
permissionsManager := permissions.NewManager(store)
|
||||
am := DefaultAccountManager{
|
||||
Store: store,
|
||||
@@ -1812,6 +1848,34 @@ func TestDefaultAccountManager_GetCurrentUserInfo(t *testing.T) {
|
||||
userAuth: auth.UserAuth{AccountId: account1.Id, UserId: "service-user"},
|
||||
expectedErr: status.NewPermissionDeniedError(),
|
||||
},
|
||||
{
|
||||
name: "pending approval names the owner",
|
||||
userAuth: auth.UserAuth{AccountId: account3.Id, UserId: "pending-user"},
|
||||
expectedErr: status.NewUserPendingApprovalByOwnerError("ow****r@example.com"),
|
||||
},
|
||||
{
|
||||
name: "pending approval without an owner address",
|
||||
userAuth: auth.UserAuth{AccountId: account4.Id, UserId: "pending-user-without-owner-email"},
|
||||
expectedErr: status.NewUserPendingApprovalError(),
|
||||
},
|
||||
{
|
||||
name: "pending approval without an owner",
|
||||
userAuth: auth.UserAuth{AccountId: account5.Id, UserId: "pending-user-without-owner"},
|
||||
expectedErr: status.NewUserPendingApprovalError(),
|
||||
},
|
||||
{
|
||||
// The account claim points at an account the caller is not in. The
|
||||
// owner named has to be the one of the account holding the caller's
|
||||
// own record, never the one the claim asks for.
|
||||
name: "pending approval ignores a mismatched account claim",
|
||||
userAuth: auth.UserAuth{AccountId: account6.Id, UserId: "pending-user"},
|
||||
expectedErr: status.NewUserPendingApprovalByOwnerError("ow****r@example.com"),
|
||||
},
|
||||
{
|
||||
name: "blocked user answers before the account claim is validated",
|
||||
userAuth: auth.UserAuth{AccountId: account6.Id, UserId: "blocked-user"},
|
||||
expectedErr: status.NewUserBlockedError(),
|
||||
},
|
||||
{
|
||||
name: "owner user",
|
||||
userAuth: auth.UserAuth{AccountId: account1.Id, UserId: "account1Owner"},
|
||||
|
||||
Reference in New Issue
Block a user