Merge main into poc/certificate-posture

This commit is contained in:
Viktor Liu
2026-10-05 19:09:43 +02:00
390 changed files with 28443 additions and 14545 deletions
+23 -4
View File
@@ -10,6 +10,7 @@ import (
"io"
"strings"
"text/tabwriter"
"unicode"
"github.com/spf13/cobra"
@@ -68,8 +69,8 @@ func runDisconnectAll(ctx context.Context, s store.Store, out io.Writer, in io.R
toDisconnect := 0
w := tabwriter.NewWriter(out, 0, 0, 2, ' ', 0)
_, _ = fmt.Fprintln(w, "ID\tCLUSTER\tIP\tACCOUNT\tSTATUS\tLAST SEEN")
_, _ = fmt.Fprintln(w, "--\t-------\t--\t-------\t------\t---------")
_, _ = fmt.Fprintln(w, "ID\tCLUSTER\tIP\tVERSION\tACCOUNT\tSTATUS\tLAST SEEN")
_, _ = fmt.Fprintln(w, "--\t-------\t--\t-------\t-------\t------\t---------")
for _, p := range proxies {
if p.Status != rpproxy.StatusDisconnected {
@@ -80,11 +81,16 @@ func runDisconnectAll(ctx context.Context, s store.Store, out io.Writer, in io.R
if p.AccountID != nil {
account = *p.AccountID
}
version := "-"
if p.Version != "" {
version = sanitizeReportedValue(p.Version)
}
_, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\t%s\n",
p.ID,
_, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\t%s\t%s\n",
sanitizeReportedValue(p.ID),
p.ClusterAddress,
p.IPAddress,
version,
account,
p.Status,
p.LastSeen.Format("2006-01-02 15:04:05"),
@@ -139,3 +145,16 @@ func confirmDisconnectAll(out io.Writer, in io.Reader) (bool, error) {
return strings.EqualFold(strings.TrimSpace(scanner.Text()), disconnectAllConfirmation), nil
}
// sanitizeReportedValue replaces non-printable characters in a value the proxy
// reports about itself. Both the id and the version arrive unvalidated over
// gRPC, so a tab would forge a column, a carriage return or ANSI escape would
// redraw the operator's terminal, and U+202E would reverse the rest of the line.
func sanitizeReportedValue(s string) string {
return strings.Map(func(r rune) rune {
if unicode.IsPrint(r) {
return r
}
return '\uFFFD'
}, s)
}
+39
View File
@@ -35,6 +35,7 @@ func seedProxies(t *testing.T, ctx context.Context, s store.Store) {
SessionID: "session-1",
ClusterAddress: "cluster-a.example.com",
IPAddress: "10.0.0.1",
Version: "0.60.0",
LastSeen: time.Now(),
Status: rpproxy.StatusConnected,
},
@@ -89,6 +90,7 @@ func TestRunDisconnectAllWithConfirmation(t *testing.T) {
require.Contains(t, output, "proxy-2")
require.Contains(t, output, "proxy-3")
require.Contains(t, output, "cluster-a.example.com")
require.Contains(t, output, "0.60.0")
require.Contains(t, output, "account-1")
require.Contains(t, output, "Type \"disconnect all proxies\" to continue")
require.Contains(t, output, "Force-marked 2 of 3 reverse proxy instance(s) as disconnected.")
@@ -178,3 +180,40 @@ func TestRunDisconnectAllEmpty(t *testing.T) {
require.NoError(t, runDisconnectAll(ctx, s, &out, strings.NewReader(""), false, false))
require.Contains(t, out.String(), "No reverse proxy instances found.")
}
func TestRunDisconnectAllEscapesProxyReportedFields(t *testing.T) {
ctx := context.Background()
s := newTestStore(t)
// A proxy reports its own id and version on connect, so both reach this
// listing unvalidated. Carriage returns, tabs and ANSI escapes would let
// a malicious proxy redraw the table or forge a row on the operator's
// terminal; U+202E would reverse the rendering of the rest of the line.
require.NoError(t, s.SaveProxy(ctx, &rpproxy.Proxy{
ID: "proxy-\r\x1b[2Kevil",
SessionID: "session-1",
ClusterAddress: "cluster-a.example.com",
IPAddress: "10.0.0.1",
Version: "0.60.0\tfake\rcolumn\u202e",
LastSeen: time.Now(),
Status: rpproxy.StatusConnected,
}))
var out bytes.Buffer
require.NoError(t, runDisconnectAll(ctx, s, &out, strings.NewReader(disconnectAllConfirmation+"\n"), true, false))
output := out.String()
for _, forbidden := range []string{"\r", "\x1b", "\u202e"} {
require.NotContains(t, output, forbidden, "listing must not carry proxy-reported control characters")
}
// The table has one data row; a smuggled tab would add a phantom column.
var dataRow string
for _, line := range strings.Split(output, "\n") {
if strings.Contains(line, "evil") {
dataRow = line
}
}
require.NotEmpty(t, dataRow, "listing should still show the proxy row")
require.NotContains(t, dataRow, "\t", "tabwriter output should not carry a smuggled column separator")
require.Contains(t, dataRow, "0.60.0", "the printable part of the version should survive")
}
@@ -44,7 +44,7 @@ func (r *repository) GetAccountNetwork(ctx context.Context, accountID string) (*
}
func (r *repository) GetAccountPeers(ctx context.Context, accountID string) ([]*peer.Peer, error) {
return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
}
func (r *repository) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) {
@@ -245,7 +245,7 @@ func computeMode(t *testing.T, ctx context.Context, mode Mode, nmData *networkma
peerGroups := maps.Keys(nmData.GetPeerGroups(peerID))
resp := mgmtgrpc.ToComponentSyncResponse(ctx, nil, nil, nil, peer, nil, nil, components, nil,
dnsDomain, nil, nmData.AccountSettings, nil, peerGroups, dnsFwdPort)
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain)
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain, false)
require.NoError(t, err, "expand envelope")
return res.NetworkMap
default:
@@ -0,0 +1,76 @@
package agentnetwork
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/store"
nbtypes "github.com/netbirdio/netbird/management/server/types"
)
// TestCleanupAccessLogs_RealStore_DeletedAccount covers a deleted account's access logs.
// The sweep is driven by settings rows, which go with the account, so without a fallback
// those logs would never expire. They get the default retention instead. A live account
// can delete its own settings row, so "no settings" must not be mistaken for "deleted":
// that account's logs are left alone, as are those of an account that keeps logs forever.
func TestCleanupAccessLogs_RealStore_DeletedAccount(t *testing.T) {
ctx := context.Background()
s, cleanup, err := store.NewTestStoreFromSQL(ctx, "", t.TempDir())
require.NoError(t, err, "real sqlite test store must come up")
defer cleanup()
const (
deletedAccountID = "acc-deleted"
keepAccountID = "acc-keep-forever"
noSettingsAccountID = "acc-live-no-settings"
)
old := time.Now().UTC().AddDate(0, 0, -(types.DefaultAccessLogRetentionDays + 10))
recent := time.Now().UTC().AddDate(0, 0, -1)
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: keepAccountID}))
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: noSettingsAccountID}))
keepSettings := types.DefaultSettings(keepAccountID)
keepSettings.Domain = "keep.gw.example.com"
keepSettings.AccessLogRetentionDays = 0
require.NoError(t, s.SaveAgentNetworkSettings(ctx, keepSettings))
mkLog := func(id, accountID string, ts time.Time) {
t.Helper()
entry := &types.AgentNetworkAccessLog{
ID: id, AccountID: accountID, ServiceID: "svc", Timestamp: ts, StatusCode: 200, Model: "gpt-4o",
}
groups := []types.AgentNetworkAccessLogGroup{{LogID: id, GroupID: "grp-eng", AccountID: accountID}}
require.NoError(t, s.CreateAgentNetworkAccessLog(ctx, entry, groups))
}
mkLog("deleted-old", deletedAccountID, old)
mkLog("deleted-recent", deletedAccountID, recent)
mkLog("keep-old", keepAccountID, old)
mkLog("no-settings-old", noSettingsAccountID, old)
m := &managerImpl{store: s}
m.cleanupAccessLogsOnce(ctx)
logIDs := func(accountID string) []string {
t.Helper()
logs, _, err := s.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID,
types.AgentNetworkAccessLogFilter{Page: 1, PageSize: 50})
require.NoError(t, err)
ids := make([]string, 0, len(logs))
for _, l := range logs {
ids = append(ids, l.ID)
}
return ids
}
assert.Equal(t, []string{"deleted-recent"}, logIDs(deletedAccountID),
"a deleted account should have logs past the default retention swept")
assert.Equal(t, []string{"keep-old"}, logIDs(keepAccountID),
"an account with retention disabled should keep its old logs")
assert.Equal(t, []string{"no-settings-old"}, logIDs(noSettingsAccountID),
"a live account without a settings row should keep its old logs")
}
@@ -80,6 +80,9 @@ type Manager interface {
ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error)
GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error)
StartAccessLogCleanup(ctx context.Context, cleanupIntervalHours int)
// RemoveAccountGateway drops the account's gateway mappings from the
// proxies. It runs as an account deletion hook.
RemoveAccountGateway(ctx context.Context, accountID string) error
RecordConsumption(ctx context.Context, accountID string, kind types.ConsumptionDimension, dimID string, windowSeconds, tokensIn, tokensOut int64, costUSD float64) error
RecordAccountBudgetUsage(ctx context.Context, accountID, userID string, groupIDs []string, tokensIn, tokensOut int64, costUSD float64) error
RecordUsage(ctx context.Context, in RecordUsageInput) error
@@ -1350,8 +1353,8 @@ func (m *managerImpl) scopeFilterToCaller(ctx context.Context, accountID, userID
// StartAccessLogCleanup launches a background sweep that periodically deletes
// each account's agent-network access-log rows older than that account's
// AccessLogRetentionDays. Usage records are never swept. A non-positive
// interval defaults to 24h.
// AccessLogRetentionDays, and the consumption counters of deleted accounts.
// Usage records are never swept. A non-positive interval defaults to 24h.
func (m *managerImpl) StartAccessLogCleanup(ctx context.Context, cleanupIntervalHours int) {
if cleanupIntervalHours <= 0 {
cleanupIntervalHours = 24
@@ -1362,21 +1365,40 @@ func (m *managerImpl) StartAccessLogCleanup(ctx context.Context, cleanupInterval
ticker := time.NewTicker(interval)
defer ticker.Stop()
m.cleanupAccessLogsOnce(ctx) // run once on startup
m.cleanupOnce(ctx) // run once on startup
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
m.cleanupAccessLogsOnce(ctx)
m.cleanupOnce(ctx)
}
}
}()
}
func (m *managerImpl) cleanupOnce(ctx context.Context) {
m.cleanupAccessLogsOnce(ctx)
m.cleanupDeletedAccountConsumption(ctx)
}
// cleanupDeletedAccountConsumption deletes the consumption counters of accounts
// that no longer exist. Best-effort: a failure is logged and retried next sweep.
func (m *managerImpl) cleanupDeletedAccountConsumption(ctx context.Context) {
deleted, err := m.store.DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx)
if err != nil {
log.WithContext(ctx).Warnf("agent-network consumption cleanup: %v", err)
return
}
if deleted > 0 {
log.WithContext(ctx).Infof("agent-network consumption cleanup: deleted %d counters of deleted accounts", deleted)
}
}
// cleanupAccessLogsOnce sweeps every account's expired access-log rows against
// its configured retention. Best-effort: a per-account failure is logged and
// the sweep continues.
// its configured retention. Deleted accounts, whose settings rows went with
// them, get the default retention. Best-effort: a per-account failure is
// logged and the sweep continues.
func (m *managerImpl) cleanupAccessLogsOnce(ctx context.Context) {
settings, err := m.store.GetAllAgentNetworkSettings(ctx, store.LockingStrengthNone)
if err != nil {
@@ -1384,18 +1406,31 @@ func (m *managerImpl) cleanupAccessLogsOnce(ctx context.Context) {
return
}
for _, s := range settings {
if s.AccessLogRetentionDays <= 0 {
continue // keep indefinitely
}
cutoff := time.Now().UTC().AddDate(0, 0, -s.AccessLogRetentionDays)
deleted, err := m.store.DeleteOldAgentNetworkAccessLogs(ctx, s.AccountID, cutoff)
if err != nil {
log.WithContext(ctx).Warnf("agent-network access-log cleanup for account %s: %v", s.AccountID, err)
continue
}
if deleted > 0 {
log.WithContext(ctx).Infof("agent-network access-log cleanup: deleted %d rows for account %s (retention %d days)", deleted, s.AccountID, s.AccessLogRetentionDays)
}
m.cleanupAccountAccessLogs(ctx, s.AccountID, s.AccessLogRetentionDays)
}
deleted, err := m.store.GetDeletedAccountIDsWithAgentNetworkAccessLogs(ctx)
if err != nil {
log.WithContext(ctx).Errorf("agent-network access-log cleanup: list deleted accounts: %v", err)
return
}
for _, accountID := range deleted {
m.cleanupAccountAccessLogs(ctx, accountID, types.DefaultAccessLogRetentionDays)
}
}
func (m *managerImpl) cleanupAccountAccessLogs(ctx context.Context, accountID string, retentionDays int) {
if retentionDays <= 0 {
return // keep indefinitely
}
cutoff := time.Now().UTC().AddDate(0, 0, -retentionDays)
deleted, err := m.store.DeleteOldAgentNetworkAccessLogs(ctx, accountID, cutoff)
if err != nil {
log.WithContext(ctx).Warnf("agent-network access-log cleanup for account %s: %v", accountID, err)
return
}
if deleted > 0 {
log.WithContext(ctx).Infof("agent-network access-log cleanup: deleted %d rows for account %s (retention %d days)", deleted, accountID, retentionDays)
}
}
@@ -1545,6 +1580,8 @@ func (*mockManager) GetUsageOverview(_ context.Context, _, _ string, _ types.Age
func (*mockManager) StartAccessLogCleanup(_ context.Context, _ int) {}
func (*mockManager) RemoveAccountGateway(_ context.Context, _ string) error { return nil }
func (*mockManager) RecordConsumption(_ context.Context, _ string, _ types.ConsumptionDimension, _ string, _, _, _ int64, _ float64) error {
return nil
}
@@ -2,8 +2,10 @@ package agentnetwork
import (
"context"
"fmt"
log "github.com/sirupsen/logrus"
goproto "google.golang.org/protobuf/proto"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/types"
@@ -81,18 +83,66 @@ func (m *managerImpl) reconcile(ctx context.Context, accountID string) {
}
m.reconcileMu.Unlock()
for _, entry := range creates {
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
m.sendMappings(ctx, accountID, creates, proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED)
m.sendMappings(ctx, accountID, updates, proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED)
m.sendMappings(ctx, accountID, deletes, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED)
}
// sendMappings sends each entry as updateType. It sends a copy: the entries'
// mappings are shared with reconcileCache, which another reconcile or
// RemoveAccountGateway may be reading, so they are never written.
func (m *managerImpl) sendMappings(ctx context.Context, accountID string, entries []syntheticMapping, updateType proto.ProxyMappingUpdateType) {
for _, entry := range entries {
update := goproto.Clone(entry.mapping).(*proto.ProxyMapping)
update.Type = updateType
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, update, entry.cluster)
}
for _, entry := range updates {
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
}
// RemoveAccountGateway tells the proxies to drop every mapping of the account's
// gateway, so a deleted account's proxy config, provider API keys included, does
// not linger in proxy memory until the next resync. It is an account deletion
// hook: it runs before the account's data is removed, the last point at which
// the mappings can be synthesised from the store. The cache alone would miss
// them, since it is per instance and empty after a restart. If the deletion
// then fails, the gateway stays down until the account's next change reconciles
// it back.
func (m *managerImpl) RemoveAccountGateway(ctx context.Context, accountID string) error {
if m.proxyController == nil {
return nil
}
for _, entry := range deletes {
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
services, err := SynthesizeServices(ctx, m.store, accountID)
if err != nil {
return fmt.Errorf("synthesise agent network services: %w", err)
}
oidcCfg := m.proxyController.GetOIDCValidationConfig()
removed := make(map[string]syntheticMapping, len(services))
for _, svc := range services {
if svc == nil || svc.ID == "" {
continue
}
removed[svc.ID] = syntheticMapping{
mapping: svc.ToProtoMapping(rpservice.Delete, "", oidcCfg),
cluster: svc.ProxyCluster,
}
}
m.reconcileMu.Lock()
for id, entry := range m.reconcileCache[accountID] {
if _, ok := removed[id]; !ok {
removed[id] = entry
}
}
delete(m.reconcileCache, accountID)
m.reconcileMu.Unlock()
entries := make([]syntheticMapping, 0, len(removed))
for _, entry := range removed {
entries = append(entries, entry)
}
m.sendMappings(ctx, accountID, entries, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED)
return nil
}
// diffMappings classifies the previous→current transition for a single
@@ -2,6 +2,8 @@ package agentnetwork
import (
"context"
"sync"
"sync/atomic"
"testing"
"go.uber.org/mock/gomock"
@@ -12,6 +14,7 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/management/status"
)
func newReconcileMgr(t *testing.T, ctrl *gomock.Controller) (*managerImpl, *store.MockStore, *proxy.MockController) {
@@ -287,3 +290,154 @@ func TestDiffMappings_RemovedServiceIsDeletedOnItsOwnCluster(t *testing.T) {
assert.Equal(t, "brave-otter.gateway.example.com", deletes[0].cluster)
}
}
// TestRemoveAccountGateway_EmitsRemovedFromStore — account deletion runs on an
// instance that may never have reconciled the account, so its cache is empty.
// The mappings are synthesised from the store, still intact before the delete,
// and each is sent as REMOVED to the cluster that serves it.
func TestRemoveAccountGateway_EmitsRemovedFromStore(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
provider := newReconcileTestProvider()
policy := newReconcileTestPolicy(provider.ID, "grp-eng")
expectReconcileSynthInputs(mockStore, ctx, []*types.Provider{provider}, []*types.Policy{policy}, []*types.Guardrail{})
mockProxy.EXPECT().GetOIDCValidationConfig().Return(proxy.OIDCValidationConfig{})
var sent []*proto.ProxyMapping
mockProxy.EXPECT().
SendServiceUpdateToCluster(ctx, "acct-1", gomock.Any(), "eu.proxy.netbird.io").
Do(func(_ context.Context, _ string, m *proto.ProxyMapping, _ string) {
sent = append(sent, m)
})
require.NoError(t, mgr.RemoveAccountGateway(ctx, "acct-1"))
require.Len(t, sent, 1, "the account's one gateway mapping must be removed")
assert.Equal(t, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED, sent[0].Type, "the update must be a removal")
assert.Equal(t, "agent-net-svc-acct-1", sent[0].Id, "the removal must name the account's gateway service")
}
// TestRemoveAccountGateway_AlsoRemovesCachedMappings — a mapping this instance
// last sent but the store no longer synthesises (here, one on another cluster)
// is removed too, and the account's cache entry is cleared.
func TestRemoveAccountGateway_AlsoRemovesCachedMappings(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
mgr.reconcileCache["acct-1"] = map[string]syntheticMapping{
"stale-svc": {mapping: &proto.ProxyMapping{Id: "stale-svc"}, cluster: "us.proxy.netbird.io"},
}
// Settings but no providers: the store synthesises nothing.
mockStore.EXPECT().
GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").
Return(newReconcileTestSettings(), nil)
mockStore.EXPECT().
GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "acct-1").
Return([]*types.Provider{}, nil)
mockProxy.EXPECT().GetOIDCValidationConfig().Return(proxy.OIDCValidationConfig{})
var sent []*proto.ProxyMapping
mockProxy.EXPECT().
SendServiceUpdateToCluster(ctx, "acct-1", gomock.Any(), "us.proxy.netbird.io").
Do(func(_ context.Context, _ string, m *proto.ProxyMapping, _ string) {
sent = append(sent, m)
})
require.NoError(t, mgr.RemoveAccountGateway(ctx, "acct-1"))
require.Len(t, sent, 1, "the cached mapping must be removed from its own cluster")
assert.Equal(t, "stale-svc", sent[0].Id)
assert.Equal(t, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED, sent[0].Type)
mgr.reconcileMu.Lock()
_, present := mgr.reconcileCache["acct-1"]
mgr.reconcileMu.Unlock()
assert.False(t, present, "the deleted account's cache entry must be cleared")
}
// TestRemoveAccountGateway_SynthFailureAbortsDeletion — if the mappings cannot
// be read, nothing is sent and the error is returned, which as an account
// deletion hook keeps the account rather than leaving its gateway running.
func TestRemoveAccountGateway_SynthFailureAbortsDeletion(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mgr, mockStore, _ := newReconcileMgr(t, ctrl)
mockStore.EXPECT().
GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").
Return(nil, status.Errorf(status.Internal, "store unavailable"))
assert.Error(t, mgr.RemoveAccountGateway(ctx, "acct-1"), "a failed synthesis must fail the hook")
}
func TestRemoveAccountGateway_NilProxyController_NoOp(t *testing.T) {
mgr := &managerImpl{reconcileCache: make(map[string]map[string]syntheticMapping)}
// Must not panic and must not query the store.
assert.NoError(t, mgr.RemoveAccountGateway(context.Background(), "acct-1"))
}
// TestReconcile_ConcurrentWithGatewayChanges — while an account's gateway
// flaps (its policy is removed and re-added between reads), concurrent
// reconciles and RemoveAccountGateway share the cached mappings: one caches a
// mapping and sends it, another finds it gone and sends its removal. Run under
// -race: neither path may write a cached mapping, only copies of it.
func TestReconcile_ConcurrentWithGatewayChanges(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
// gomock serialises every call on the controller's mutex, which would give
// the race detector the ordering the code under test lacks. The sends go
// through a fake that takes no lock.
mgr.proxyController = unsyncedSender{MockController: mockProxy}
provider := newReconcileTestProvider()
policy := newReconcileTestPolicy(provider.ID, "grp-eng")
var reads atomic.Int64
mockStore.EXPECT().GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").Return(newReconcileTestSettings(), nil).AnyTimes()
mockStore.EXPECT().GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "acct-1").Return([]*types.Provider{provider}, nil).AnyTimes()
mockStore.EXPECT().GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, "acct-1").
DoAndReturn(func(context.Context, store.LockingStrength, string) ([]*types.Policy, error) {
if reads.Add(1)%2 == 0 {
return []*types.Policy{}, nil
}
return []*types.Policy{policy}, nil
}).AnyTimes()
mockStore.EXPECT().GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, "acct-1").Return([]*types.Guardrail{}, nil).AnyTimes()
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(1)
go func(remove bool) {
defer wg.Done()
for j := 0; j < 50; j++ {
if remove && j%10 == 0 {
_ = mgr.RemoveAccountGateway(ctx, "acct-1")
continue
}
mgr.reconcile(ctx, "acct-1")
}
}(i == 0)
}
wg.Wait()
}
// unsyncedSender answers the calls reconcile makes on every pass without any
// locking, so concurrent callers are not ordered by the fake itself.
type unsyncedSender struct {
*proxy.MockController
}
func (unsyncedSender) GetOIDCValidationConfig() proxy.OIDCValidationConfig {
return proxy.OIDCValidationConfig{}
}
func (unsyncedSender) SendServiceUpdateToCluster(context.Context, string, *proto.ProxyMapping, string) {}
@@ -97,7 +97,7 @@ func (m *managerImpl) GetAllPeers(ctx context.Context, accountID, userID string)
return m.store.GetUserPeers(ctx, store.LockingStrengthNone, accountID, userID)
}
return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
}
func (m *managerImpl) GetPeerAccountID(ctx context.Context, peerID string) (string, error) {
@@ -9,6 +9,7 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/management/server/geolocation"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
@@ -18,14 +19,16 @@ import (
)
type managerImpl struct {
repo accesslogs.Repository
store store.Store
permissionsManager permissions.Manager
geo geolocation.Geolocation
cleanupCancel context.CancelFunc
}
func NewManager(store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager {
func NewManager(repo accesslogs.Repository, store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager {
return &managerImpl{
repo: repo,
store: store,
permissionsManager: permissionsManager,
geo: geo,
@@ -54,7 +57,7 @@ func (m *managerImpl) SaveAccessLog(ctx context.Context, logEntry *accesslogs.Ac
}
}
if err := m.store.CreateAccessLog(ctx, logEntry); err != nil {
if err := m.repo.Create(ctx, logEntry); err != nil {
log.WithContext(ctx).WithFields(log.Fields{
"service_id": logEntry.ServiceID,
"method": logEntry.Method,
@@ -82,7 +85,7 @@ func (m *managerImpl) GetAllAccessLogs(ctx context.Context, accountID, userID st
log.WithContext(ctx).Warnf("failed to resolve user filters: %v", err)
}
logs, totalCount, err := m.store.GetAccountAccessLogs(ctx, store.LockingStrengthNone, accountID, *filter)
logs, totalCount, err := m.repo.ListByAccount(ctx, db.LockingStrengthNone, accountID, *filter)
if err != nil {
return nil, 0, err
}
@@ -98,7 +101,7 @@ func (m *managerImpl) CleanupOldAccessLogs(ctx context.Context, retentionDays in
}
cutoffTime := time.Now().AddDate(0, 0, -retentionDays)
deletedCount, err := m.store.DeleteOldAccessLogs(ctx, cutoffTime)
deletedCount, err := m.repo.DeleteOlderThan(ctx, cutoffTime)
if err != nil {
log.WithContext(ctx).Errorf("failed to cleanup old access logs: %v", err)
return 0, err
@@ -5,27 +5,27 @@ import (
"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/server/store"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
)
func TestCleanupOldAccessLogs(t *testing.T) {
tests := []struct {
name string
retentionDays int
setupMock func(*store.MockStore)
setupMock func(*accesslogs.MockRepository)
expectedCount int64
expectedError bool
}{
{
name: "cleanup logs older than retention period",
retentionDays: 30,
setupMock: func(mockStore *store.MockStore) {
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
setupMock: func(mockRepo *accesslogs.MockRepository) {
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) {
expectedCutoff := time.Now().AddDate(0, 0, -30)
timeDiff := olderThan.Sub(expectedCutoff)
@@ -41,9 +41,9 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "no logs to cleanup",
retentionDays: 30,
setupMock: func(mockStore *store.MockStore) {
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
setupMock: func(mockRepo *accesslogs.MockRepository) {
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(0), nil)
},
expectedCount: 0,
@@ -52,8 +52,8 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "zero retention days skips cleanup",
retentionDays: 0,
setupMock: func(mockStore *store.MockStore) {
// No expectations - DeleteOldAccessLogs should not be called
setupMock: func(mockRepo *accesslogs.MockRepository) {
// No expectations - DeleteOlderThan should not be called
},
expectedCount: 0,
expectedError: false,
@@ -61,8 +61,8 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "negative retention days skips cleanup",
retentionDays: -10,
setupMock: func(mockStore *store.MockStore) {
// No expectations - DeleteOldAccessLogs should not be called
setupMock: func(mockRepo *accesslogs.MockRepository) {
// No expectations - DeleteOlderThan should not be called
},
expectedCount: 0,
expectedError: false,
@@ -74,11 +74,11 @@ func TestCleanupOldAccessLogs(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
tt.setupMock(mockStore)
mockRepo := accesslogs.NewMockRepository(ctrl)
tt.setupMock(mockRepo)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -98,10 +98,10 @@ func TestCleanupWithExactBoundary(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) {
expectedCutoff := time.Now().AddDate(0, 0, -30)
timeDiff := olderThan.Sub(expectedCutoff)
@@ -110,7 +110,7 @@ func TestCleanupWithExactBoundary(t *testing.T) {
})
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -125,11 +125,11 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
// No expectations - cleanup should not run
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -139,22 +139,22 @@ func TestStartPeriodicCleanup(t *testing.T) {
time.Sleep(100 * time.Millisecond)
// If DeleteOldAccessLogs was called, the test will fail due to unexpected call
// If DeleteOlderThan was called, the test will fail due to unexpected call
})
t.Run("periodic cleanup runs immediately on start", func(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(2), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -171,15 +171,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(1), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -198,15 +198,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(0), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -223,15 +223,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(3), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -249,15 +249,15 @@ func TestStopPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(1), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -0,0 +1,135 @@
package manager
import (
"context"
"strings"
"time"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/shared/management/status"
)
type sqlRepository struct {
conn *db.Conn
db *gorm.DB
}
// NewRepository returns the access log repository backed by conn.
func NewRepository(conn *db.Conn) accesslogs.Repository {
return &sqlRepository{conn: conn, db: conn.DB(nil)}
}
func (r *sqlRepository) WithTx(tx *db.Tx) accesslogs.Repository {
return &sqlRepository{conn: r.conn, db: r.conn.DB(tx)}
}
func (r *sqlRepository) Create(ctx context.Context, entry *accesslogs.AccessLogEntry) error {
if err := r.db.Create(entry).Error; err != nil {
log.WithContext(ctx).WithFields(log.Fields{
"service_id": entry.ServiceID,
"method": entry.Method,
"host": entry.Host,
"path": entry.Path,
}).Errorf("failed to create access log entry in store: %v", err)
return status.Errorf(status.Internal, "failed to create access log entry in store")
}
return nil
}
// ListByAccount returns one page of an account's access logs together with the
// total number of entries matching the filter.
func (r *sqlRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) {
var totalCount int64
countQuery := applyFilters(r.db.Model(&accesslogs.AccessLogEntry{}).Where("account_id = ?", accountID), filter)
if err := countQuery.Count(&totalCount).Error; err != nil {
log.WithContext(ctx).Errorf("failed to count access logs: %v", err)
return nil, 0, status.Errorf(status.Internal, "failed to count access logs")
}
query := applyFilters(r.db.Where("account_id = ?", accountID), filter)
sortOrder := strings.ToUpper(filter.GetSortOrder())
for _, column := range strings.Split(filter.GetSortColumn(), ",") {
if column = strings.TrimSpace(column); column != "" {
query = query.Order(column + " " + sortOrder)
}
}
query = query.Limit(filter.GetLimit()).Offset(filter.GetOffset())
if lockStrength != db.LockingStrengthNone {
query = query.Clauses(clause.Locking{Strength: string(lockStrength)})
}
var logs []*accesslogs.AccessLogEntry
if err := query.Find(&logs).Error; err != nil {
log.WithContext(ctx).Errorf("failed to get access logs from store: %v", err)
return nil, 0, status.Errorf(status.Internal, "failed to get access logs from store")
}
return logs, totalCount, nil
}
func (r *sqlRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) {
result := r.db.Where("timestamp < ?", olderThan).Delete(&accesslogs.AccessLogEntry{})
if result.Error != nil {
log.WithContext(ctx).Errorf("failed to delete old access logs: %v", result.Error)
return 0, status.Errorf(status.Internal, "failed to delete old access logs")
}
return result.RowsAffected, nil
}
func applyFilters(query *gorm.DB, filter accesslogs.AccessLogFilter) *gorm.DB {
if filter.Search != nil {
searchPattern := "%" + *filter.Search + "%"
query = query.Where(
"id LIKE ? OR location_connection_ip LIKE ? OR host LIKE ? OR path LIKE ? OR CONCAT(host, path) LIKE ? OR user_id IN (SELECT id FROM users WHERE email LIKE ? OR name LIKE ?)",
searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern,
)
}
if filter.SourceIP != nil {
query = query.Where("location_connection_ip = ?", *filter.SourceIP)
}
if filter.Host != nil {
query = query.Where("host = ?", *filter.Host)
}
if filter.Path != nil {
query = query.Where("path LIKE ?", "%"+*filter.Path+"%")
}
if filter.UserID != nil {
query = query.Where("user_id = ?", *filter.UserID)
}
if filter.Method != nil {
query = query.Where("method = ?", *filter.Method)
}
if filter.Status != nil {
switch *filter.Status {
case "success":
query = query.Where("(status_code >= ? AND status_code < ?)", 200, 400)
case "failed":
query = query.Where("((status_code >= ? AND status_code < ?) OR status_code >= ?)", 100, 200, 400)
}
}
if filter.StatusCode != nil {
query = query.Where("status_code = ?", *filter.StatusCode)
}
if filter.StartDate != nil {
query = query.Where("timestamp >= ?", *filter.StartDate)
}
if filter.EndDate != nil {
query = query.Where("timestamp <= ?", *filter.EndDate)
}
return query
}
@@ -0,0 +1,125 @@
package manager
import (
"context"
"errors"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/management/internals/shared/db/dbtest"
)
func newTestRepository(t *testing.T) (accesslogs.Repository, *db.Conn) {
conn := dbtest.NewConn(t, &accesslogs.AccessLogEntry{})
return NewRepository(conn), conn
}
func newEntry(id, accountID, method string, age time.Duration) *accesslogs.AccessLogEntry {
return &accesslogs.AccessLogEntry{
ID: id,
AccountID: accountID,
Method: method,
Host: "app.example.com",
Path: "/",
StatusCode: 200,
Timestamp: time.Now().Add(-age),
}
}
func TestSqlRepository_ListByAccount(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
for _, entry := range []*accesslogs.AccessLogEntry{
newEntry("a1", "acc-a", "GET", 3*time.Hour),
newEntry("a2", "acc-a", "POST", 2*time.Hour),
newEntry("a3", "acc-a", "GET", time.Hour),
newEntry("b1", "acc-b", "GET", time.Hour),
} {
require.NoError(t, repo.Create(ctx, entry))
}
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 2})
require.NoError(t, err)
assert.EqualValues(t, 3, total)
require.Len(t, logs, 2)
assert.Equal(t, "a3", logs[0].ID)
assert.Equal(t, "a2", logs[1].ID)
method := "GET"
logs, total, err = repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Method: &method, SortOrder: "asc"})
require.NoError(t, err)
assert.EqualValues(t, 2, total)
require.Len(t, logs, 2)
assert.Equal(t, "a1", logs[0].ID)
assert.Equal(t, "a3", logs[1].ID)
}
func TestSqlRepository_DeleteOlderThan(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
require.NoError(t, repo.Create(ctx, newEntry("old", "acc", "GET", 48*time.Hour)))
require.NoError(t, repo.Create(ctx, newEntry("new", "acc", "GET", time.Hour)))
deleted, err := repo.DeleteOlderThan(ctx, time.Now().Add(-24*time.Hour))
require.NoError(t, err)
assert.EqualValues(t, 1, deleted)
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.EqualValues(t, 1, total)
require.Len(t, logs, 1)
assert.Equal(t, "new", logs[0].ID)
}
func TestSqlRepository_CreateInsideTransactionRollsBack(t *testing.T) {
repo, conn := newTestRepository(t)
ctx := context.Background()
failure := errors.New("abort")
err := conn.RunInTx(ctx, func(tx *db.Tx) error {
txRepo := repo.WithTx(tx)
require.NoError(t, txRepo.Create(ctx, newEntry("tx", "acc", "GET", 0)))
_, total, err := txRepo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.EqualValues(t, 1, total)
return failure
})
require.ErrorIs(t, err, failure)
_, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Zero(t, total)
}
func TestSqlRepository_ListByAccount_StatusFilter(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
statusCodes := map[string]int{"l4": 0, "info": 101, "ok": 200, "notfound": 404}
for id, code := range statusCodes {
entry := newEntry(id, "acc", "GET", time.Hour)
entry.StatusCode = code
require.NoError(t, repo.Create(ctx, entry))
}
foreign := newEntry("foreign", "other", "GET", time.Hour)
foreign.StatusCode = 500
require.NoError(t, repo.Create(ctx, foreign))
listIDs := func(status string) []string {
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Status: &status, SortBy: "status_code", SortOrder: "asc"})
require.NoError(t, err)
require.EqualValues(t, len(logs), total)
ids := make([]string, 0, len(logs))
for _, entry := range logs {
ids = append(ids, entry.ID)
}
return ids
}
assert.Equal(t, []string{"info", "notfound"}, listIDs("failed"))
assert.Equal(t, []string{"ok"}, listIDs("success"))
}
@@ -0,0 +1,18 @@
package accesslogs
import (
"context"
"time"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
//go:generate go tool mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod
// Repository persists reverse proxy access log entries.
type Repository interface {
WithTx(tx *db.Tx) Repository
Create(ctx context.Context, entry *AccessLogEntry) error
ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error)
DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error)
}
@@ -0,0 +1,102 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./repository.go
//
// Generated by this command:
//
// mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod
//
// Package accesslogs is a generated GoMock package.
package accesslogs
import (
context "context"
reflect "reflect"
time "time"
db "github.com/netbirdio/netbird/management/internals/shared/db"
gomock "go.uber.org/mock/gomock"
)
// MockRepository is a mock of Repository interface.
type MockRepository struct {
ctrl *gomock.Controller
recorder *MockRepositoryMockRecorder
isgomock struct{}
}
// MockRepositoryMockRecorder is the mock recorder for MockRepository.
type MockRepositoryMockRecorder struct {
mock *MockRepository
}
// NewMockRepository creates a new mock instance.
func NewMockRepository(ctrl *gomock.Controller) *MockRepository {
mock := &MockRepository{ctrl: ctrl}
mock.recorder = &MockRepositoryMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockRepository) EXPECT() *MockRepositoryMockRecorder {
return m.recorder
}
// Create mocks base method.
func (m *MockRepository) Create(ctx context.Context, entry *AccessLogEntry) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Create", ctx, entry)
ret0, _ := ret[0].(error)
return ret0
}
// Create indicates an expected call of Create.
func (mr *MockRepositoryMockRecorder) Create(ctx, entry any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Create", reflect.TypeOf((*MockRepository)(nil).Create), ctx, entry)
}
// DeleteOlderThan mocks base method.
func (m *MockRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "DeleteOlderThan", ctx, olderThan)
ret0, _ := ret[0].(int64)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// DeleteOlderThan indicates an expected call of DeleteOlderThan.
func (mr *MockRepositoryMockRecorder) DeleteOlderThan(ctx, olderThan any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteOlderThan", reflect.TypeOf((*MockRepository)(nil).DeleteOlderThan), ctx, olderThan)
}
// ListByAccount mocks base method.
func (m *MockRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ListByAccount", ctx, lockStrength, accountID, filter)
ret0, _ := ret[0].([]*AccessLogEntry)
ret1, _ := ret[1].(int64)
ret2, _ := ret[2].(error)
return ret0, ret1, ret2
}
// ListByAccount indicates an expected call of ListByAccount.
func (mr *MockRepositoryMockRecorder) ListByAccount(ctx, lockStrength, accountID, filter any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListByAccount", reflect.TypeOf((*MockRepository)(nil).ListByAccount), ctx, lockStrength, accountID, filter)
}
// WithTx mocks base method.
func (m *MockRepository) WithTx(tx *db.Tx) Repository {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "WithTx", tx)
ret0, _ := ret[0].(Repository)
return ret0
}
// WithTx indicates an expected call of WithTx.
func (mr *MockRepositoryMockRecorder) WithTx(tx any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WithTx", reflect.TypeOf((*MockRepository)(nil).WithTx), tx)
}
@@ -0,0 +1,80 @@
package manager
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/gorilla/mux"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/activity"
nbcontext "github.com/netbirdio/netbird/management/server/context"
nbstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/auth"
)
func TestDeleteDomain_ServiceDependencies(t *testing.T) {
for _, tt := range []struct {
name string
domainName string
serviceHost string
accountID string
enabled bool
protected bool
}{
{"exact", "example.com", "example.com", accountA, true, true},
{"subdomain", "example.com", "deep.app.example.com", accountA, true, true},
{"disabled", "example.com", "app.example.com", accountA, false, true},
// A service is authorized by its own account's registration, so another
// account's service under this namespace is not a dependency of it.
{"other account", "example.com", "app.example.com", accountB, true, false},
{"case and trailing dot", "example.com", "APP.EXAMPLE.COM.", accountA, true, true},
{"suffix boundary", "example.com", "notexample.com", accountA, true, false},
{"literal underscore", "a_b.example.com", "app.a_b.example.com", accountA, true, true},
{"underscore wildcard", "a_b.example.com", "app.axb.example.com", accountA, true, false},
} {
t.Run(tt.name, func(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
events := captureDomainEvents(env)
d, err := env.store.CreateCustomDomain(ctx, accountA, tt.domainName, testCluster, true)
require.NoError(t, err)
svc := &rpservice.Service{
ID: "dependent", AccountID: tt.accountID, Domain: tt.serviceHost,
Enabled: tt.enabled, ProxyCluster: testCluster,
}
require.NoError(t, env.store.CreateService(ctx, svc))
router := mux.NewRouter()
RegisterEndpoints(router, env.manager)
deleteDomain := func() *httptest.ResponseRecorder {
req := httptest.NewRequest(http.MethodDelete, "/domains/"+d.ID, nil)
req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: accountA, UserId: accountAUser})
response := httptest.NewRecorder()
router.ServeHTTP(response, req)
return response
}
response := deleteDomain()
if tt.protected {
require.Equal(t, http.StatusPreconditionFailed, response.Code, "dependent services must block deletion: %s", response.Body.String())
assert.NotContains(t, response.Body.String(), tt.accountID, "the error must not reveal the service's account")
assert.NotNil(t, storedDomain(t, env.store, accountA, d.Domain), "the namespace must remain reserved")
assert.Empty(t, events.get(), "rejected deletion must not emit DomainDeleted")
stored, err := env.store.GetServiceByID(ctx, nbstore.LockingStrengthNone, tt.accountID, svc.ID)
require.NoError(t, err)
assert.Equal(t, svc.Enabled, stored.Enabled, "rejected deletion must preserve the service")
require.NoError(t, env.store.DeleteService(ctx, tt.accountID, svc.ID))
response = deleteDomain()
}
require.Equal(t, http.StatusNoContent, response.Code, "deletion must succeed without dependencies: %s", response.Body.String())
assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "the registration must be deleted")
captured := events.get()
require.Len(t, captured, 1, "only successful deletion may emit an event")
assert.Equal(t, activity.DomainDeleted, captured[0].Activity, "the event must describe the successful deletion")
})
}
}
@@ -357,6 +357,26 @@ func (m Manager) DeriveClusterFromDomain(ctx context.Context, accountID, domain
return "", fmt.Errorf("domain %s does not match any available proxy cluster", domain)
}
// ValidateServiceDomain holds custom domain authorization through a service write transaction.
func (m Manager) ValidateServiceDomain(ctx context.Context, tx nbstore.Store, accountID, serviceDomain, cluster string) error {
if _, ok := ExtractClusterFromFreeDomain(serviceDomain, []string{cluster}); ok {
return nil
}
name, err := nbdomain.FromString(serviceDomain)
if err != nil {
return status.Errorf(status.InvalidArgument, "invalid service domain: %v", err)
}
customDomains, err := tx.LockCustomDomains(ctx, accountID, name)
if err != nil {
return err
}
target, match := extractClusterFromCustomDomains(serviceDomain, customDomains)
if match != customDomainValidated || target != cluster {
return status.Errorf(status.PreconditionFailed, "custom domain authorization changed; retry the service operation")
}
return nil
}
func (m Manager) getClusterAllowList(ctx context.Context, accountID string) ([]string, error) {
byopAddresses, err := m.proxyManager.GetActiveClusterAddressesForAccount(ctx, accountID)
if err != nil {
@@ -99,7 +99,7 @@ func setupDomainTest(t *testing.T) *domainTestEnv {
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", nil, nil)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", "", nil, nil)
require.NoError(t, err)
resolver := &stubResolver{cnames: make(map[string]string)}
@@ -11,7 +11,7 @@ import (
// Manager defines the interface for proxy operations
type Manager interface {
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error)
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error)
Disconnect(ctx context.Context, proxyID, sessionID string) error
Heartbeat(ctx context.Context, p *Proxy) error
GetActiveClusterAddresses(ctx context.Context) ([]string, error)
@@ -20,6 +20,8 @@ type Manager interface {
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool
CleanupStale(ctx context.Context, inactivityDuration time.Duration) error
GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error)
CountAccountProxies(ctx context.Context, accountID string) (int64, error)
@@ -8,6 +8,7 @@ import (
"go.opentelemetry.io/otel/metric"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
nbversion "github.com/netbirdio/netbird/version"
)
// store defines the interface for proxy persistence operations
@@ -22,6 +23,8 @@ type store interface {
GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error)
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error)
CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error)
@@ -29,6 +32,8 @@ type store interface {
DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error
}
const minSessionCodeVersion = "0.81.0"
// Manager handles all proxy operations
type Manager struct {
store store
@@ -50,7 +55,7 @@ func NewManager(store store, meter metric.Meter) (*Manager, error) {
// Connect registers a new proxy connection in the database.
// capabilities may be nil for old proxies that do not report them.
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
now := time.Now()
var caps proxy.Capabilities
if capabilities != nil {
@@ -61,6 +66,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
SessionID: sessionID,
ClusterAddress: clusterAddress,
IPAddress: ipAddress,
Version: truncateVersion(version),
AccountID: accountID,
LastSeen: now,
ConnectedAt: &now,
@@ -78,6 +84,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
"sessionID": sessionID,
"clusterAddress": clusterAddress,
"ipAddress": ipAddress,
"version": p.Version,
}).Info("proxy connected")
return p, nil
@@ -143,6 +150,27 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string)
return m.store.GetClusterSupportsPrivate(ctx, clusterAddr)
}
// ClusterAllProxiesPrivate reports whether every active proxy claims the private capability (nil = unreported).
func (m Manager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
return m.store.GetClusterAllProxiesPrivate(ctx, clusterAddr)
}
// ClusterSupportsSessionCode reports whether all active proxies support session codes.
func (m Manager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
versions, err := m.store.GetActiveProxyVersions(ctx, clusterAddr)
if err != nil || len(versions) == 0 {
return false
}
for _, version := range versions {
if supported, err := nbversion.MeetsMinVersion(minSessionCodeVersion, version); err != nil || !supported {
return false
}
}
return true
}
// CleanupStale removes proxies that haven't sent heartbeat in the specified duration
func (m *Manager) CleanupStale(ctx context.Context, inactivityDuration time.Duration) error {
if err := m.store.CleanupStaleProxies(ctx, inactivityDuration); err != nil {
@@ -184,3 +212,13 @@ func (m *Manager) DeleteAccountCluster(ctx context.Context, clusterAddress, acco
}
return nil
}
// truncateVersion cuts a proxy-reported version to the column width so an
// oversized value cannot fail the save and block the connect.
func truncateVersion(version string) string {
runes := []rune(version)
if len(runes) <= proxy.MaxVersionLength {
return version
}
return string(runes[:proxy.MaxVersionLength])
}
@@ -4,8 +4,10 @@ import (
"context"
"errors"
"fmt"
"strings"
"testing"
"time"
"unicode/utf8"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -20,6 +22,7 @@ type mockStore struct {
updateProxyHeartbeatFunc func(ctx context.Context, p *proxy.Proxy) error
getActiveProxyClusterAddressesFunc func(ctx context.Context) ([]string, error)
getActiveProxyClusterAddressesForAccFunc func(ctx context.Context, accountID string) ([]string, error)
getActiveProxyVersionsFunc func(ctx context.Context, clusterAddress string) ([]string, error)
cleanupStaleProxiesFunc func(ctx context.Context, d time.Duration) error
getProxyByAccountIDFunc func(ctx context.Context, accountID string) (*proxy.Proxy, error)
countProxiesByAccountIDFunc func(ctx context.Context, accountID string) (int64, error)
@@ -102,6 +105,15 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo
func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool {
return nil
}
func (m *mockStore) GetClusterAllProxiesPrivate(_ context.Context, _ string) *bool {
return nil
}
func (m *mockStore) GetActiveProxyVersions(ctx context.Context, clusterAddress string) ([]string, error) {
if m.getActiveProxyVersionsFunc != nil {
return m.getActiveProxyVersionsFunc(ctx, clusterAddress)
}
return nil, nil
}
func newTestManager(s store) *Manager {
meter := noop.NewMeterProvider().Meter("test")
@@ -112,6 +124,34 @@ func newTestManager(s store) *Manager {
return m
}
func TestClusterSupportsSessionCode(t *testing.T) {
tests := []struct {
name string
versions []string
storeErr error
want bool
}{
{name: "all supported", versions: []string{"0.81.0", "0.81.2"}, want: true},
{name: "one old proxy", versions: []string{"0.81.0", "0.80.0"}},
{name: "missing version", versions: []string{"0.81.0", ""}},
{name: "no active proxies"},
{name: "store error", storeErr: errors.New("db error")},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
s := &mockStore{
getActiveProxyVersionsFunc: func(_ context.Context, _ string) ([]string, error) {
return tt.versions, tt.storeErr
},
}
got := newTestManager(s).ClusterSupportsSessionCode(context.Background(), "cluster.example.com")
assert.Equal(t, tt.want, got)
})
}
}
func TestConnect_WithAccountID(t *testing.T) {
accountID := "acc-123"
@@ -124,7 +164,7 @@ func TestConnect_WithAccountID(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", &accountID, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "0.60.0", &accountID, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
@@ -132,6 +172,7 @@ func TestConnect_WithAccountID(t *testing.T) {
assert.Equal(t, "session-1", savedProxy.SessionID)
assert.Equal(t, "cluster.example.com", savedProxy.ClusterAddress)
assert.Equal(t, "10.0.0.1", savedProxy.IPAddress)
assert.Equal(t, "0.60.0", savedProxy.Version, "reported proxy version should be stored")
assert.Equal(t, &accountID, savedProxy.AccountID)
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
assert.NotNil(t, savedProxy.ConnectedAt)
@@ -147,7 +188,7 @@ func TestConnect_WithoutAccountID(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", nil, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", "", nil, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
@@ -155,6 +196,29 @@ func TestConnect_WithoutAccountID(t *testing.T) {
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
}
func TestConnect_TruncatesOversizedVersion(t *testing.T) {
var savedProxy *proxy.Proxy
s := &mockStore{
saveProxyFunc: func(_ context.Context, p *proxy.Proxy) error {
savedProxy = p
return nil
},
}
// Multi-byte runes make sure the cut counts characters, as varchar does,
// and never splits a rune into invalid UTF-8.
version := strings.Repeat("ü", proxy.MaxVersionLength+10)
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", version, nil, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
assert.Equal(t, proxy.MaxVersionLength, utf8.RuneCountInString(savedProxy.Version), "stored version should be cut to the column width")
assert.True(t, utf8.ValidString(savedProxy.Version), "stored version should remain valid UTF-8")
assert.True(t, strings.HasPrefix(version, savedProxy.Version), "stored version should be a prefix of the reported one")
}
func TestConnect_StoreError(t *testing.T) {
s := &mockStore{
saveProxyFunc: func(_ context.Context, _ *proxy.Proxy) error {
@@ -163,7 +227,7 @@ func TestConnect_StoreError(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", nil, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "", nil, nil)
assert.Error(t, err)
}
@@ -56,6 +56,20 @@ func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *go
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration)
}
// ClusterAllProxiesPrivate mocks base method.
func (m *MockManager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ClusterAllProxiesPrivate", ctx, clusterAddr)
ret0, _ := ret[0].(*bool)
return ret0
}
// ClusterAllProxiesPrivate indicates an expected call of ClusterAllProxiesPrivate.
func (mr *MockManagerMockRecorder) ClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterAllProxiesPrivate", reflect.TypeOf((*MockManager)(nil).ClusterAllProxiesPrivate), ctx, clusterAddr)
}
// ClusterRequireSubdomain mocks base method.
func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool {
m.ctrl.T.Helper()
@@ -112,19 +126,33 @@ func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
}
// Connect mocks base method.
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
// ClusterSupportsSessionCode mocks base method.
func (m *MockManager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
ret := m.ctrl.Call(m, "ClusterSupportsSessionCode", ctx, clusterAddr)
ret0, _ := ret[0].(bool)
return ret0
}
// ClusterSupportsSessionCode indicates an expected call of ClusterSupportsSessionCode.
func (mr *MockManagerMockRecorder) ClusterSupportsSessionCode(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsSessionCode", reflect.TypeOf((*MockManager)(nil).ClusterSupportsSessionCode), ctx, clusterAddr)
}
// Connect mocks base method.
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
ret0, _ := ret[0].(*Proxy)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// Connect indicates an expected call of Connect.
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call {
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
}
// CountAccountProxies mocks base method.
@@ -9,6 +9,9 @@ const (
StatusDisconnected = "disconnected"
)
// MaxVersionLength is the width of the Version column, in characters.
const MaxVersionLength = 255
// Capabilities describes what a proxy can handle, as reported via gRPC.
// Nil fields mean the proxy never reported this capability.
type Capabilities struct {
@@ -31,6 +34,7 @@ type Proxy struct {
SessionID string `gorm:"type:varchar(36)"`
ClusterAddress string `gorm:"type:varchar(255);not null;index:idx_proxy_cluster_status"`
IPAddress string `gorm:"type:varchar(45)"`
Version string `gorm:"type:varchar(255)"`
AccountID *string `gorm:"type:varchar(255);index:idx_proxy_account_id"`
LastSeen time.Time `gorm:"not null;index:idx_proxy_last_seen"`
ConnectedAt *time.Time
@@ -1,6 +1,7 @@
package proxytoken
import (
"context"
"encoding/json"
"net/http"
"time"
@@ -18,13 +19,29 @@ import (
"github.com/netbirdio/netbird/shared/management/status"
)
// RevocationGuard vetoes the tenant-facing revocation of a proxy access
// token. Implementations are supplied by integrations; none is installed by
// default, so every token the caller's account owns may be revoked. It is
// consulted after the ownership check and before the token is revoked. A
// returned status error is written with util.WriteError: its type selects the
// HTTP status and its message is shown to the caller, so it must not carry
// internal detail. Any other error is reported as a generic internal error.
type RevocationGuard interface {
CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error
}
type handler struct {
store store.Store
permissionsManager permissions.Manager
// revocationGuard vetoes revocations. Optional — when nil every owned
// token may be revoked.
revocationGuard RevocationGuard
}
func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, router *mux.Router) {
h := &handler{store: s, permissionsManager: permissionsManager}
// RegisterEndpoints registers the proxy token endpoints. revocationGuard is
// optional; pass nil for no revocation policy.
func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, revocationGuard RevocationGuard, router *mux.Router) {
h := &handler{store: s, permissionsManager: permissionsManager, revocationGuard: revocationGuard}
router.HandleFunc("/reverse-proxies/proxy-tokens", h.listTokens).Methods("GET", "OPTIONS")
router.HandleFunc("/reverse-proxies/proxy-tokens", h.createToken).Methods("POST", "OPTIONS")
router.HandleFunc("/reverse-proxies/proxy-tokens/{tokenId}", h.revokeToken).Methods("DELETE", "OPTIONS")
@@ -154,6 +171,13 @@ func (h *handler) revokeToken(w http.ResponseWriter, r *http.Request) {
return
}
if h.revocationGuard != nil {
if err := h.revocationGuard.CheckProxyAccessTokenRevocation(ctx, token); err != nil {
util.WriteError(ctx, err, w)
return
}
}
if err := h.store.RevokeProxyAccessToken(ctx, tokenID); err != nil {
util.WriteErrorResponse("failed to revoke token", http.StatusInternalServerError, w)
return
@@ -4,6 +4,7 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"testing"
@@ -22,6 +23,7 @@ import (
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/auth"
"github.com/netbirdio/netbird/shared/management/http/api"
"github.com/netbirdio/netbird/shared/management/status"
)
func authContext(accountID, userID string) context.Context {
@@ -273,3 +275,152 @@ func TestRevokeToken_ManagementWideToken(t *testing.T) {
h.revokeToken(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
}
type revocationGuardFunc func(ctx context.Context, token *types.ProxyAccessToken) error
func (f revocationGuardFunc) CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error {
return f(ctx, token)
}
func TestRevokeToken_GuardRefuses(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
// No RevokeProxyAccessToken expectation: a refused revocation must not
// reach the store.
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
var checked *types.ProxyAccessToken
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(_ context.Context, token *types.ProxyAccessToken) error {
checked = token
return status.Errorf(status.PreconditionFailed, "token is in use")
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusPreconditionFailed, w.Code)
assert.Contains(t, w.Body.String(), "token is in use")
require.NotNil(t, checked)
assert.Equal(t, "tok-1", checked.ID)
}
func TestRevokeToken_GuardAllows(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
mockStore.EXPECT().RevokeProxyAccessToken(gomock.Any(), "tok-1").Return(nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
return nil
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusOK, w.Code)
}
func TestRevokeToken_GuardFailure(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
// No RevokeProxyAccessToken expectation: a guard that cannot decide must
// not let the revocation through.
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
return errors.New("connection refused")
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusInternalServerError, w.Code)
assert.Contains(t, w.Body.String(), "internal server error")
assert.NotContains(t, w.Body.String(), "connection refused")
}
func TestRevokeToken_GuardNotConsultedForForeignToken(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
otherAccount := "acc-other"
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &otherAccount,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), "acc-123", "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
// A foreign token must read as not found, not reveal through the guard's
// answer that it belongs to some account's managed proxy.
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
t.Fatal("guard consulted for a token the caller does not own")
return nil
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext("acc-123", "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
}
@@ -30,7 +30,7 @@ func withRealDomainManager(t *testing.T, mgr *Manager, testStore store.Store) {
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", nil, nil)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", "", nil, nil)
require.NoError(t, err)
accountMgr := &mock_server.MockAccountManager{
@@ -125,3 +125,53 @@ func TestUpdateService_RefusesMoveToUnvalidatedDomain(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, "app.proven.example.com", stored.Domain, "the service must keep its original domain")
}
func TestCreateService_DomainDeletedBeforeWrite(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
d, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true)
require.NoError(t, err)
svc := newTestService("app.proven.example.com")
require.NoError(t, mgr.initializeServiceForCreate(ctx, testAccountID, svc))
// Delete after the initial authorization check, before the service transaction starts.
require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID))
err = mgr.persistNewService(ctx, testAccountID, svc)
require.Error(t, err, "an earlier validation result must not authorize a deleted registration")
sErr, ok := status.FromError(err)
require.True(t, ok, "the caller must receive a typed precondition error")
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the service must require current domain authorization")
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, err)
assert.Empty(t, services, "the failed write must not leave a service")
}
func TestUpdateService_DomainDeletedBeforeWrite(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "original.example.com", validationTestCluster, true)
require.NoError(t, err)
d, err := testStore.CreateCustomDomain(ctx, testAccountID, "destination.example.com", validationTestCluster, true)
require.NoError(t, err)
svc, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.original.example.com"))
require.NoError(t, err)
moved := svc.Copy()
moved.Domain = "app.destination.example.com"
cluster, err := mgr.resolveEffectiveCluster(ctx, testAccountID, moved)
require.NoError(t, err)
require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID))
err = testStore.ExecuteInTransaction(ctx, func(tx store.Store) error {
return mgr.executeServiceUpdate(ctx, tx, testAccountID, moved, &serviceUpdateInfo{}, nil, cluster)
})
require.Error(t, err, "a domain deleted after cluster resolution must reject the update")
sErr, ok := status.FromError(err)
require.True(t, ok, "the caller must receive a typed precondition error")
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the move must require current domain authorization")
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, svc.ID)
require.NoError(t, err)
assert.Equal(t, svc.Domain, stored.Domain, "the service must retain its authorized domain")
}
@@ -74,6 +74,7 @@ const unknownHostPlaceholder = "unknown"
// ClusterDeriver derives the proxy cluster from a domain.
type ClusterDeriver interface {
DeriveClusterFromDomain(ctx context.Context, accountID, domain string) (string, error)
ValidateServiceDomain(ctx context.Context, tx store.Store, accountID, domain, cluster string) error
GetClusterDomains() []string
}
@@ -83,6 +84,7 @@ type CapabilityProvider interface {
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
}
type Manager struct {
@@ -331,7 +333,14 @@ func (m *Manager) persistNewService(ctx context.Context, accountID string, svc *
return err
}
if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil {
return err
}
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
return err
}
if svc.Domain != "" {
if err := m.checkDomainAvailable(ctx, transaction, svc.Domain, ""); err != nil {
return err
@@ -365,6 +374,43 @@ func (m *Manager) clusterCustomPorts(ctx context.Context, svc *service.Service)
return m.capabilities.ClusterSupportsCustomPorts(ctx, svc.ProxyCluster)
}
// validatePrivateClusterTargets rejects cluster and direct upstream targets unless
// every active proxy in the service's cluster reports the private capability. The
// mapping reaches all proxies in the cluster, so one non-private proxy would serve
// these targets too. An unreported capability is treated as unsupported. Must be
// called outside a transaction, like clusterCustomPorts.
func (m *Manager) validatePrivateClusterTargets(ctx context.Context, targets []*service.Target, cluster string) error {
target := firstPrivateClusterTarget(targets)
if target == nil {
return nil
}
if private := m.capabilities.ClusterAllProxiesPrivate(ctx, cluster); private != nil && *private {
return nil
}
if target.TargetType == service.TargetTypeCluster {
return status.Errorf(status.InvalidArgument,
"target_type %q requires a proxy cluster with private mode enabled, cluster %s does not support it",
service.TargetTypeCluster, cluster)
}
return status.Errorf(status.InvalidArgument,
"direct_upstream requires a proxy cluster with private mode enabled, cluster %s does not support it", cluster)
}
// firstPrivateClusterTarget returns the first target that only a private cluster may serve.
func firstPrivateClusterTarget(targets []*service.Target) *service.Target {
for _, target := range targets {
if target == nil {
continue
}
if target.TargetType == service.TargetTypeCluster || target.Options.DirectUpstream {
return target
}
}
return nil
}
// ensureL4Port auto-assigns a listen port when needed and validates cluster support.
// customPorts must be pre-computed via clusterCustomPorts before entering a transaction.
func (m *Manager) ensureL4Port(ctx context.Context, tx store.Store, svc *service.Service, customPorts *bool, serviceUpdate bool) error {
@@ -460,7 +506,14 @@ func (m *Manager) persistNewEphemeralService(ctx context.Context, accountID, pee
return err
}
if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil {
return err
}
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
return err
}
if err := m.validateEphemeralPreconditions(ctx, transaction, accountID, peerID, svc); err != nil {
return err
}
@@ -577,6 +630,10 @@ func (m *Manager) persistServiceUpdate(ctx context.Context, accountID string, se
return nil, err
}
if err := m.validatePrivateClusterTargets(ctx, service.Targets, effectiveCluster); err != nil {
return nil, err
}
// Validate subdomain requirement *before* the transaction: the underlying
// capability lookup talks to the main DB pool, and SQLite's single-connection
// pool would self-deadlock if this ran while the tx already held the only
@@ -622,6 +679,9 @@ func (m *Manager) resolveEffectiveCluster(ctx context.Context, accountID string,
}
func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.Store, accountID string, service *service.Service, updateInfo *serviceUpdateInfo, customPorts *bool, effectiveCluster string) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, service, effectiveCluster); err != nil {
return err
}
existingService, err := transaction.GetServiceByID(ctx, store.LockingStrengthUpdate, accountID, service.ID)
if err != nil {
return err
@@ -677,6 +737,13 @@ func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.St
return nil
}
func (m *Manager) validateServiceDomain(ctx context.Context, tx store.Store, accountID string, svc *service.Service, cluster string) error {
if m.clusterDeriver == nil {
return nil
}
return m.clusterDeriver.ValidateServiceDomain(ctx, tx, accountID, svc.Domain, cluster)
}
// validateL4PortDiffOnClusterDiff checks if custom L4 ports are configured and validates port changes across clusters.
// It ensures no port changes if custom ports are unsupported for a given cluster and protocol mode.
// Returns an error if validation fails, otherwise returns nil.
@@ -433,8 +433,8 @@ func TestDeletePeerService_SourcePeerValidation(t *testing.T) {
newProxyServer := func(t *testing.T) *nbgrpc.ProxyServiceServer {
t.Helper()
tokenStore := nbgrpc.NewOneTimeTokenStore(context.Background(), testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(context.Background(), testCacheStore(t))
srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(context.Background(), testCacheStore(t))
srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
return srv
}
@@ -655,6 +655,10 @@ func (d *testClusterDeriver) GetClusterDomains() []string {
return d.domains
}
func (d *testClusterDeriver) ValidateServiceDomain(context.Context, store.Store, string, string, string) error {
return nil
}
const (
testAccountID = "test-account"
testPeerID = "test-peer-1"
@@ -722,8 +726,8 @@ func setupIntegrationTest(t *testing.T) (*Manager, store.Store) {
}
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
@@ -1146,8 +1150,8 @@ func TestDeleteService_DeletesTargets(t *testing.T) {
mockAcct := account.NewMockManager(ctrl)
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
@@ -0,0 +1,218 @@
package manager
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/metric/noop"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
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/store"
"github.com/netbirdio/netbird/shared/management/status"
)
// setupPrivateClusterTest wires the real proxy manager as the capability
// provider and connects one proxy to testCluster reporting the given private
// capability. A nil private connects no proxy, so the capability is unreported.
func setupPrivateClusterTest(t *testing.T, private *bool) (*Manager, store.Store) {
t.Helper()
mgr, testStore := setupIntegrationTest(t)
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
mgr.capabilities = proxyMgr
if private != nil {
connectTestProxy(t, proxyMgr, "proxy-1", &proxy.Capabilities{Private: private})
}
return mgr, testStore
}
func connectTestProxy(t *testing.T, proxyMgr *proxymanager.Manager, proxyID string, caps *proxy.Capabilities) {
t.Helper()
_, err := proxyMgr.Connect(context.Background(), proxyID, "session-"+proxyID, testCluster, "127.0.0.1", "", nil, caps)
require.NoError(t, err)
}
func clusterTarget() *rpservice.Target {
return &rpservice.Target{
TargetId: testCluster,
TargetType: rpservice.TargetTypeCluster,
Host: "backend.lan",
Port: 8080,
Protocol: "http",
Enabled: true,
Options: rpservice.TargetOptions{DirectUpstream: true},
}
}
func directUpstreamPeerTarget() *rpservice.Target {
return &rpservice.Target{
TargetId: testPeerID,
TargetType: rpservice.TargetTypePeer,
Host: "backend.lan",
Port: 8080,
Protocol: "http",
Enabled: true,
Options: rpservice.TargetOptions{DirectUpstream: true},
}
}
func TestCreateService_PrivateClusterTargets(t *testing.T) {
tests := []struct {
name string
private *bool
target *rpservice.Target
wantErr string
}{
{name: "cluster target on private cluster", private: boolPtr(true), target: clusterTarget()},
{name: "direct upstream on private cluster", private: boolPtr(true), target: directUpstreamPeerTarget()},
{name: "cluster target on non-private cluster", private: boolPtr(false), target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
{name: "direct upstream on non-private cluster", private: boolPtr(false), target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
{name: "cluster target with unreported capability", private: nil, target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
{name: "direct upstream with unreported capability", private: nil, target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupPrivateClusterTest(t, tc.private)
svc := newTestService("app.test.netbird.io")
svc.Targets = []*rpservice.Target{tc.target}
_, err := mgr.CreateService(ctx, testAccountID, testUserID, svc)
services, listErr := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, listErr)
if tc.wantErr == "" {
require.NoError(t, err)
assert.Len(t, services, 1, "the service should be persisted")
return
}
require.Error(t, err)
assert.Contains(t, err.Error(), tc.wantErr)
sErr, ok := status.FromError(err)
require.True(t, ok, "the caller must receive a typed error")
assert.Equal(t, status.InvalidArgument, sErr.Type(), "the rejection should be an invalid argument")
assert.Empty(t, services, "a rejected service must not be persisted")
})
}
}
// A cluster where only some proxies run in private mode must not accept these
// targets: the mapping is delivered to every proxy in the cluster, so the
// non-private ones would serve the target from their host network as well.
func TestCreateService_MixedClusterRejectsPrivateTargets(t *testing.T) {
tests := []struct {
name string
secondCaps *proxy.Capabilities
}{
{name: "second proxy reports not private", secondCaps: &proxy.Capabilities{Private: boolPtr(false)}},
{name: "second proxy predates capability reporting", secondCaps: nil},
}
for _, tc := range tests {
for _, target := range []*rpservice.Target{clusterTarget(), directUpstreamPeerTarget()} {
t.Run(tc.name+"/"+string(target.TargetType), func(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupPrivateClusterTest(t, boolPtr(true))
connectTestProxy(t, mgr.capabilities.(*proxymanager.Manager), "proxy-2", tc.secondCaps)
svc := newTestService("app.test.netbird.io")
svc.Targets = []*rpservice.Target{target}
_, err := mgr.CreateService(ctx, testAccountID, testUserID, svc)
require.Error(t, err, "a cluster with a non-private proxy must not accept the target")
assert.Contains(t, err.Error(), "requires a proxy cluster with private mode enabled")
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, err)
assert.Empty(t, services, "a rejected service must not be persisted")
})
}
}
}
func TestCreateService_RegularTargetIgnoresPrivateCapability(t *testing.T) {
ctx := context.Background()
mgr, _ := setupPrivateClusterTest(t, boolPtr(false))
_, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
require.NoError(t, err, "a peer target without direct upstream must not need a private cluster")
}
func TestUpdateService_PrivateClusterTargets(t *testing.T) {
tests := []struct {
name string
target *rpservice.Target
wantErr string
}{
{name: "switch to cluster target", target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
{name: "enable direct upstream", target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupPrivateClusterTest(t, boolPtr(false))
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
require.NoError(t, err)
updated := newTestService("app.test.netbird.io")
updated.ID = created.ID
updated.AccountID = testAccountID
updated.Targets = []*rpservice.Target{tc.target}
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated)
require.Error(t, err)
assert.Contains(t, err.Error(), tc.wantErr)
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
require.NoError(t, err)
require.Len(t, stored.Targets, 1)
assert.Equal(t, rpservice.TargetTypePeer, stored.Targets[0].TargetType, "the stored target must be unchanged")
assert.False(t, stored.Targets[0].Options.DirectUpstream, "the stored target must keep direct upstream disabled")
})
}
}
func TestUpdateService_PrivateClusterAllowsClusterTarget(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupPrivateClusterTest(t, boolPtr(true))
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
require.NoError(t, err)
updated := newTestService("app.test.netbird.io")
updated.ID = created.ID
updated.AccountID = testAccountID
updated.Targets = []*rpservice.Target{clusterTarget()}
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated)
require.NoError(t, err)
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
require.NoError(t, err)
require.Len(t, stored.Targets, 1)
assert.Equal(t, rpservice.TargetTypeCluster, stored.Targets[0].TargetType, "the cluster target should be stored")
}
func TestValidatePrivateClusterTargets_NoLookupWithoutPrivateTargets(t *testing.T) {
ctrl := gomock.NewController(t)
// No ClusterAllProxiesPrivate expectation: a lookup would fail the test.
mgr := &Manager{capabilities: proxy.NewMockManager(ctrl)}
targets := []*rpservice.Target{{TargetId: testPeerID, TargetType: rpservice.TargetTypePeer}}
require.NoError(t, mgr.validatePrivateClusterTargets(context.Background(), targets, testCluster))
}
@@ -3,15 +3,16 @@ package manager
import (
"context"
"fmt"
"slices"
"github.com/netbirdio/netbird/management/internals/modules/zones"
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/affectedpeers"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/status"
)
@@ -69,6 +70,9 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
}
zone = zones.NewZone(accountID, zone.Name, zone.Domain, zone.Enabled, zone.EnableSearchDomain, zone.DistributionGroups)
var snap *affectedpeers.Snapshot
change := affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
existingZone, err := transaction.GetZoneByDomain(ctx, accountID, zone.Domain)
if err != nil {
@@ -88,7 +92,15 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
}
if err = transaction.CreateZone(ctx, zone); err != nil {
return fmt.Errorf("failed to create zone: %w", err)
return fmt.Errorf("create zone: %w", err)
}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
return fmt.Errorf("increment network serial: %w", err)
}
return nil
@@ -99,6 +111,8 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneCreated, zone.EventMeta())
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return zone, nil
}
@@ -111,21 +125,26 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
return nil, status.NewPermissionDeniedError()
}
zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID)
if err != nil {
return nil, fmt.Errorf("failed to get zone: %w", err)
}
if zone.Domain != updatedZone.Domain {
return nil, status.Errorf(status.InvalidArgument, "zone domain cannot be updated")
}
zone.Name = updatedZone.Name
zone.Enabled = updatedZone.Enabled
zone.EnableSearchDomain = updatedZone.EnableSearchDomain
zone.DistributionGroups = updatedZone.DistributionGroups
var zone *zones.Zone
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID)
if err != nil {
return fmt.Errorf("get zone: %w", err)
}
if zone.Domain != updatedZone.Domain {
return status.Errorf(status.InvalidArgument, "zone domain cannot be updated")
}
oldGroups := zone.DistributionGroups
zone.Name = updatedZone.Name
zone.Enabled = updatedZone.Enabled
zone.EnableSearchDomain = updatedZone.EnableSearchDomain
zone.DistributionGroups = updatedZone.DistributionGroups
for _, groupID := range zone.DistributionGroups {
_, err = transaction.GetGroupByID(ctx, store.LockingStrengthNone, accountID, groupID)
if err != nil {
@@ -134,7 +153,16 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
}
if err = transaction.UpdateZone(ctx, zone); err != nil {
return fmt.Errorf("failed to update zone: %w", err)
return fmt.Errorf("update zone: %w", err)
}
change = affectedpeers.Change{DistributionGroupIDs: slices.Concat(zone.DistributionGroups, oldGroups)}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
return fmt.Errorf("increment network serial: %w", err)
}
return nil
@@ -145,7 +173,7 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneUpdated, zone.EventMeta())
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationUpdate})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return zone, nil
}
@@ -159,13 +187,23 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID
return status.NewPermissionDeniedError()
}
zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
if err != nil {
return fmt.Errorf("failed to get zone: %w", err)
}
var zone *zones.Zone
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
var eventsToStore []func()
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
if err != nil {
return fmt.Errorf("get zone: %w", err)
}
// Load before delete: the post-delete state no longer references the groups.
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
records, err := transaction.GetZoneDNSRecords(ctx, store.LockingStrengthNone, accountID, zoneID)
if err != nil {
return fmt.Errorf("failed to get records: %w", err)
@@ -207,7 +245,7 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID
event()
}
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationDelete})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return nil
}
@@ -9,11 +9,11 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/zones/records"
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/affectedpeers"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/status"
)
@@ -65,6 +65,8 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
}
var zone *zones.Zone
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
record = records.NewRecord(accountID, zoneID, record.Name, record.Type, record.Content, record.TTL)
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
@@ -82,6 +84,11 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
return fmt.Errorf("failed to create dns record: %w", err)
}
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
err = transaction.IncrementNetworkSerial(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to increment network serial: %w", err)
@@ -96,7 +103,7 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
meta := record.EventMeta(zone.ID, zone.Name)
m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordCreated, meta)
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationCreate})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return record, nil
}
@@ -112,6 +119,8 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
var zone *zones.Zone
var record *records.Record
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
@@ -141,6 +150,11 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
return fmt.Errorf("failed to update dns record: %w", err)
}
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
err = transaction.IncrementNetworkSerial(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to increment network serial: %w", err)
@@ -155,7 +169,7 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
meta := record.EventMeta(zone.ID, zone.Name)
m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordUpdated, meta)
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationUpdate})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return record, nil
}
@@ -171,6 +185,8 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
var record *records.Record
var zone *zones.Zone
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
@@ -188,6 +204,11 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
return fmt.Errorf("failed to delete dns record: %w", err)
}
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
err = transaction.IncrementNetworkSerial(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to increment network serial: %w", err)
@@ -202,7 +223,7 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
meta := record.EventMeta(zone.ID, zone.Name)
m.accountManager.StoreEvent(ctx, userID, recordID, accountID, activity.DNSRecordDeleted, meta)
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationDelete})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return nil
}
+24 -12
View File
@@ -32,17 +32,18 @@ import (
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"
"github.com/netbirdio/netbird/management/internals/shared/db"
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"
nbcache "github.com/netbirdio/netbird/management/server/cache"
nbContext "github.com/netbirdio/netbird/management/server/context"
nbhttp "github.com/netbirdio/netbird/management/server/http"
"github.com/netbirdio/netbird/management/server/http/middleware"
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/telemetry"
mgmtProto "github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/ratelimit"
"github.com/netbirdio/netbird/util/crypt"
)
@@ -84,9 +85,20 @@ func (s *BaseServer) CacheStore() nbcache.Store {
})
}
// DBConn opens the database connection shared by the store and the domain repositories.
func (s *BaseServer) DBConn() *db.Conn {
return Create(s, func() *db.Conn {
conn, err := store.OpenConn(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir)
if err != nil {
log.Fatalf("failed to open database connection: %v", err)
}
return conn
})
}
func (s *BaseServer) Store() store.Store {
return Create(s, func() store.Store {
store, err := store.NewStore(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir, s.Metrics(), false)
store, err := store.NewSqlStore(context.Background(), s.DBConn(), s.Metrics(), false)
if err != nil {
log.Fatalf("failed to create store: %v", err)
}
@@ -147,7 +159,7 @@ func (s *BaseServer) EventStore() activity.Store {
func (s *BaseServer) APIHandler() http.Handler {
return Create(s, func() http.Handler {
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager())
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager(), nil)
if err != nil {
log.Fatalf("failed to create API handler: %v", err)
}
@@ -171,10 +183,10 @@ func (s *BaseServer) Router() *mux.Router {
})
}
func (s *BaseServer) RateLimiter() *middleware.APIRateLimiter {
return Create(s, func() *middleware.APIRateLimiter {
cfg, enabled := middleware.RateLimiterConfigFromEnv()
limiter := middleware.NewAPIRateLimiter(cfg)
func (s *BaseServer) RateLimiter() *ratelimit.APIRateLimiter {
return Create(s, func() *ratelimit.APIRateLimiter {
cfg, enabled := ratelimit.RateLimiterConfigFromEnv()
limiter := ratelimit.NewAPIRateLimiter(cfg)
limiter.SetEnabled(enabled)
return limiter
})
@@ -236,7 +248,7 @@ func (s *BaseServer) GRPCServer() *grpc.Server {
func (s *BaseServer) ReverseProxyGRPCServer() *nbgrpc.ProxyServiceServer {
return Create(s, func() *nbgrpc.ProxyServiceServer {
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.PKCEVerifierStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.SingleUseStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
s.AfterInit(func(s *BaseServer) {
proxyService.SetServiceManager(s.ServiceManager())
proxyService.SetActivityManager(s.ProxyActivityManager())
@@ -293,9 +305,9 @@ func (s *BaseServer) ProxyTokenStore() *nbgrpc.OneTimeTokenStore {
})
}
func (s *BaseServer) PKCEVerifierStore() *nbgrpc.PKCEVerifierStore {
return Create(s, func() *nbgrpc.PKCEVerifierStore {
return nbgrpc.NewPKCEVerifierStore(context.Background(), s.CacheStore())
func (s *BaseServer) SingleUseStore() *nbgrpc.SingleUseStore {
return Create(s, func() *nbgrpc.SingleUseStore {
return nbgrpc.NewSingleUseStore(context.Background(), s.CacheStore())
})
}
@@ -308,7 +320,7 @@ func (s *BaseServer) ProxyActivityManager() proxyactivity.Manager {
func (s *BaseServer) AccessLogsManager() accesslogs.Manager {
return Create(s, func() accesslogs.Manager {
accessLogManager := accesslogsmanager.NewManager(s.Store(), s.PermissionsManager(), s.GeoLocationManager())
accessLogManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(s.DBConn()), s.Store(), s.PermissionsManager(), s.GeoLocationManager())
accessLogManager.StartPeriodicCleanup(
context.Background(),
s.Config.ReverseProxy.AccessLogRetentionDays,
+1
View File
@@ -103,6 +103,7 @@ func (s *BaseServer) AccountManager() account.Manager {
s.AfterInit(func(s *BaseServer) {
accountManager.SetServiceManager(s.ServiceManager())
accountManager.AddAccountDeletionHook(s.AgentNetworkManager().RemoveAccountGateway)
})
return accountManager
+18
View File
@@ -23,6 +23,8 @@ import (
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/metrics"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/lifecycle"
"github.com/netbirdio/netbird/shared/profiling"
"github.com/netbirdio/netbird/util/wsproxy"
wsproxyserver "github.com/netbirdio/netbird/util/wsproxy/server"
"github.com/netbirdio/netbird/version"
@@ -36,6 +38,8 @@ const (
DefaultSelfHostedDomain = "netbird.selfhosted"
ContainerKeyBaseServer = "baseServer"
applicationName = "management"
)
type Server interface {
@@ -82,6 +86,8 @@ type BaseServer struct {
errCh chan error
wg sync.WaitGroup
cancel context.CancelFunc
lifecycle.StopHandlers
}
// Config holds the configuration parameters for creating a new server
@@ -117,6 +123,9 @@ func NewServer(cfg *Config) *BaseServer {
}
s.container[ContainerKeyBaseServer] = s
stopProfiling := profiling.Start(applicationName)
s.OnStop(stopProfiling)
return s
}
@@ -126,6 +135,14 @@ func (s *BaseServer) AfterInit(fn func(s *BaseServer)) {
// Start begins listening for HTTP requests on the configured address
func (s *BaseServer) Start(ctx context.Context) error {
if err := s.start(ctx); err != nil {
s.RunStopHandlers()
return err
}
return nil
}
func (s *BaseServer) start(ctx context.Context) error {
srvCtx, cancel := context.WithCancel(ctx)
s.cancel = cancel
s.errCh = make(chan error, 4)
@@ -278,6 +295,7 @@ func (s *BaseServer) setupTLS(ctx context.Context) (bool, error) {
func (s *BaseServer) Stop() error {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
defer s.RunStopHandlers()
if s.domainCleanupStop != nil {
s.domainCleanupStop()
}
+121
View File
@@ -0,0 +1,121 @@
package db
import (
"context"
"fmt"
"os"
"runtime"
"strconv"
"time"
"github.com/jackc/pgx/v5/pgxpool"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
)
const (
defaultTransactionTimeout = 5 * time.Minute
connMaxLifetime = time.Hour
connMaxIdleTime = 3 * time.Minute
)
// TxMetrics receives the duration of every committed top-level transaction.
type TxMetrics interface {
CountTransactionDuration(duration time.Duration)
}
// Conn is the database connection shared by all repositories: one gorm handle,
// the pgx pool of a Postgres deployment and the engine they talk to.
type Conn struct {
db *gorm.DB
pool *pgxpool.Pool
engine Engine
txTimeout time.Duration
metrics TxMetrics
}
// NewConn takes ownership of an open gorm handle and pool once it returns
// without error, applying the connection limits and transaction timeout
// configured through the environment.
func NewConn(ctx context.Context, gormDB *gorm.DB, engine Engine, pool *pgxpool.Pool) (*Conn, error) {
sqlDB, err := gormDB.DB()
if err != nil {
return nil, err
}
txTimeout := defaultTransactionTimeout
if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" {
if parsed, err := time.ParseDuration(v); err == nil {
txTimeout = parsed
}
}
log.WithContext(ctx).Infof("Setting transaction timeout to %v", txTimeout)
conns := runtime.NumCPU()
configuredConns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS"))
connsConfigured := err == nil
if connsConfigured {
conns = configuredConns
}
if engine == SqliteStoreEngine {
if connsConfigured {
log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1")
}
conns = 1
}
sqlDB.SetMaxOpenConns(conns)
sqlDB.SetMaxIdleConns(conns)
sqlDB.SetConnMaxLifetime(connMaxLifetime)
sqlDB.SetConnMaxIdleTime(connMaxIdleTime)
log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v",
conns, conns, connMaxLifetime, connMaxIdleTime)
return &Conn{db: gormDB, pool: pool, engine: engine, txTimeout: txTimeout}, nil
}
// DB returns the handle a query must run on: the transaction when tx is set,
// otherwise the shared connection.
func (c *Conn) DB(tx *Tx) *gorm.DB {
if tx != nil {
return tx.db
}
return c.db
}
// Pool returns the pgx pool for read paths that bypass gorm. It is nil on
// engines other than Postgres and inside a transaction, where the pool would
// not see the uncommitted writes.
func (c *Conn) Pool(tx *Tx) *pgxpool.Pool {
if tx != nil {
return nil
}
return c.pool
}
func (c *Conn) Engine() Engine {
return c.engine
}
// SetTxMetrics registers the sink that receives transaction durations.
func (c *Conn) SetTxMetrics(metrics TxMetrics) {
c.metrics = metrics
}
// AutoMigrate creates or updates the tables of the given models.
func (c *Conn) AutoMigrate(models ...any) error {
return c.db.AutoMigrate(models...)
}
// Close releases the gorm connection and the pgx pool.
func (c *Conn) Close() error {
if c.pool != nil {
c.pool.Close()
}
sqlDB, err := c.db.DB()
if err != nil {
return fmt.Errorf("get db: %w", err)
}
return sqlDB.Close()
}
+148
View File
@@ -0,0 +1,148 @@
package db
import (
"context"
"errors"
"path/filepath"
"testing"
"time"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
type testRow struct {
ID uint `gorm:"primaryKey"`
Name string
}
func openTestConn(t *testing.T) *Conn {
t.Helper()
conn, err := OpenSqliteFile(context.Background(), t.TempDir(), SqliteFileName)
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, conn.Close()) })
require.NoError(t, conn.AutoMigrate(&testRow{}))
return conn
}
func countRows(t *testing.T, conn *Conn) int64 {
t.Helper()
var count int64
require.NoError(t, conn.DB(nil).Model(&testRow{}).Count(&count).Error)
return count
}
func TestNewConn_ReadsTransactionTimeoutFromEnv(t *testing.T) {
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "1s")
conn := openTestConn(t)
assert.Equal(t, time.Second, conn.txTimeout)
assert.Equal(t, SqliteStoreEngine, conn.Engine())
}
func TestRunInTx_CommitsOnSuccess(t *testing.T) {
conn := openTestConn(t)
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
})
require.NoError(t, err)
assert.EqualValues(t, 1, countRows(t, conn))
}
func TestRunInTx_RollsBackOnError(t *testing.T) {
conn := openTestConn(t)
failure := errors.New("boom")
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
return failure
})
require.ErrorIs(t, err, failure)
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_RollsBackOnPanic(t *testing.T) {
conn := openTestConn(t)
require.Panics(t, func() {
_ = conn.RunInTx(context.Background(), func(tx *Tx) error {
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
panic("boom")
})
})
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_FailsWhenTimeoutExceeded(t *testing.T) {
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "50ms")
conn := openTestConn(t)
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
time.Sleep(100 * time.Millisecond)
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
})
require.ErrorIs(t, err, context.DeadlineExceeded)
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_ReportsDurationToMetrics(t *testing.T) {
conn := openTestConn(t)
metrics := &recordingMetrics{}
conn.SetTxMetrics(metrics)
require.NoError(t, conn.RunInTx(context.Background(), func(*Tx) error { return nil }))
assert.Equal(t, 1, metrics.calls)
}
func TestConn_DBSelectsTransactionHandle(t *testing.T) {
conn := openTestConn(t)
assert.Same(t, conn.db, conn.DB(nil))
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
assert.Same(t, tx.db, conn.DB(tx))
assert.NotSame(t, conn.db, conn.DB(tx))
return nil
})
require.NoError(t, err)
}
func TestConn_PoolIsUnavailableInsideTransaction(t *testing.T) {
conn := openTestConn(t)
conn.pool = &pgxpool.Pool{}
defer func() { conn.pool = nil }()
assert.Same(t, conn.pool, conn.Pool(nil))
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
assert.Nil(t, conn.Pool(tx))
return nil
})
require.NoError(t, err)
}
type recordingMetrics struct {
calls int
}
func (m *recordingMetrics) CountTransactionDuration(time.Duration) {
m.calls++
}
func TestNewConn_MaxOpenConnsFromEnv(t *testing.T) {
t.Setenv("NB_SQL_MAX_OPEN_CONNS", "7")
gormDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "store.db")), GormConfig())
require.NoError(t, err)
conn, err := NewConn(context.Background(), gormDB, PostgresStoreEngine, nil)
require.NoError(t, err)
t.Cleanup(func() { _ = conn.Close() })
sqlDB, err := conn.DB(nil).DB()
require.NoError(t, err)
assert.Equal(t, 7, sqlDB.Stats().MaxOpenConnections)
sqliteDB, err := openTestConn(t).DB(nil).DB()
require.NoError(t, err)
assert.Equal(t, 1, sqliteDB.Stats().MaxOpenConnections)
}
@@ -0,0 +1,23 @@
package dbtest
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
// NewConn opens a fresh SQLite database in a temporary directory, migrates the
// given models and closes the connection when the test ends. It ignores
// NB_STORE_ENGINE_SQLITE_FILE, so a developer's configured database is never
// touched, and is safe to call from parallel tests.
func NewConn(t testing.TB, models ...any) *db.Conn {
t.Helper()
conn, err := db.OpenSqliteFile(context.Background(), t.TempDir(), db.SqliteFileName)
require.NoError(t, err)
t.Cleanup(func() { _ = conn.Close() })
require.NoError(t, conn.AutoMigrate(models...))
return conn
}
@@ -0,0 +1,31 @@
package dbtest
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
func TestNewConn_IgnoresSqliteFileOverride(t *testing.T) {
override := filepath.Join(t.TempDir(), "configured.db")
t.Setenv("NB_STORE_ENGINE_SQLITE_FILE", override)
conn := NewConn(t)
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
_, err := os.Stat(override)
require.ErrorIs(t, err, os.ErrNotExist)
}
func TestNewConn_Parallel(t *testing.T) {
t.Parallel()
conn := NewConn(t)
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
}
+10
View File
@@ -0,0 +1,10 @@
package db
// Engine identifies the SQL engine behind a Conn.
type Engine string
const (
SqliteStoreEngine Engine = "sqlite"
PostgresStoreEngine Engine = "postgres"
MysqlStoreEngine Engine = "mysql"
)
+12
View File
@@ -0,0 +1,12 @@
package db
// LockingStrength is the row lock a query holds until its transaction ends.
type LockingStrength string
const (
LockingStrengthUpdate LockingStrength = "UPDATE"
LockingStrengthShare LockingStrength = "SHARE"
LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE"
LockingStrengthKeyShare LockingStrength = "KEY SHARE"
LockingStrengthNone LockingStrength = "NONE"
)
+176
View File
@@ -0,0 +1,176 @@
package db
import (
"context"
"fmt"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"time"
"github.com/jackc/pgx/v5/pgxpool"
"gorm.io/driver/mysql"
"gorm.io/driver/postgres"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// SqliteFileName is the default SQLite database file inside the data directory.
const SqliteFileName = "store.db"
// PoolConfig sizes the pgx pool a Postgres deployment uses for the read paths
// that bypass gorm.
type PoolConfig struct {
MaxConns int32
MinConns int32
MaxConnLifetime time.Duration
HealthCheckPeriod time.Duration
}
var DefaultPoolConfig = PoolConfig{
MaxConns: 30,
MinConns: 1,
MaxConnLifetime: 60 * time.Minute,
HealthCheckPeriod: time.Minute,
}
// GormConfig is the configuration every engine is opened with.
func GormConfig() *gorm.Config {
return &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
CreateBatchSize: 400,
}
}
// OpenSqlite opens the SQLite database in dataDir, or the file named by
// NB_STORE_ENGINE_SQLITE_FILE.
func OpenSqlite(ctx context.Context, dataDir string) (*Conn, error) {
storeFile := SqliteFileName
if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" {
storeFile = envFile
}
return OpenSqliteFile(ctx, dataDir, storeFile)
}
// OpenSqliteFile opens the SQLite database storeFile, resolved against dataDir
// when relative. storeFile may carry SQLite URI query parameters.
func OpenSqliteFile(ctx context.Context, dataDir, storeFile string) (*Conn, error) {
// Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc")
filePath, query, hasQuery := strings.Cut(storeFile, "?")
connStr := filePath
if !filepath.IsAbs(filePath) {
connStr = filepath.Join(dataDir, filePath)
}
// Compose query parameters. User-provided ?_busy_timeout (or its mattn alias
// ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at
// most that long on a lock instead of blocking the only Go-side connection.
// mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so
// the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared
// stays the default on non-Windows for the same reason as before.
parsed, _ := url.ParseQuery(query)
var defaults []string
if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" {
defaults = append(defaults, "_busy_timeout=30000")
}
if !hasQuery && runtime.GOOS != "windows" {
// To avoid `The process cannot access the file because it is being used by another process` on Windows
defaults = append(defaults, "cache=shared")
}
parts := defaults
if hasQuery {
parts = append(parts, query)
}
if len(parts) > 0 {
connStr += "?" + strings.Join(parts, "&")
}
gormDB, err := gorm.Open(sqlite.Open(connStr), GormConfig())
if err != nil {
return nil, err
}
conn, err := NewConn(ctx, gormDB, SqliteStoreEngine, nil)
if err != nil {
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
// OpenPostgres opens a Postgres database through gorm and a pgx pool sized by pool.
func OpenPostgres(ctx context.Context, dsn string, pool PoolConfig) (*Conn, error) {
gormDB, err := gorm.Open(postgres.Open(dsn), GormConfig())
if err != nil {
return nil, err
}
pgxPool, err := newPgxPool(ctx, dsn, pool)
if err != nil {
closeGorm(gormDB)
return nil, err
}
conn, err := NewConn(ctx, gormDB, PostgresStoreEngine, pgxPool)
if err != nil {
pgxPool.Close()
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
// MysqlDSN adds the connection parameters every MySQL handle needs, keeping
// the options already present in dsn.
func MysqlDSN(dsn string) string {
separator := "?"
if strings.Contains(dsn, "?") {
separator = "&"
}
return dsn + separator + "charset=utf8&parseTime=True&loc=Local"
}
// OpenMysql opens a MySQL database through gorm.
func OpenMysql(ctx context.Context, dsn string) (*Conn, error) {
gormDB, err := gorm.Open(mysql.Open(MysqlDSN(dsn)), GormConfig())
if err != nil {
return nil, err
}
conn, err := NewConn(ctx, gormDB, MysqlStoreEngine, nil)
if err != nil {
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
func newPgxPool(ctx context.Context, dsn string, cfg PoolConfig) (*pgxpool.Pool, error) {
config, err := pgxpool.ParseConfig(dsn)
if err != nil {
return nil, fmt.Errorf("unable to parse database config: %w", err)
}
config.MaxConns = cfg.MaxConns
config.MinConns = cfg.MinConns
config.MaxConnLifetime = cfg.MaxConnLifetime
config.HealthCheckPeriod = cfg.HealthCheckPeriod
pool, err := pgxpool.NewWithConfig(ctx, config)
if err != nil {
return nil, fmt.Errorf("unable to create connection pool: %w", err)
}
if err := pool.Ping(ctx); err != nil {
pool.Close()
return nil, fmt.Errorf("unable to ping database: %w", err)
}
return pool, nil
}
func closeGorm(gormDB *gorm.DB) {
if sqlDB, err := gormDB.DB(); err == nil {
_ = sqlDB.Close()
}
}
@@ -0,0 +1,12 @@
package db
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestMysqlDSN(t *testing.T) {
assert.Equal(t, "user:pw@tcp(host:3306)/db?charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db"))
assert.Equal(t, "user:pw@tcp(host:3306)/db?tls=true&charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db?tls=true"))
}
@@ -0,0 +1,105 @@
package db
import (
"context"
"errors"
"fmt"
"runtime/debug"
"time"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
)
// Tx is an open transaction handed to repository calls; nil means autocommit.
type Tx struct {
db *gorm.DB
}
// RunInTx runs fn in one transaction that commits when fn returns nil and rolls
// back otherwise, bounded by the configured transaction timeout.
func (c *Conn) RunInTx(ctx context.Context, fn func(tx *Tx) error) error {
timeoutCtx, cancel := context.WithTimeout(ctx, c.txTimeout)
defer cancel()
startTime := time.Now()
tx := c.db.WithContext(timeoutCtx).Begin()
if tx.Error != nil {
return tx.Error
}
defer func() {
if r := recover(); r != nil {
tx.Rollback()
panic(r)
}
}()
if err := c.applyStatementTimeouts(tx); err != nil {
tx.Rollback()
return err
}
err := c.withForeignKeyChecksDisabled(tx, func() error {
return fn(&Tx{db: tx})
})
if err != nil {
tx.Rollback()
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction", startTime)
return err
}
if err := tx.Commit().Error; err != nil {
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction commit", startTime)
return err
}
log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime))
if c.metrics != nil {
c.metrics.CountTransactionDuration(time.Since(startTime))
}
return nil
}
func (c *Conn) applyStatementTimeouts(tx *gorm.DB) error {
if c.engine != PostgresStoreEngine {
return nil
}
if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil {
return fmt.Errorf("failed to set statement timeout: %w", err)
}
if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil {
return fmt.Errorf("failed to set lock timeout: %w", err)
}
return nil
}
// withForeignKeyChecksDisabled runs fn with MySQL's FK checks off, which avoids
// deadlocks on MySQL and Aurora without needing SUPER privilege. The setting is
// session-scoped and survives a rollback, so it is turned back on whenever fn
// returns or panics; otherwise the pooled connection would keep it disabled.
func (c *Conn) withForeignKeyChecksDisabled(tx *gorm.DB, fn func() error) (err error) {
if c.engine != MysqlStoreEngine {
return fn()
}
if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil {
return fmt.Errorf("failed to disable FK checks: %w", err)
}
defer func() {
restoreErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error
if restoreErr == nil {
return
}
if err == nil {
err = fmt.Errorf("failed to re-enable FK checks: %w", restoreErr)
return
}
log.WithContext(tx.Statement.Context).Warnf("failed to re-enable FK checks after failed transaction: %v", restoreErr)
}()
return fn()
}
func (c *Conn) logIfTimedOut(ctx, timeoutCtx context.Context, err error, phase string, startTime time.Time) {
if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) {
log.WithContext(ctx).Warnf("%s exceeded %s timeout after %v, stack: %s", phase, c.txTimeout, time.Since(startTime), debug.Stack())
}
}
@@ -14,6 +14,7 @@ const (
baseBlockDuration = 10 * time.Minute // Duration for which a peer is banned after exceeding the reconnection limit
reconnLimitForBan = 30 // Number of reconnections within the reconnTreshold that triggers a ban
metaChangeLimit = 5 // Number of reconnections with different metadata that triggers a ban of one peer
maxBanLevel = 6 // Highest ban level; the ban duration doubles per level up to this one
)
type lfConfig struct {
@@ -21,6 +22,7 @@ type lfConfig struct {
baseBlockDuration time.Duration
reconnLimitForBan int
metaChangeLimit int
maxBanLevel int
}
func initCfg() *lfConfig {
@@ -29,6 +31,7 @@ func initCfg() *lfConfig {
baseBlockDuration: baseBlockDuration,
reconnLimitForBan: reconnLimitForBan,
metaChangeLimit: metaChangeLimit,
maxBanLevel: maxBanLevel,
}
}
@@ -102,11 +105,18 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
return
}
if state.isBanned && now.After(state.banExpiresAt) {
if state.isBanned {
if now.Before(state.banExpiresAt) {
return
}
state.isBanned = false
}
if state.banLevel > 0 && now.Sub(state.lastSeen) > (2*l.cfg.baseBlockDuration) {
quietSince := state.lastSeen
if state.banExpiresAt.After(quietSince) {
quietSince = state.banExpiresAt
}
if state.banLevel > 0 && now.Sub(quietSince) > (2*l.cfg.baseBlockDuration) {
state.banLevel = 0
}
@@ -124,10 +134,17 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
return
}
if now.Sub(state.sessionStart) >= l.cfg.reconnThreshold {
state.sessionStart = now
state.sessionCounter = 0
}
state.sessionCounter++
if state.sessionCounter > l.cfg.reconnLimitForBan && now.Sub(state.sessionStart) < l.cfg.reconnThreshold {
if state.sessionCounter > l.cfg.reconnLimitForBan {
state.isBanned = true
state.banLevel++
if state.banLevel < l.cfg.maxBanLevel {
state.banLevel++
}
backoffFactor := math.Pow(2, float64(state.banLevel-1))
duration := time.Duration(float64(l.cfg.baseBlockDuration) * backoffFactor)
@@ -20,6 +20,7 @@ func testAdvancedCfg() *lfConfig {
baseBlockDuration: 100 * time.Millisecond,
reconnLimitForBan: 3,
metaChangeLimit: 2,
maxBanLevel: 3,
}
}
@@ -157,6 +158,187 @@ func (s *LoginFilterTestSuite) TestMetaChangeIsAllowedAfterWindowResets() {
s.Equal(1, s.filter.logged[pubKey].metaChangeCounter, "meta change counter should reset")
}
func (s *LoginFilterTestSuite) TestReconnectStormAfterQuietPeriodTriggersBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
s.filter.addLogin(pubKey, meta)
s.Require().Contains(s.filter.logged, pubKey)
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
s.filter.addLogin(pubKey, meta)
s.Equal(1, s.filter.logged[pubKey].sessionCounter, "expired window should restart the count")
for i := 1; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
s.filter.addLogin(pubKey, meta)
s.False(s.filter.allowLogin(pubKey, meta))
s.True(s.filter.logged[pubKey].isBanned)
}
func (s *LoginFilterTestSuite) TestReconnectStormAfterBanExpiresTriggersBanAgain() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.Require().True(s.filter.logged[pubKey].isBanned)
expired := time.Now().Add(-(s.filter.cfg.baseBlockDuration + time.Second))
s.filter.logged[pubKey].banExpiresAt = expired
s.filter.logged[pubKey].sessionStart = expired
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(2, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestSlowReconnectsAcrossWindowsDoNotBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
for i := 0; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
}
func (s *LoginFilterTestSuite) TestBanLevelEscalatesWhenStormResumesRightAfterBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
banTime := time.Now().Add(-3 * s.filter.cfg.baseBlockDuration)
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
isBanned: true,
banLevel: 1,
banExpiresAt: time.Now().Add(-time.Millisecond),
sessionStart: banTime,
lastSeen: banTime,
}
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(2, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestBanLevelResetsAfterQuietPeriodFollowingBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
quiet := 2*s.filter.cfg.baseBlockDuration + time.Second
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
banLevel: 2,
banExpiresAt: time.Now().Add(-s.filter.cfg.baseBlockDuration),
lastSeen: time.Now().Add(-2 * quiet),
}
s.filter.addLogin(pubKey, meta)
s.Equal(2, s.filter.logged[pubKey].banLevel, "ban ended more recently than the quiet period")
s.filter.logged[pubKey].banExpiresAt = time.Now().Add(-quiet)
s.filter.logged[pubKey].lastSeen = time.Now().Add(-2 * quiet)
s.filter.addLogin(pubKey, meta)
s.Equal(0, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestBanDurationIsCappedAtMaxLevel() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
maxLevel := s.filter.cfg.maxBanLevel
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
banLevel: maxLevel,
sessionStart: time.Now(),
lastSeen: time.Now(),
}
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(maxLevel, s.filter.logged[pubKey].banLevel)
expected := s.filter.cfg.baseBlockDuration << (maxLevel - 1)
s.InDelta(expected, s.filter.logged[pubKey].banExpiresAt.Sub(s.filter.logged[pubKey].lastSeen), float64(time.Millisecond))
}
func (s *LoginFilterTestSuite) TestEstablishedPeerReconnectingOnceIsAllowed() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
longAgo := time.Now().Add(-time.Hour)
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
sessionCounter: 1,
sessionStart: longAgo,
lastSeen: longAgo,
metaChangeWindowStart: longAgo,
metaChangeCounter: 1,
}
s.True(s.filter.allowLogin(pubKey, meta))
s.filter.addLogin(pubKey, meta)
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
s.Equal(1, s.filter.logged[pubKey].sessionCounter)
}
func (s *LoginFilterTestSuite) TestLoginsDuringActiveBanDoNotExtendIt() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.Require().True(s.filter.logged[pubKey].isBanned)
expiresAt := time.Now().Add(time.Hour)
s.filter.logged[pubKey].banExpiresAt = expiresAt
lastSeen := s.filter.logged[pubKey].lastSeen
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(1, s.filter.logged[pubKey].banLevel)
s.Equal(expiresAt, s.filter.logged[pubKey].banExpiresAt)
s.Equal(lastSeen, s.filter.logged[pubKey].lastSeen)
s.Equal(0, s.filter.logged[pubKey].sessionCounter)
}
func BenchmarkHashingMethods(b *testing.B) {
meta := nbpeer.PeerSystemMeta{
WtVersion: "1.25.1",
@@ -1,55 +0,0 @@
package grpc
import (
"context"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/store"
log "github.com/sirupsen/logrus"
nbcache "github.com/netbirdio/netbird/management/server/cache"
)
// PKCEVerifierStore manages PKCE verifiers for OAuth flows.
// Supports both in-memory and Redis storage via NB_IDP_CACHE_REDIS_ADDRESS env var.
type PKCEVerifierStore struct {
cache nbcache.Store
ctx context.Context
}
// NewPKCEVerifierStore creates a PKCE verifier store using the provided shared cache store.
func NewPKCEVerifierStore(ctx context.Context, cacheStore nbcache.Store) *PKCEVerifierStore {
return &PKCEVerifierStore{
cache: cacheStore,
ctx: ctx,
}
}
// Store saves a PKCE verifier associated with an OAuth state parameter.
// The verifier is stored with the specified TTL and will be automatically deleted after expiration.
func (s *PKCEVerifierStore) Store(state, verifier string, ttl time.Duration) error {
if err := s.cache.Set(s.ctx, state, verifier, store.WithExpiration(ttl)); err != nil {
return fmt.Errorf("failed to store PKCE verifier: %w", err)
}
log.Debugf("Stored PKCE verifier for state (expires in %s)", ttl)
return nil
}
// LoadAndDelete retrieves and removes a PKCE verifier for the given state.
// Returns the verifier and true if found, or empty string and false if not found.
// This enforces single-use semantics for PKCE verifiers.
func (s *PKCEVerifierStore) LoadAndDelete(state string) (string, bool) {
verifier, found, err := s.cache.GetDel(s.ctx, state)
if err != nil {
log.Warnf("Failed to consume PKCE verifier: %v", err)
return "", false
}
if !found {
log.Debug("PKCE verifier not found for state")
return "", false
}
return verifier, true
}
+74 -27
View File
@@ -27,8 +27,6 @@ import (
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/peers"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
@@ -42,6 +40,7 @@ import (
"github.com/netbirdio/netbird/management/server/users"
proxyauth "github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/shared/hash/argon2id"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/proto"
nbstatus "github.com/netbirdio/netbird/shared/management/status"
)
@@ -102,7 +101,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.
@@ -141,8 +141,8 @@ type ProxyServiceServer struct {
// OIDC configuration for proxy authentication
oidcConfig ProxyOIDCConfig
// Store for PKCE verifiers
pkceVerifierStore *PKCEVerifierStore
// singleUseStore backs both PKCE verifiers and OIDC session exchange codes.
singleUseStore *SingleUseStore
// tokenTTL is the lifetime of one-time tokens generated for proxy
// authentication. Defaults to defaultProxyTokenTTL when zero.
@@ -157,6 +157,13 @@ type ProxyServiceServer struct {
const pkceVerifierTTL = 10 * time.Minute
const sessionCodeTTL = 60 * time.Second
const sessionCodeCacheNamespace = "proxy:session"
// The signed nonce binds the handoff mode without changing the state format.
const sessionCodeNoncePrefix = "code."
const defaultProxyTokenTTL = 5 * time.Minute
const defaultSnapshotBatchSize = 500
@@ -207,13 +214,13 @@ func enforceAccountScope(ctx context.Context, requestAccountID string) error {
}
// NewProxyServiceServer creates a new proxy service server.
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, pkceStore *PKCEVerifierStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, singleUseStore *SingleUseStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
ctx, cancel := context.WithCancel(context.Background())
s := &ProxyServiceServer{
accessLogManager: accessLogMgr,
oidcConfig: oidcConfig,
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
peersManager: peersManager,
usersManager: usersManager,
idpManager: idpManager,
@@ -242,9 +249,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.
@@ -304,6 +312,16 @@ func (s *ProxyServiceServer) proxyConnectAuthorizer() ProxyConnectAuthorizer {
return s.connectAuthorizer
}
// GenerateSessionCode creates a single-use code for the given session token.
func (s *ProxyServiceServer) GenerateSessionCode(sessionToken string) (code string, ok bool) {
code, err := s.singleUseStore.Generate(sessionCodeCacheNamespace, sessionToken, sessionCodeTTL)
if err != nil {
log.WithError(err).Error("failed to generate proxy session code")
return "", false
}
return code, true
}
// CheckLLMPolicyLimits is the pre-flight policy gate the proxy calls before
// forwarding an LLM request upstream. Delegates to the agent-network selector,
// which scores applicable policies by remaining headroom and returns the
@@ -412,6 +430,7 @@ func (s *ProxyServiceServer) SetProxyController(proxyController proxy.Controller
type proxyConnectParams struct {
proxyID string
address string
version string
capabilities *proto.ProxyCapabilities
}
@@ -422,6 +441,7 @@ func (s *ProxyServiceServer) GetMappingUpdate(req *proto.GetMappingUpdateRequest
return err
}
params.capabilities = req.GetCapabilities()
params.version = req.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
stream: stream,
@@ -455,6 +475,7 @@ func (s *ProxyServiceServer) SyncMappings(stream proto.ProxyService_SyncMappings
return err
}
params.capabilities = init.GetCapabilities()
params.version = init.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
syncStream: stream,
@@ -566,7 +587,7 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params
}
}
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, accountID, caps)
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, params.version, accountID, caps)
if err != nil {
cancel()
if accountID != nil {
@@ -1223,6 +1244,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 +1256,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
@@ -1522,18 +1552,20 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
log.WithContext(ctx).Errorf("failed to get account services: %v", err)
return nil, status.Errorf(codes.FailedPrecondition, "get account services: %v", err)
}
var found bool
var matchedService *rpservice.Service
for _, service := range services {
if service.Domain == redirectURL.Hostname() {
found = true
matchedService = service
break
}
}
if !found {
if matchedService == nil {
log.WithContext(ctx).Debugf("OIDC redirect URL %q does not match any service domain", redirectURL.Hostname())
return nil, status.Errorf(codes.FailedPrecondition, "service not found in store")
}
useSessionCode := s.proxyManager.ClusterSupportsSessionCode(ctx, matchedService.ProxyCluster)
provider, err := oidc.NewProvider(ctx, s.oidcConfig.Issuer)
if err != nil {
log.WithContext(ctx).Errorf("failed to create OIDC provider: %v", err)
@@ -1553,15 +1585,18 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
return nil, status.Errorf(codes.Internal, "generate nonce: %v", err)
}
nonceB64 := base64.URLEncoding.EncodeToString(nonce)
if useSessionCode {
nonceB64 = sessionCodeNoncePrefix + nonceB64
}
// Using an HMAC here to avoid redirection state being modified.
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
payload := redirectURL.String() + "|" + nonceB64
hmacSum := s.generateHMAC(payload)
state := fmt.Sprintf("%s|%s|%s", base64.URLEncoding.EncodeToString([]byte(redirectURL.String())), nonceB64, hmacSum)
codeVerifier := oauth2.GenerateVerifier()
if err := s.pkceVerifierStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil {
if err := s.singleUseStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil {
log.WithContext(ctx).Errorf("failed to store PKCE verifier: %v", err)
return nil, status.Errorf(codes.Internal, "store PKCE verifier: %v", err)
}
@@ -1598,15 +1633,12 @@ func (s *ProxyServiceServer) generateHMAC(input string) string {
return hex.EncodeToString(mac.Sum(nil))
}
// ValidateState validates the state parameter from an OAuth callback.
// Returns the original redirect URL if valid, or an error if invalid.
// The HMAC is verified before consuming the PKCE verifier to prevent
// an attacker from invalidating a legitimate user's auth flow.
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, err error) {
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
// ValidateState validates and consumes an OIDC state.
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, useSessionCode bool, err error) {
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
parts := strings.Split(state, "|")
if len(parts) != 3 {
return "", "", errors.New("invalid state format")
return "", "", false, errors.New("invalid state format")
}
encodedURL := parts[0]
@@ -1615,7 +1647,7 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
redirectURLBytes, err := base64.URLEncoding.DecodeString(encodedURL)
if err != nil {
return "", "", fmt.Errorf("invalid state encoding: %w", err)
return "", "", false, fmt.Errorf("invalid state encoding: %w", err)
}
redirectURL = string(redirectURLBytes)
@@ -1623,16 +1655,17 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
expectedHMAC := s.generateHMAC(payload)
if !hmac.Equal([]byte(providedHMAC), []byte(expectedHMAC)) {
return "", "", errors.New("invalid state signature")
return "", "", false, errors.New("invalid state signature")
}
useSessionCode = strings.HasPrefix(nonce, sessionCodeNoncePrefix)
// Consume the PKCE verifier only after HMAC validation passes.
verifier, ok := s.pkceVerifierStore.LoadAndDelete(state)
verifier, ok := s.singleUseStore.LoadAndDelete(state)
if !ok {
return "", "", errors.New("no verifier for state")
return "", "", false, errors.New("no verifier for state")
}
return verifier, redirectURL, nil
return verifier, redirectURL, useSessionCode, nil
}
// Denied reasons reported to the proxy when access is refused because of the
@@ -1837,7 +1870,20 @@ func (s *ProxyServiceServer) getAccountServiceByDomain(ctx context.Context, acco
// ValidateSession validates a session token and checks if the user has access to the domain.
func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.ValidateSessionRequest) (*proto.ValidateSessionResponse, error) {
domain := req.GetDomain()
sessionToken := req.GetSessionToken()
sessionToken := req.GetSessionToken() //nolint:staticcheck
// A one-time code from the OIDC callback is redeemed here for the durable
// token, so the token never travels in a redirect URL. The redeemed token
// is returned to the proxy (mintedToken) to install as the session cookie.
mintedToken := ""
if code := req.GetSessionCode(); code != "" {
redeemed, found := s.singleUseStore.LoadAndDelete(singleUseCacheKey(sessionCodeCacheNamespace, code))
if !found {
return deniedSessionResponse("invalid or expired session code"), nil
}
sessionToken = redeemed
mintedToken = redeemed
}
if domain == "" || sessionToken == "" {
return deniedSessionResponse("missing domain or session_token"), nil
@@ -1910,6 +1956,7 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
UserEmail: user.Email,
PeerGroupIds: groupIDs,
PeerGroupNames: groupNames,
SessionToken: mintedToken,
}, nil
}
@@ -0,0 +1,93 @@
package grpc
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/shared/management/proto"
)
const (
versionTestProxyID = "proxy-a"
versionTestCluster = "cluster.example.com"
versionTestVersion = "0.60.0"
)
// hangupStream cancels its context on the first Send, emulating a proxy that
// disconnects right after receiving the initial snapshot. The legacy stream
// carries no proxy-to-management messages, so this is the only way for
// GetMappingUpdate to return.
type hangupStream struct {
recordingStream
ctx context.Context
cancel context.CancelFunc
}
func (s *hangupStream) Send(m *proto.GetMappingUpdateResponse) error {
s.cancel()
return s.recordingStream.Send(m)
}
func (s *hangupStream) Context() context.Context { return s.ctx }
// newVersionTestServer wires a server whose proxy manager only accepts a
// Connect carrying versionTestVersion, so a dropped or mangled version fails
// the test as an unexpected call.
func newVersionTestServer(t *testing.T) *ProxyServiceServer {
t.Helper()
ctrl := gomock.NewController(t)
svcMgr := rpservice.NewMockManager(ctrl)
svcMgr.EXPECT().GetGlobalServices(gomock.Any()).Return(nil, nil)
proxyMgr := proxy.NewMockManager(ctrl)
proxyMgr.EXPECT().
Connect(gomock.Any(), versionTestProxyID, gomock.Any(), versionTestCluster, gomock.Any(), versionTestVersion, gomock.Any(), gomock.Any()).
Return(&proxy.Proxy{ID: versionTestProxyID, Version: versionTestVersion}, nil)
proxyMgr.EXPECT().Disconnect(gomock.Any(), versionTestProxyID, gomock.Any()).Return(nil)
s := newSnapshotTestServer(t, 10)
s.serviceManager = svcMgr
s.proxyManager = proxyMgr
return s
}
func TestSyncMappings_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
// The init carries the version, the ack acknowledges the empty snapshot,
// and the exhausted fake stream then ends the RPC.
stream := &syncRecordingStream{
recvMsgs: []*proto.SyncMappingsRequest{
{Msg: &proto.SyncMappingsRequest_Init{Init: &proto.SyncMappingsInit{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}}},
ackMsg(),
},
}
err := s.SyncMappings(stream)
require.ErrorContains(t, err, "no more recv messages")
}
func TestGetMappingUpdate_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
stream := &hangupStream{ctx: ctx, cancel: cancel}
err := s.GetMappingUpdate(&proto.GetMappingUpdateRequest{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}, stream)
require.ErrorIs(t, err, context.Canceled)
}
@@ -0,0 +1,101 @@
package grpc
import (
"sync"
"time"
"golang.org/x/time/rate"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/types/known/durationpb"
)
const (
credentialVerificationInterval = 6 * time.Second
credentialVerificationBurst = 5
credentialVerificationMaxServices = 4096
credentialVerificationIdleTimeout = 15 * time.Minute
credentialVerificationCleanupInterval = time.Minute
)
type credentialAccountID string
type credentialServiceID string
type credentialVerificationKey struct {
accountID credentialAccountID
serviceID credentialServiceID
}
type credentialVerificationBudget struct {
limiter *rate.Limiter
lastUsed time.Time
}
// The zero value is ready to use. Budgets are local to this Management process;
// proxy replicas reaching this process share a service's verification budget.
type credentialVerificationLimiter struct {
mu sync.Mutex
now func() time.Time
services map[credentialVerificationKey]*credentialVerificationBudget
nextCleanup time.Time
closed bool
}
func (l *credentialVerificationLimiter) allow(key credentialVerificationKey) error {
l.mu.Lock()
defer l.mu.Unlock()
if l.closed {
return status.Error(codes.Unavailable, "credential verification is closed")
}
now := time.Now()
if l.now != nil {
now = l.now()
}
l.cleanup(now)
budget := l.services[key]
if budget == nil {
if len(l.services) >= credentialVerificationMaxServices {
return credentialVerificationThrottled(credentialVerificationCleanupInterval)
}
if l.services == nil {
l.services = make(map[credentialVerificationKey]*credentialVerificationBudget)
}
budget = &credentialVerificationBudget{limiter: rate.NewLimiter(rate.Every(credentialVerificationInterval), credentialVerificationBurst)}
l.services[key] = budget
}
budget.lastUsed = now
if budget.limiter.AllowN(now, 1) {
return nil
}
delay := max(time.Nanosecond, time.Duration((1-budget.limiter.TokensAt(now))*float64(credentialVerificationInterval)))
return credentialVerificationThrottled(delay)
}
func (l *credentialVerificationLimiter) cleanup(now time.Time) {
if now.Before(l.nextCleanup) {
return
}
l.nextCleanup = now.Add(credentialVerificationCleanupInterval)
for key, budget := range l.services {
if now.Sub(budget.lastUsed) >= credentialVerificationIdleTimeout {
delete(l.services, key)
}
}
}
func (l *credentialVerificationLimiter) close() {
l.mu.Lock()
defer l.mu.Unlock()
l.closed = true
l.services = nil
}
func credentialVerificationThrottled(delay time.Duration) error {
s := status.New(codes.ResourceExhausted, "too many credential verification attempts")
withRetry, err := s.WithDetails(&errdetails.RetryInfo{RetryDelay: durationpb.New(delay)})
if err != nil {
return s.Err()
}
return withRetry.Err()
}
@@ -0,0 +1,79 @@
package grpc
import (
"strconv"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
func TestCredentialVerificationRefillAndIsolation(t *testing.T) {
now := time.Now()
l := credentialVerificationLimiter{now: func() time.Time { return now }}
key := credentialVerificationKey{accountID: "account", serviceID: "service"}
for range credentialVerificationBurst {
require.NoError(t, l.allow(key))
}
err := l.allow(key)
require.Equal(t, codes.ResourceExhausted, status.Code(err), "the burst must be bounded")
now = now.Add(3 * time.Second)
err = l.allow(key)
require.Equal(t, codes.ResourceExhausted, status.Code(err), "a partially refilled token must not permit a check")
details := status.Convert(err).Details()
require.Len(t, details, 1, "throttling must provide RetryInfo")
retry, ok := details[0].(*errdetails.RetryInfo)
require.True(t, ok, "retry details must use the standard message")
assert.Equal(t, 3*time.Second, retry.RetryDelay.AsDuration(), "retry hint must reflect time until the next check")
now = now.Add(3 * time.Second)
require.NoError(t, l.allow(key))
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "only one check must refill every six seconds")
require.NoError(t, l.allow(credentialVerificationKey{accountID: "other-account", serviceID: key.serviceID}))
require.NoError(t, l.allow(credentialVerificationKey{accountID: key.accountID, serviceID: "other-service"}))
}
func TestCredentialVerificationCapacityAndExpiry(t *testing.T) {
now := time.Now()
l := credentialVerificationLimiter{now: func() time.Time { return now }}
for i := range credentialVerificationMaxServices {
require.NoError(t, l.allow(credentialVerificationKey{accountID: "account", serviceID: credentialServiceID(strconv.Itoa(i))}))
}
key := credentialVerificationKey{accountID: "account", serviceID: "new-service"}
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "capacity exhaustion must deny new checks")
now = now.Add(credentialVerificationIdleTimeout)
for range credentialVerificationBurst {
require.NoError(t, l.allow(key))
}
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "expiry must retain the normal burst bound")
}
func TestCredentialVerificationConcurrentChecksAndClose(t *testing.T) {
var l credentialVerificationLimiter
key := credentialVerificationKey{accountID: "account", serviceID: "service"}
var admitted atomic.Int32
var wg sync.WaitGroup
for range 100 {
wg.Go(func() {
if err := l.allow(key); err == nil {
admitted.Add(1)
} else {
assert.Equal(t, codes.ResourceExhausted, status.Code(err), "excess checks must be throttled")
}
})
}
wg.Wait()
assert.EqualValues(t, credentialVerificationBurst, admitted.Load(), "concurrent checks must share the burst")
for range 10 {
wg.Go(l.close)
wg.Go(func() { assert.Error(t, l.allow(key)) })
}
wg.Wait()
assert.Empty(t, l.services, "closing must release retained budgets")
assert.Equal(t, codes.Unavailable, status.Code(l.allow(key)), "checks after close must fail closed")
}
@@ -0,0 +1,18 @@
# Reverse proxy credential verification
The `ProxyService.Authenticate` RPC limits PIN and password checks before
verifying their Argon2 hashes. Both methods share one budget per account and
service: a burst of five checks, replenishing one check every six seconds
(ten per minute). Successful and failed checks consume the budget. Account
scope and service lookup run before the limiter.
Excess checks receive gRPC `ResourceExhausted` with a standard `RetryInfo` delay.
Updated proxies translate it to HTTP 429 and `Retry-After`. Older proxies show
an authentication-service error but cannot bypass the Management limit.
Budgets are held in memory per Management process and reset on restart. Proxy
replicas reaching the same Management process share its budgets. Multiple
Management processes have independent budgets; this is not a cluster-wide
limit. At most 4,096 service budgets are retained, with idle entries expiring
after fifteen minutes. Capacity exhaustion denies new checks until entries
expire. Closing the server releases the retained state.
@@ -0,0 +1,131 @@
package grpc_test
import (
"context"
"net"
"net/netip"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/peer"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
servicemanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service/manager"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/proto"
)
func credentialServer(t *testing.T) (*nbgrpc.ProxyServiceServer, context.Context, grpc.UnaryServerInterceptor) {
t.Helper()
ctx := context.Background()
s, err := store.NewStore(ctx, types.SqliteStoreEngine, t.TempDir(), nil, false)
require.NoError(t, err)
t.Cleanup(func() { assert.NoError(t, s.Close(ctx)) })
require.NoError(t, s.SaveAccount(ctx, &types.Account{Id: "account"}))
keys, err := sessionkey.GenerateKeyPair()
require.NoError(t, err)
for _, id := range []string{"service", "other-service"} {
svc := &service.Service{
ID: id, AccountID: "account", Name: id, Domain: id + ".example.com",
Enabled: true, SessionPrivateKey: keys.PrivateKey, SessionPublicKey: keys.PublicKey,
Auth: service.AuthConfig{
PinAuth: &service.PINAuthConfig{Enabled: true, Pin: "842716"},
PasswordAuth: &service.PasswordAuthConfig{Enabled: true, Password: "test-password"},
},
}
require.NoError(t, svc.Auth.HashSecrets())
require.NoError(t, s.CreateService(ctx, svc))
}
account := "account"
token, err := types.CreateNewProxyAccessToken("test proxy", time.Hour, &account, "admin")
require.NoError(t, err)
require.NoError(t, s.SaveProxyAccessToken(ctx, &token.ProxyAccessToken))
ctx = metadata.NewIncomingContext(ctx, metadata.Pairs("authorization", "Bearer "+string(token.PlainToken)))
ctx = peer.NewContext(ctx, &peer.Peer{Addr: net.TCPAddrFromAddrPort(netip.MustParseAddrPort("192.0.2.1:443"))})
server := nbgrpc.NewProxyServiceServer(nil, nil, nil, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
t.Cleanup(server.Close)
server.SetServiceManager(servicemanager.NewManager(s, nil, nil, nil, nil, nil))
interceptor, _, closeInterceptor := nbgrpc.NewProxyAuthInterceptors(s)
t.Cleanup(closeInterceptor)
return server, ctx, interceptor
}
func TestAuthenticateCredentialRateLimit(t *testing.T) {
server, ctx, interceptor := credentialServer(t)
authenticate := func(req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
response, err := interceptor(ctx, req, &grpc.UnaryServerInfo{FullMethod: "/management.ProxyService/Authenticate"}, func(ctx context.Context, req any) (any, error) {
return server.Authenticate(ctx, req.(*proto.AuthenticateRequest))
})
if err != nil {
return nil, err
}
return response.(*proto.AuthenticateResponse), nil
}
for i := range 5 {
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service"}
if i%2 == 0 {
req.Request = &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}
} else {
req.Request = &proto.AuthenticateRequest_Password{Password: &proto.PasswordRequest{Password: "wrong-password"}}
}
resp, err := authenticate(req)
require.NoError(t, err)
assert.False(t, resp.GetSuccess(), "incorrect PINs and passwords must be denied")
assert.Empty(t, resp.GetSessionToken(), "incorrect credentials must not issue a token")
}
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "842716"}}}
resp, err := authenticate(req)
assert.Nil(t, resp, "a throttled verification must not return a session")
require.Equal(t, codes.ResourceExhausted, status.Code(err), "PIN and password checks must share a service budget even with a valid proxy token")
details := status.Convert(err).Details()
require.Len(t, details, 1, "throttled responses must include a retry hint")
retry, ok := details[0].(*errdetails.RetryInfo)
require.True(t, ok, "the hint must use the standard RetryInfo message")
assert.Positive(t, retry.RetryDelay.AsDuration(), "the retry delay must be positive")
assert.LessOrEqual(t, retry.RetryDelay.AsDuration(), 6*time.Second, "the service must replenish one verification every six seconds")
req.AccountId = "another-account"
_, err = authenticate(req)
assert.Equal(t, codes.PermissionDenied, status.Code(err), "account scope must still be enforced before throttling")
req.AccountId = "account"
req.Id = "other-service"
resp, err = authenticate(req)
require.NoError(t, err)
assert.True(t, resp.GetSuccess(), "one service's throttle must not block another service")
assert.NotEmpty(t, resp.GetSessionToken(), "valid credentials on another service must issue a session")
}
func TestAuthenticateCredentialConcurrentLimit(t *testing.T) {
server, _, _ := credentialServer(t)
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}}
var checked, throttled atomic.Int32
var wg sync.WaitGroup
for range 20 {
wg.Go(func() {
resp, err := server.Authenticate(context.Background(), req)
switch status.Code(err) {
case codes.OK:
checked.Add(1)
assert.False(t, resp.GetSuccess(), "incorrect credentials must be denied")
case codes.ResourceExhausted:
throttled.Add(1)
default:
assert.NoError(t, err)
}
})
}
wg.Wait()
assert.EqualValues(t, 5, checked.Load(), "only the burst budget may reach concurrent credential verification")
assert.EqualValues(t, 15, throttled.Load(), "excess concurrent checks must be throttled")
}
+43 -19
View File
@@ -129,11 +129,11 @@ func drainEmpty(ch chan *proto.GetMappingUpdateResponse) bool {
func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -186,11 +186,11 @@ func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) {
func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -220,11 +220,11 @@ func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
func TestSendServiceUpdate_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -272,13 +272,13 @@ func generateState(s *ProxyServiceServer, redirectURL string) string {
func TestOAuthState_NeverTheSame(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
redirectURL := "https://app.example.com/callback"
@@ -300,20 +300,20 @@ func TestOAuthState_NeverTheSame(t *testing.T) {
func TestValidateState_RejectsOldTwoPartFormat(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
// Old format had only 2 parts: base64(url)|hmac
err := s.pkceVerifierStore.Store("base64url|hmac", "test", 10*time.Minute)
err := s.singleUseStore.Store("base64url|hmac", "test", 10*time.Minute)
require.NoError(t, err)
_, _, err = s.ValidateState("base64url|hmac")
_, _, _, err = s.ValidateState("base64url|hmac")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid state format")
}
@@ -372,24 +372,48 @@ func TestEnforceAccountScope_AllowsNoTokenInContext(t *testing.T) {
func TestValidateState_RejectsInvalidHMAC(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
// Store with tampered HMAC
err := s.pkceVerifierStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
err := s.singleUseStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
require.NoError(t, err)
_, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
_, _, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid state signature")
}
func TestSessionCodeCannotConsumeOIDCState(t *testing.T) {
const verifier = "pkce-verifier"
store := NewSingleUseStore(context.Background(), testCacheStore(t))
server := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
singleUseStore: store,
}
state := generateState(server, "https://service.example.com/callback")
require.NoError(t, store.Store(state, verifier, time.Minute))
response, err := server.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
SessionCode: state,
})
require.NoError(t, err)
assert.False(t, response.GetValid())
gotVerifier, _, _, err := server.ValidateState(state)
require.NoError(t, err)
assert.Equal(t, verifier, gotVerifier)
}
func TestSendServiceUpdateToCluster_FiltersOnCapability(t *testing.T) {
tokenStore := NewOneTimeTokenStore(context.Background(), testCacheStore(t))
@@ -0,0 +1,67 @@
package grpc
import (
"context"
"crypto/rand"
"encoding/base64"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/store"
log "github.com/sirupsen/logrus"
nbcache "github.com/netbirdio/netbird/management/server/cache"
)
// SingleUseStore stores short-lived values that can be retrieved only once.
type SingleUseStore struct {
cache nbcache.Store
ctx context.Context
}
// NewSingleUseStore creates a single-use value store over the shared cache.
func NewSingleUseStore(ctx context.Context, cacheStore nbcache.Store) *SingleUseStore {
return &SingleUseStore{
cache: cacheStore,
ctx: ctx,
}
}
// Store saves value under key with the given TTL, after which it is evicted.
func (s *SingleUseStore) Store(key, value string, ttl time.Duration) error {
if err := s.cache.Set(s.ctx, key, value, store.WithExpiration(ttl)); err != nil {
return fmt.Errorf("store single-use value: %w", err)
}
return nil
}
// Generate stores a value under a namespaced random key and returns the random key.
func (s *SingleUseStore) Generate(namespace, value string, ttl time.Duration) (string, error) {
buf := make([]byte, 32)
if _, err := rand.Read(buf); err != nil {
return "", fmt.Errorf("generate single-use key: %w", err)
}
key := base64.RawURLEncoding.EncodeToString(buf)
if err := s.Store(singleUseCacheKey(namespace, key), value, ttl); err != nil {
return "", err
}
return key, nil
}
func singleUseCacheKey(namespace, key string) string {
return namespace + ":" + key
}
// LoadAndDelete retrieves and removes the value for a key.
func (s *SingleUseStore) LoadAndDelete(key string) (string, bool) {
value, found, err := s.cache.GetDel(s.ctx, key)
if err != nil {
log.Warnf("failed to consume single-use value: %v", err)
return "", false
}
if !found {
return "", false
}
return value, true
}
@@ -6,7 +6,7 @@ import (
"time"
)
func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
func TestSingleUseStoreLoadAndDelete(t *testing.T) {
const (
state = "state"
verifier = "verifier"
@@ -14,7 +14,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
)
t.Run("exactly one concurrent caller consumes the verifier", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, time.Minute); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -50,7 +50,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("replayed state is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, time.Minute); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -64,7 +64,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("unknown state is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if got, found := store.LoadAndDelete("never-stored"); found {
t.Fatalf("unknown state should not resolve, got %q", got)
@@ -72,7 +72,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("expired verifier is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, 50*time.Millisecond); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -83,3 +83,40 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
}
})
}
func TestSingleUseStore_GenerateAndConsumeOnce(t *testing.T) {
const namespace = "test"
s := NewSingleUseStore(context.Background(), testCacheStore(t))
key, err := s.Generate(namespace, "the-value", time.Minute)
if err != nil {
t.Fatalf("generate: %v", err)
}
if key == "" || key == "the-value" {
t.Fatalf("unexpected key %q", key)
}
value, found := s.LoadAndDelete(singleUseCacheKey(namespace, key))
if !found || value != "the-value" {
t.Fatalf("expected to load the stored value, got %q found=%v", value, found)
}
if _, found := s.LoadAndDelete(singleUseCacheKey(namespace, key)); found {
t.Fatal("value must be consumed on first LoadAndDelete")
}
}
func TestSingleUseStore_GenerateUniqueKeys(t *testing.T) {
s := NewSingleUseStore(context.Background(), testCacheStore(t))
a, err := s.Generate("test", "v", time.Minute)
if err != nil {
t.Fatalf("generate a: %v", err)
}
b, err := s.Generate("test", "v", time.Minute)
if err != nil {
t.Fatalf("generate b: %v", err)
}
if a == b {
t.Fatal("generated keys must be distinct")
}
}
@@ -40,9 +40,9 @@ func setupValidateSessionTest(t *testing.T) *validateSessionTestSetup {
proxyManager := &testValidateSessionProxyManager{}
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
proxyService := NewProxyServiceServer(nil, tokenStore, pkceStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil)
proxyService := NewProxyServiceServer(nil, tokenStore, singleUseStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil)
proxyService.SetServiceManager(serviceManager)
createTestProxies(t, ctx, testStore)
@@ -570,7 +570,7 @@ func (m *testValidateSessionServiceManager) DeleteAccountCluster(_ context.Conte
type testValidateSessionProxyManager struct{}
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
return nil, nil
}
@@ -634,6 +634,10 @@ func (m *testValidateSessionProxyManager) ClusterSupportsPrivate(_ context.Conte
return nil
}
func (m *testValidateSessionProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
return false
}
type testValidateSessionUsersManager struct {
store store.Store
}
@@ -662,3 +666,47 @@ func (m *testValidateSessionUsersManager) GetUserWithGroups(ctx context.Context,
}
return user, groups, nil
}
func TestValidateSession_RedeemsSessionCode(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "testProxyId")
require.NoError(t, err)
token := createSessionToken(t, proxy.SessionPrivateKey, "allowedUserId", "test-proxy.example.com")
code, ok := setup.proxyService.GenerateSessionCode(token)
require.True(t, ok)
require.NotEqual(t, token, code, "code must not be the token itself")
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: code,
})
require.NoError(t, err)
assert.True(t, resp.Valid, "redeemed code should authorize the user")
assert.Equal(t, "allowedUserId", resp.UserId)
assert.Equal(t, token, resp.GetSessionToken(), "response must carry the durable token for the cookie")
// Single-use: the same code must not redeem again.
resp2, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: code,
})
require.NoError(t, err)
assert.False(t, resp2.Valid, "a consumed code must be rejected")
assert.Empty(t, resp2.GetSessionToken())
}
func TestValidateSession_InvalidSessionCode(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: "does-not-exist",
})
require.NoError(t, err)
assert.False(t, resp.Valid)
assert.Empty(t, resp.GetSessionToken())
}
+125 -33
View File
@@ -112,6 +112,9 @@ type DefaultAccountManager struct {
permissionsManager permissions.Manager
disableDefaultPolicy bool
deletionHooksMu sync.RWMutex
deletionHooks []account.DeletionHook
}
var _ account.Manager = (*DefaultAccountManager)(nil)
@@ -120,6 +123,32 @@ func (am *DefaultAccountManager) SetServiceManager(serviceManager service.Manage
am.serviceManager = serviceManager
}
// AddAccountDeletionHook registers hook to run on every account deletion. Hooks run in
// registration order, and the first one to fail stops the rest and aborts the deletion.
// It panics on a nil hook: dropping one silently would skip that hook's cleanup on every
// deletion, so the wiring bug surfaces at startup instead.
func (am *DefaultAccountManager) AddAccountDeletionHook(hook account.DeletionHook) {
if hook == nil {
panic("nil account deletion hook")
}
am.deletionHooksMu.Lock()
defer am.deletionHooksMu.Unlock()
am.deletionHooks = append(am.deletionHooks, hook)
}
func (am *DefaultAccountManager) runAccountDeletionHooks(ctx context.Context, accountID string) error {
am.deletionHooksMu.RLock()
hooks := slices.Clone(am.deletionHooks)
am.deletionHooksMu.RUnlock()
for _, hook := range hooks {
if err := hook(ctx, accountID); err != nil {
return fmt.Errorf("account deletion hook: %w", err)
}
}
return nil
}
func isUniqueConstraintError(err error) bool {
switch {
case strings.Contains(err.Error(), "(SQLSTATE 23505)"),
@@ -305,6 +334,9 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco
var groupChangesAffectPeers bool
var reloadReverseProxy bool
var effectiveOldNetworkRange netip.Prefix
var ipv6Changed bool
var ipv6Snap *affectedpeers.Snapshot
var ipv6Change affectedpeers.Change
err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
var groupsUpdated bool
@@ -350,10 +382,10 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco
}
if ipv6SettingsChanged(oldSettings, newSettings) {
if err = am.updatePeerIPv6Addresses(ctx, transaction, accountID, newSettings); err != nil {
if ipv6Change, err = am.applyIPv6SettingsChange(ctx, transaction, accountID, oldSettings, newSettings); err != nil {
return err
}
updateAccountPeers = true
ipv6Changed = true
}
if oldSettings.RoutingPeerDNSResolutionEnabled != newSettings.RoutingPeerDNSResolutionEnabled ||
@@ -390,12 +422,20 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco
return err
}
if updateAccountPeers || groupsUpdated {
if updateAccountPeers || groupsUpdated || ipv6Changed {
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
return err
}
}
// A full account refresh already covers the IPv6 change, so the affected-peers
// snapshot is only needed when nothing account-wide changed.
if ipv6Changed && !updateAccountPeers && !groupChangesAffectPeers {
if ipv6Snap, err = affectedpeers.Load(ctx, transaction, accountID, ipv6Change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
}
return nil
})
if err != nil {
@@ -457,13 +497,34 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco
}
}
if updateAccountPeers || extraSettingsChanged || groupChangesAffectPeers {
go am.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceAccountSettings, Operation: types.UpdateOperationUpdate})
switch {
case updateAccountPeers || extraSettingsChanged || groupChangesAffectPeers:
go am.UpdateAccountPeers(context.WithoutCancel(ctx), accountID, types.UpdateReason{Resource: types.UpdateResourceAccountSettings, Operation: types.UpdateOperationUpdate})
case ipv6Snap != nil:
am.ExpandAndUpdateAffected(ctx, accountID, ipv6Snap, ipv6Change)
}
return newSettings, nil
}
// applyIPv6SettingsChange reconciles peer IPv6 addresses for new IPv6 settings and
// returns the affected-peers change: peers whose address changed refresh together
// with every peer that reaches them. On a range change every peer holding an address
// also refreshes itself, since its interface prefix comes from the account range even
// when its address stays inside the new one.
func (am *DefaultAccountManager) applyIPv6SettingsChange(ctx context.Context, transaction store.Store, accountID string, oldSettings, newSettings *types.Settings) (affectedpeers.Change, error) {
result, err := am.updatePeerIPv6Addresses(ctx, transaction, accountID, newSettings)
if err != nil {
return affectedpeers.Change{}, err
}
change := affectedpeers.Change{ChangedPeerIDs: result.changed}
if oldSettings.NetworkRangeV6 != newSettings.NetworkRangeV6 {
change.OutputPeerIDs = result.withIPv6
}
return change, nil
}
func ipv6SettingsChanged(old, updated *types.Settings) bool {
if old.NetworkRangeV6 != updated.NetworkRangeV6 {
return true
@@ -889,6 +950,10 @@ func (am *DefaultAccountManager) DeleteAccount(ctx context.Context, accountID, u
return status.Errorf(status.Internal, "failed to build user infos for account %s: %v", accountID, err)
}
if err = am.runAccountDeletionHooks(ctx, accountID); err != nil {
return err
}
if err = am.deleteAccountUsers(ctx, accountID, userID, account.Users, userInfosMap); err != nil {
return err
}
@@ -1709,9 +1774,11 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth
change.LinkGroups = allGroupChanges
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges); err != nil {
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges)
if err != nil {
return fmt.Errorf("reconcile IPv6 for group changes: %w", err)
}
change.ChangedPeerIDs = append(change.ChangedPeerIDs, ipv6Changed...)
if err = transaction.IncrementNetworkSerial(ctx, userAuth.AccountId); err != nil {
return fmt.Errorf("error incrementing network serial: %w", err)
@@ -2301,7 +2368,8 @@ func (am *DefaultAccountManager) propagateUserGroupMemberships(ctx context.Conte
return false, false, err
}
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, updatedGroups); err != nil {
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, updatedGroups)
if err != nil {
return false, false, fmt.Errorf("reconcile IPv6 for group changes: %w", err)
}
@@ -2310,7 +2378,7 @@ func (am *DefaultAccountManager) propagateUserGroupMemberships(ctx context.Conte
return false, false, fmt.Errorf("error checking if group changes affect peers: %w", err)
}
return len(updatedGroups) > 0, peersAffected, nil
return len(updatedGroups) > 0, peersAffected || len(ipv6Changed) > 0, nil
}
// propagateAutoGroupsForUsers adds each user's peers to their AutoGroups where not already present.
@@ -2358,7 +2426,7 @@ func (am *DefaultAccountManager) reallocateAccountPeerIPs(ctx context.Context, t
return err
}
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "")
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "")
if err != nil {
return err
}
@@ -2395,7 +2463,7 @@ func (am *DefaultAccountManager) reallocateAccountPeerIPs(ctx context.Context, t
// v6 address get one allocated. When disabled, all v6 addresses are cleared.
// When the v6 range changes, all v6 addresses are reallocated.
func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transaction store.Store, accountID, peerID string, newIPv6 netip.Addr) error {
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "")
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "", "")
if err != nil {
return fmt.Errorf("get peers: %w", err)
}
@@ -2407,56 +2475,78 @@ func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transac
return nil
}
func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) error {
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "")
// ipv6Reassignment reports the outcome of an IPv6 address reconciliation.
type ipv6Reassignment struct {
// changed are the peers whose IPv6 address was assigned, removed or reallocated.
changed []string
// withIPv6 are all peers holding an IPv6 address after the reconciliation.
withIPv6 []string
}
func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) (ipv6Reassignment, error) {
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "")
if err != nil {
return fmt.Errorf("get peers: %w", err)
return ipv6Reassignment{}, fmt.Errorf("get peers: %w", err)
}
network, err := transaction.GetAccountNetwork(ctx, store.LockingStrengthUpdate, accountID)
if err != nil {
return fmt.Errorf("get network: %w", err)
return ipv6Reassignment{}, fmt.Errorf("get network: %w", err)
}
if err := am.ensureIPv6Subnet(ctx, transaction, accountID, settings, network); err != nil {
return err
return ipv6Reassignment{}, err
}
allowedPeers, err := am.buildIPv6AllowedPeers(ctx, transaction, accountID, settings)
if err != nil {
return err
return ipv6Reassignment{}, err
}
v6Prefix, err := netip.ParsePrefix(network.NetV6.String())
if err != nil {
return fmt.Errorf("parse IPv6 prefix: %w", err)
return ipv6Reassignment{}, fmt.Errorf("parse IPv6 prefix: %w", err)
}
if err := am.assignPeerIPv6Addresses(ctx, transaction, accountID, peers, network, allowedPeers, v6Prefix); err != nil {
return err
changed, err := am.assignPeerIPv6Addresses(ctx, transaction, accountID, peers, network, allowedPeers, v6Prefix)
if err != nil {
return ipv6Reassignment{}, err
}
log.WithContext(ctx).Infof("updated IPv6 addresses for %d peers in account %s (groups=%d)",
len(peers), accountID, len(settings.IPv6EnabledGroups))
result := ipv6Reassignment{changed: changed}
for _, peer := range peers {
if peer.IPv6.IsValid() {
result.withIPv6 = append(result.withIPv6, peer.ID)
}
}
return nil
log.WithContext(ctx).Infof("updated IPv6 addresses for %d of %d peers in account %s (groups=%d)",
len(changed), len(peers), accountID, len(settings.IPv6EnabledGroups))
return result, nil
}
// reconcileIPv6ForGroupChanges checks whether the given group IDs overlap with
// the account's IPv6EnabledGroups. If they do, it runs a full IPv6 address
// reconciliation so that peers gaining or losing membership in an IPv6-enabled
// group get their addresses assigned or removed.
func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) error {
// group get their addresses assigned or removed. It returns the peers whose IPv6
// address changed, which callers pass as changed peers so every peer that can
// reach them refreshes.
func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) ([]string, error) {
settings, err := transaction.GetAccountSettings(ctx, store.LockingStrengthNone, accountID)
if err != nil {
return fmt.Errorf("get account settings: %w", err)
return nil, fmt.Errorf("get account settings: %w", err)
}
if !ipv6ReconcileNeeded(settings, groupIDs) {
return nil
return nil, nil
}
return am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings)
result, err := am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings)
if err != nil {
return nil, err
}
return result.changed, nil
}
// ipv6ReconcileNeeded reports whether changes to the given groups trigger an IPv6
@@ -2495,7 +2585,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses(
ctx context.Context, transaction store.Store, accountID string,
peers []*nbpeer.Peer, network *types.Network,
allowedPeers map[string]struct{}, v6Prefix netip.Prefix,
) error {
) ([]string, error) {
takenV6 := make(map[netip.Addr]struct{})
for _, peer := range peers {
if _, ok := allowedPeers[peer.ID]; ok && peer.IPv6.IsValid() && network.NetV6.Contains(peer.IPv6.AsSlice()) {
@@ -2503,6 +2593,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses(
}
}
var changed []string
for _, peer := range peers {
_, allowed := allowedPeers[peer.ID]
oldIPv6 := peer.IPv6
@@ -2512,7 +2603,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses(
} else if !peer.IPv6.IsValid() || !network.NetV6.Contains(peer.IPv6.AsSlice()) {
newIP, err := allocateIPv6WithRetry(v6Prefix, takenV6, peer.ID)
if err != nil {
return err
return nil, err
}
peer.IPv6 = newIP
}
@@ -2522,10 +2613,11 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses(
}
if err := transaction.SavePeer(ctx, accountID, peer); err != nil {
return fmt.Errorf("save peer %s: %w", peer.ID, err)
return nil, fmt.Errorf("save peer %s: %w", peer.ID, err)
}
changed = append(changed, peer.ID)
}
return nil
return changed, nil
}
func allocateIPv6WithRetry(prefix netip.Prefix, taken map[netip.Addr]struct{}, peerID string) (netip.Addr, error) {
@@ -2569,7 +2661,7 @@ func (am *DefaultAccountManager) buildIPv6AllowedPeers(ctx context.Context, tran
// Embedded proxy peers sit outside regular group membership but must
// participate in any v6-enabled overlay to reach v6-only peers.
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
if err != nil {
return nil, fmt.Errorf("get peers: %w", err)
}
@@ -2640,7 +2732,7 @@ func (am *DefaultAccountManager) updatePeerIPInTransaction(ctx context.Context,
return nil
}
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "")
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "", "")
if err != nil {
return fmt.Errorf("get account peers: %w", err)
}
@@ -0,0 +1,14 @@
package account
import "context"
// DeletionHook runs when an account is deleted, after the caller's permission to delete
// it has been checked and before any of its users or data are removed. It lets code that
// keeps per-account state outside the store tear that state down while the account still
// exists.
//
// A hook that returns an error aborts the deletion and the account is kept. The caller
// sees the error, so a hook that wants a specific response returns a status error. A
// retried deletion runs every hook again, and a later step can still fail after the hooks
// succeed, so a hook must be idempotent and must tolerate the account surviving it.
type DeletionHook func(ctx context.Context, accountID string) error
+1 -1
View File
@@ -62,7 +62,7 @@ type Manager interface {
GetUserByID(ctx context.Context, id string) (*types.User, error)
GetUserFromUserAuth(ctx context.Context, userAuth auth.UserAuth) (*types.User, error)
ListUsers(ctx context.Context, accountID string) ([]*types.User, error)
GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error)
GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error)
MarkPeerConnected(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error
MarkPeerDisconnected(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error
DeletePeer(ctx context.Context, accountID, peerID, userID string) error
+4 -4
View File
@@ -982,18 +982,18 @@ func (mr *MockManagerMockRecorder) GetPeerNetwork(ctx, peerID any) *gomock.Call
}
// GetPeers mocks base method.
func (m *MockManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*peer.Peer, error) {
func (m *MockManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*peer.Peer, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetPeers", ctx, accountID, userID, nameFilter, ipFilter)
ret := m.ctrl.Call(m, "GetPeers", ctx, accountID, userID, nameFilter, ipFilter, macFilter)
ret0, _ := ret[0].([]*peer.Peer)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetPeers indicates an expected call of GetPeers.
func (mr *MockManagerMockRecorder) GetPeers(ctx, accountID, userID, nameFilter, ipFilter any) *gomock.Call {
func (mr *MockManagerMockRecorder) GetPeers(ctx, accountID, userID, nameFilter, ipFilter, macFilter any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeers", reflect.TypeOf((*MockManager)(nil).GetPeers), ctx, accountID, userID, nameFilter, ipFilter)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeers", reflect.TypeOf((*MockManager)(nil).GetPeers), ctx, accountID, userID, nameFilter, ipFilter, macFilter)
}
// GetPolicy mocks base method.
+104 -9
View File
@@ -958,6 +958,101 @@ func TestAccountManager_DeleteAccount(t *testing.T) {
assert.Len(t, pats, 0)
}
func TestAccountManager_DeleteAccount_RunsDeletionHooks(t *testing.T) {
manager, _, err := createManager(t)
require.NoError(t, err)
ownerID := "account_creator"
account, err := createAccount(manager, "test_account", ownerID, "")
require.NoError(t, err)
// Each hook records its call and checks the account is still in the store, which is
// the point of running before deletion: a hook must be able to read what it cleans up.
var calls []string
hook := func(name string) nbAccount.DeletionHook {
return func(ctx context.Context, accountID string) error {
calls = append(calls, name+":"+accountID)
_, err := manager.Store.GetAccount(ctx, accountID)
assert.NoError(t, err, "account should still exist while hook %s runs", name)
return nil
}
}
manager.AddAccountDeletionHook(hook("first"))
manager.AddAccountDeletionHook(hook("second"))
require.NoError(t, manager.DeleteAccount(context.Background(), account.Id, ownerID))
assert.Equal(t, []string{"first:" + account.Id, "second:" + account.Id}, calls,
"hooks should run once each, in registration order, with the deleted account's ID")
_, err = manager.Store.GetAccount(context.Background(), account.Id)
assert.Error(t, err, "account should be deleted after the hooks succeed")
}
func TestAccountManager_DeleteAccount_DeletionHookErrorAbortsDeletion(t *testing.T) {
manager, _, err := createManager(t)
require.NoError(t, err)
ownerID := "account_creator"
account, err := createAccount(manager, "test_account", ownerID, "")
require.NoError(t, err)
manager.AddAccountDeletionHook(func(context.Context, string) error {
return status.Errorf(status.PreconditionFailed, "teardown refused")
})
secondCalled := false
manager.AddAccountDeletionHook(func(context.Context, string) error {
secondCalled = true
return nil
})
err = manager.DeleteAccount(context.Background(), account.Id, ownerID)
require.Error(t, err)
// The hook's status type has to survive the wrapping, since the HTTP layer maps it
// to the response code.
sErr, ok := status.FromError(err)
require.True(t, ok, "error should carry the hook's status error, got %v", err)
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "status type should be the hook's")
assert.False(t, secondCalled, "hooks after a failing one should not run")
_, err = manager.Store.GetAccount(context.Background(), account.Id)
assert.NoError(t, err, "account should survive a failing hook")
_, err = manager.Store.GetUserByUserID(context.Background(), store.LockingStrengthNone, ownerID)
assert.NoError(t, err, "account owner should survive a failing hook")
}
func TestAccountManager_AddAccountDeletionHook_RejectsNil(t *testing.T) {
manager, _, err := createManager(t)
require.NoError(t, err)
assert.PanicsWithValue(t, "nil account deletion hook", func() {
manager.AddAccountDeletionHook(nil)
}, "registering a nil hook should panic instead of breaking a later deletion")
}
func TestAccountManager_DeleteAccount_DeletionHooksSkippedWithoutPermission(t *testing.T) {
manager, _, err := createManager(t)
require.NoError(t, err)
ownerID := "account_creator"
account, err := createAccount(manager, "test_account", ownerID, "")
require.NoError(t, err)
adminID := "regular_admin"
account.Users[adminID] = types.NewAdminUser(adminID)
require.NoError(t, manager.Store.SaveAccount(context.Background(), account))
called := false
manager.AddAccountDeletionHook(func(context.Context, string) error {
called = true
return nil
})
err = manager.DeleteAccount(context.Background(), account.Id, adminID)
require.Error(t, err, "only the owner may delete the account")
assert.False(t, called, "hooks should not run for a caller who may not delete the account")
}
func BenchmarkTest_GetAccountWithclaims(b *testing.B) {
claims := auth.UserAuth{
Domain: "example.com",
@@ -2462,7 +2557,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_PeerApproval(t *testing.T)
_, err = manager.UpdateAccountSettings(ctx, accountID, userID, newSettings)
require.NoError(t, err)
accountPeers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
accountPeers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
require.NoError(t, err)
for _, peer := range accountPeers {
@@ -4462,7 +4557,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
})
require.NoError(t, err)
peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "")
peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "")
require.NoError(t, err)
require.Len(t, peers, len(before))
for _, p := range peers {
@@ -4480,7 +4575,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
})
require.NoError(t, err)
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "")
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "")
require.NoError(t, err)
for _, p := range peers {
assert.Equal(t, before[p.ID], p.IP, "peer %s IP should not change for host-bit-set equivalent range", p.ID)
@@ -4494,7 +4589,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
})
require.NoError(t, err)
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "")
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "")
require.NoError(t, err)
for _, p := range peers {
assert.Equal(t, before[p.ID], p.IP, "peer %s IP should not change when NetworkRange omitted", p.ID)
@@ -4510,7 +4605,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
})
require.NoError(t, err)
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "")
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "")
require.NoError(t, err)
for _, p := range peers {
assert.True(t, newRange.Contains(p.IP), "peer %s should be in new range %s, got %s", p.ID, newRange, p.IP)
@@ -4528,7 +4623,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin
require.NoError(t, err)
require.NotEmpty(t, settings.IPv6EnabledGroups, "new account should have IPv6 enabled for All group")
peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
require.NoError(t, err)
for _, p := range peers {
assert.True(t, p.IPv6.IsValid(), "peer %s should have IPv6 with All group enabled", p.ID)
@@ -4556,7 +4651,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin
assert.Equal(t, []string{partialGroup.ID}, updatedSettings.IPv6EnabledGroups)
// peer1 and peer2 should have IPv6; peer3 should not.
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
require.NoError(t, err)
peerMap := make(map[string]*nbpeer.Peer, len(peers))
for _, p := range peers {
@@ -4576,7 +4671,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin
require.NoError(t, err)
assert.Empty(t, updatedSettings.IPv6EnabledGroups)
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
require.NoError(t, err)
for _, p := range peers {
assert.False(t, p.IPv6.IsValid(), "peer %s should have no IPv6 when groups cleared", p.ID)
@@ -4591,7 +4686,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin
})
require.NoError(t, err)
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
require.NoError(t, err)
peerMap = make(map[string]*nbpeer.Peer, len(peers))
for _, p := range peers {
@@ -0,0 +1,243 @@
package server
import (
"context"
"net/netip"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
)
const (
ipv6GroupA = "ipv6-grp-a"
ipv6GroupB = "ipv6-grp-b"
ipv6GroupC = "ipv6-grp-c"
ipv6GroupD = "ipv6-grp-d"
)
// ipv6AffectedTest holds three peers: peer1 in group A, peer2 in group B, peer3 in
// group C, with a single A<->B policy. peer3 is unrelated to peer1 and peer2. Group D
// is empty and referenced by nothing.
type ipv6AffectedTest struct {
manager *DefaultAccountManager
accountID string
peer1, peer2, peer3 *nbpeer.Peer
updMsg1, updMsg2, updMsg3 <-chan *network_map.UpdateMessage
}
func setupIPv6AffectedTest(t *testing.T, ipv6Groups []string) *ipv6AffectedTest {
t.Helper()
manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t)
ctx := context.Background()
accountID := account.Id
policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID)
require.NoError(t, err)
for _, p := range policies {
require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID))
}
for _, g := range []*types.Group{
{ID: ipv6GroupA, Name: "IPv6-A", Peers: []string{peer1.ID}},
{ID: ipv6GroupB, Name: "IPv6-B", Peers: []string{peer2.ID}},
{ID: ipv6GroupC, Name: "IPv6-C", Peers: []string{peer3.ID}},
{ID: ipv6GroupD, Name: "IPv6-D"},
} {
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, g))
}
_, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{
Enabled: true,
Rules: []*types.PolicyRule{{
Enabled: true,
Sources: []string{ipv6GroupA},
Destinations: []string{ipv6GroupB},
Bidirectional: true,
Action: types.PolicyTrafficActionAccept,
}},
}, true)
require.NoError(t, err)
// New accounts enable IPv6 for the All group; start from the requested groups.
updateIPv6TestSettings(t, manager, accountID, func(s *types.Settings) {
s.IPv6EnabledGroups = ipv6Groups
})
tc := &ipv6AffectedTest{
manager: manager,
accountID: accountID,
peer1: peer1,
peer2: peer2,
peer3: peer3,
}
tc.updMsg1 = updateManager.CreateChannel(ctx, peer1.ID)
tc.updMsg2 = updateManager.CreateChannel(ctx, peer2.ID)
tc.updMsg3 = updateManager.CreateChannel(ctx, peer3.ID)
t.Cleanup(func() {
updateManager.CloseChannel(ctx, peer1.ID)
updateManager.CloseChannel(ctx, peer2.ID)
updateManager.CloseChannel(ctx, peer3.ID)
})
// The setup changes above dispatch asynchronously and can land after the
// channels open, so drop them before the test acts.
drainPeerUpdates(tc.updMsg1)
drainPeerUpdates(tc.updMsg2)
drainPeerUpdates(tc.updMsg3)
return tc
}
// updateIPv6TestSettings applies mutate to a copy of the current settings, so only
// the mutated fields differ from what is stored.
func updateIPv6TestSettings(t *testing.T, manager *DefaultAccountManager, accountID string, mutate func(*types.Settings)) {
t.Helper()
ctx := context.Background()
current, err := manager.Store.GetAccountSettings(ctx, store.LockingStrengthNone, accountID)
require.NoError(t, err)
updated := current.Copy()
mutate(updated)
_, err = manager.UpdateAccountSettings(ctx, accountID, userID, updated)
require.NoError(t, err)
}
func (tc *ipv6AffectedTest) peerIPv6(t *testing.T, peerID string) netip.Addr {
t.Helper()
peer, err := tc.manager.Store.GetPeerByID(context.Background(), store.LockingStrengthNone, tc.accountID, peerID)
require.NoError(t, err)
return peer.IPv6
}
func TestAffectedPeers_IPv6GroupEnabled_RefreshesOnlyReachablePeers(t *testing.T) {
tc := setupIPv6AffectedTest(t, nil)
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
s.IPv6EnabledGroups = []string{ipv6GroupA}
})
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address")
peerShouldReceiveUpdate(t, tc.updMsg1)
peerShouldReceiveUpdate(t, tc.updMsg2)
peerShouldNotReceiveUpdate(t, tc.updMsg3)
}
func TestAffectedPeers_IPv6GroupDisabled_RefreshesOnlyReachablePeers(t *testing.T) {
tc := setupIPv6AffectedTest(t, []string{ipv6GroupA})
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should start with an IPv6 address")
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
s.IPv6EnabledGroups = []string{}
})
require.False(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should lose its IPv6 address")
peerShouldReceiveUpdate(t, tc.updMsg1)
peerShouldReceiveUpdate(t, tc.updMsg2)
peerShouldNotReceiveUpdate(t, tc.updMsg3)
}
// Widening the IPv6 range keeps peer addresses, but each holder's interface prefix
// comes from the range, so holders refresh while peers that only reach them do not.
func TestAffectedPeers_IPv6RangeWidened_RefreshesAddressHolders(t *testing.T) {
tc := setupIPv6AffectedTest(t, []string{ipv6GroupA})
oldIPv6 := tc.peerIPv6(t, tc.peer1.ID)
require.True(t, oldIPv6.IsValid(), "peer1 should start with an IPv6 address")
// The range is allocated on the account network; settings may leave it empty.
network, err := tc.manager.Store.GetAccountNetwork(context.Background(), store.LockingStrengthNone, tc.accountID)
require.NoError(t, err)
current := prefixFromIPNet(network.NetV6)
require.True(t, current.IsValid(), "account should have an IPv6 range")
widened := netip.PrefixFrom(current.Addr(), current.Bits()-8).Masked()
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
s.NetworkRangeV6 = widened
})
require.Equal(t, oldIPv6, tc.peerIPv6(t, tc.peer1.ID), "peer1 should keep its address inside the widened range")
peerShouldReceiveUpdate(t, tc.updMsg1)
peerShouldNotReceiveUpdate(t, tc.updMsg2)
peerShouldNotReceiveUpdate(t, tc.updMsg3)
}
func TestAffectedPeers_IPv4RangeChange_RefreshesWholeAccount(t *testing.T) {
tc := setupIPv6AffectedTest(t, nil)
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
s.NetworkRange = netip.MustParsePrefix("100.70.0.0/16")
})
peerShouldReceiveUpdate(t, tc.updMsg1)
peerShouldReceiveUpdate(t, tc.updMsg2)
peerShouldReceiveUpdate(t, tc.updMsg3)
}
func TestAffectedPeers_IPv6WithAccountWideChange_RefreshesWholeAccount(t *testing.T) {
tc := setupIPv6AffectedTest(t, nil)
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
s.IPv6EnabledGroups = []string{ipv6GroupA}
s.LazyConnectionEnabled = !s.LazyConnectionEnabled
})
peerShouldReceiveUpdate(t, tc.updMsg1)
peerShouldReceiveUpdate(t, tc.updMsg2)
peerShouldReceiveUpdate(t, tc.updMsg3)
}
// Joining an IPv6-enabled group that no policy references gives peer1 an address.
// peer2 reaches peer1 through group A, not through the joined group, and must still
// learn the new address.
func TestAffectedPeers_GroupAddPeerIPv6_RefreshesPeersReachingThroughOtherGroups(t *testing.T) {
tc := setupIPv6AffectedTest(t, []string{ipv6GroupD})
require.NoError(t, tc.manager.GroupAddPeer(context.Background(), tc.accountID, ipv6GroupD, tc.peer1.ID))
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address")
peerShouldReceiveUpdate(t, tc.updMsg1)
peerShouldReceiveUpdate(t, tc.updMsg2)
peerShouldNotReceiveUpdate(t, tc.updMsg3)
}
func TestAffectedPeers_UpdateGroupIPv6_RefreshesPeersReachingThroughOtherGroups(t *testing.T) {
tc := setupIPv6AffectedTest(t, []string{ipv6GroupD})
require.NoError(t, tc.manager.UpdateGroup(context.Background(), tc.accountID, userID, &types.Group{
ID: ipv6GroupD,
Name: "IPv6-D",
Peers: []string{tc.peer1.ID},
}))
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address")
peerShouldReceiveUpdate(t, tc.updMsg1)
peerShouldReceiveUpdate(t, tc.updMsg2)
peerShouldNotReceiveUpdate(t, tc.updMsg3)
}
// Deleting an IPv6-enabled group removes its members' addresses after the
// pre-delete snapshot was taken.
func TestAffectedPeers_DeleteIPv6Group_RefreshesFormerMembersAndReachablePeers(t *testing.T) {
tc := setupIPv6AffectedTest(t, []string{ipv6GroupD})
ctx := context.Background()
require.NoError(t, tc.manager.GroupAddPeer(ctx, tc.accountID, ipv6GroupD, tc.peer1.ID))
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address")
drainPeerUpdates(tc.updMsg1)
drainPeerUpdates(tc.updMsg2)
drainPeerUpdates(tc.updMsg3)
require.NoError(t, tc.manager.DeleteGroup(ctx, tc.accountID, userID, ipv6GroupD))
require.False(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should lose its IPv6 address")
peerShouldReceiveUpdate(t, tc.updMsg1)
peerShouldReceiveUpdate(t, tc.updMsg2)
peerShouldNotReceiveUpdate(t, tc.updMsg3)
}
@@ -108,11 +108,13 @@ func TestAffectedPeers_SaveUser_OnlyAffectedPeersUpdated(t *testing.T) {
})
t.Run("auto group change reassigning IPv6 refreshes the changed peers and their observers", func(t *testing.T) {
account, err := manager.Store.GetAccount(ctx, accountID)
require.NoError(t, err)
account.Settings.IPv6EnabledGroups = []string{"ug-v6"}
require.NoError(t, manager.Store.SaveAccount(ctx, account))
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "ug-v6", Name: "ug-v6"}))
// Apply through the settings API so the reconciliation that strips the other
// peers' addresses happens here, leaving the target as the only peer the
// user update reassigns.
updateIPv6TestSettings(t, manager, accountID, func(s *types.Settings) {
s.IPv6EnabledGroups = []string{"ug-v6"}
})
drainPeerUpdates(updTarget)
drainPeerUpdates(upd2)
@@ -0,0 +1,145 @@
package server
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
"github.com/netbirdio/netbird/management/internals/modules/zones"
"github.com/netbirdio/netbird/management/internals/modules/zones/records"
"github.com/netbirdio/netbird/management/server/affectedpeers"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
)
const affectedZoneDomain = "zone.test"
// createAffectedZone stores a zone distributed to the given groups, optionally with
// one A record so the network map actually ships it.
func createAffectedZone(t *testing.T, s store.Store, accountID, domain string, enabled, withRecord bool, groups []string) *zones.Zone {
t.Helper()
ctx := context.Background()
zone := zones.NewZone(accountID, domain, domain, enabled, false, groups)
require.NoError(t, s.CreateZone(ctx, zone))
if withRecord {
record := records.NewRecord(accountID, zone.ID, "host."+domain, records.RecordTypeA, "10.0.0.1", 300)
require.NoError(t, s.CreateDNSRecord(ctx, record))
}
return zone
}
func TestCollectGroupChange_ZoneLinked(t *testing.T) {
_, s, accountID, _, groupIDs := setupAffectedPeersTest(t)
ctx := context.Background()
createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]})
groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]})
assert.Contains(t, groups, groupIDs[0], "group distributed a zone should be affected by its own change")
groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]})
assert.Empty(t, groups, "group not referenced by any zone should not be affected")
}
func TestCollectGroupChange_UnshippedZoneNotLinked(t *testing.T) {
_, s, accountID, _, groupIDs := setupAffectedPeersTest(t)
ctx := context.Background()
// Disabled zone and zone without records are never shipped by the network map.
createAffectedZone(t, s, accountID, "disabled."+affectedZoneDomain, false, true, []string{groupIDs[0]})
createAffectedZone(t, s, accountID, "empty."+affectedZoneDomain, true, false, []string{groupIDs[1]})
groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0], groupIDs[1]})
assert.Empty(t, groups, "groups referenced only by unshipped zones should not be affected")
}
func TestResolveAffectedPeers_ZoneGroupMembershipChange(t *testing.T) {
_, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t)
createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]})
// Same change shape UpdateGroup builds: the group changed as a whole and peer1
// left it, so peer1 must refresh to drop the zone.
change := affectedpeers.Change{
ChangedGroupIDs: []string{groupIDs[0]},
RemovedPeersByGroup: map[string][]string{groupIDs[0]: {peerIDs[1]}},
}
result := resolveAffected(t, s, accountID, change)
assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result, "current and removed members of the zone group should be affected")
}
func TestResolveAffectedPeers_ZoneDistributionChange(t *testing.T) {
_, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t)
// Zone create/update/delete passes old and new distribution groups.
change := affectedpeers.Change{DistributionGroupIDs: []string{groupIDs[0], groupIDs[2]}}
result := resolveAffected(t, s, accountID, change)
assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2]}, result, "only members of the distribution groups should be affected")
}
// TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone verifies that adding a peer
// to a group referenced only by a zone pushes the zone to the new member and leaves
// unrelated peers alone.
func TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone(t *testing.T) {
manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t)
ctx := context.Background()
accountID := account.Id
policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID)
require.NoError(t, err)
for _, p := range policies {
require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID))
}
zoneGroup := &types.Group{ID: "zone-grp", Name: "ZoneGroup", Peers: []string{peer1.ID}}
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, zoneGroup))
createAffectedZone(t, manager.Store, accountID, affectedZoneDomain, true, true, []string{zoneGroup.ID})
updMsg1 := updateManager.CreateChannel(ctx, peer1.ID)
updMsg2 := updateManager.CreateChannel(ctx, peer2.ID)
updMsg3 := updateManager.CreateChannel(ctx, peer3.ID)
t.Cleanup(func() {
updateManager.CloseChannel(ctx, peer1.ID)
updateManager.CloseChannel(ctx, peer2.ID)
updateManager.CloseChannel(ctx, peer3.ID)
})
zoneGroup.Peers = []string{peer1.ID, peer2.ID}
require.NoError(t, manager.UpdateGroup(ctx, accountID, userID, zoneGroup))
peerShouldReceiveUpdate(t, updMsg1)
msg := receivePeerUpdate(t, updMsg2)
assert.True(t, syncHasCustomZone(msg, affectedZoneDomain+"."), "new zone group member should receive the zone")
peerShouldNotReceiveUpdate(t, updMsg3)
}
func receivePeerUpdate(t *testing.T, ch <-chan *network_map.UpdateMessage) *network_map.UpdateMessage {
t.Helper()
select {
case msg := <-ch:
require.NotNil(t, msg, "update message should not be nil")
return msg
case <-time.After(peerUpdateTimeout):
require.FailNow(t, "timed out waiting for update message")
return nil
}
}
func syncHasCustomZone(msg *network_map.UpdateMessage, domain string) bool {
for _, zone := range msg.Update.GetNetworkMap().GetDNSConfig().GetCustomZones() {
if zone.GetDomain() == domain {
return true
}
}
return false
}
+27 -2
View File
@@ -22,6 +22,7 @@ import (
nbdns "github.com/netbirdio/netbird/dns"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/internals/modules/zones"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
networkTypes "github.com/netbirdio/netbird/management/server/networks/types"
@@ -50,6 +51,7 @@ type Snapshot struct {
policies []*types.Policy
routes []*route.Route
nsGroups []*nbdns.NameServerGroup
zones []*zones.Zone
dnsSettings *types.DNSSettings
routers []*routerTypes.NetworkRouter
resources []*resourceTypes.NetworkResource
@@ -127,12 +129,15 @@ func (snap *Snapshot) loadRoutesAndProxy(ctx context.Context, s store.Store, acc
return snap.loadProxyServices(ctx, s, accountID)
}
// loadDNS loads the nameserver groups and account DNS settings.
// loadDNS loads the nameserver groups, custom DNS zones and account DNS settings.
func (snap *Snapshot) loadDNS(ctx context.Context, s store.Store, accountID string) error {
var err error
if snap.nsGroups, err = s.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID); err != nil {
return err
}
if snap.zones, err = s.GetAccountZones(ctx, store.LockingStrengthNone, accountID); err != nil {
return err
}
snap.dnsSettings, err = s.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID)
return err
}
@@ -357,7 +362,7 @@ func (s policySide) opposite() policySide {
// - a changed router/resource/network sits on a NETWORK -> fold the SOURCE side of
// the policies whose destination reaches it (and the routers it implies).
//
// Routes, nameserver groups, DNS and embedded-proxy services distribute to their own
// Routes, nameserver groups, DNS zones, DNS and embedded-proxy services distribute to their own
// member peers, outside the policy graph, and are folded here too.
func (r *resolver) walk() {
for _, policy := range r.bothSidesPolicies() {
@@ -369,6 +374,7 @@ func (r *resolver) walk() {
r.collectFromPolicies()
r.collectFromRoutes()
r.collectFromNameServers()
r.collectFromZones()
r.collectFromDNSSettings()
r.collectFromNetworkRouters()
r.collectFromProxyServices()
@@ -829,6 +835,25 @@ func (r *resolver) collectFromNameServers() {
}
}
// collectFromZones folds the distribution groups of the custom DNS zones that
// reference a linked group. Like nameserver groups, a zone has no opposite side, so
// only a whole-group change folds its groups. Zones the network map does not ship
// (disabled or without records) are skipped.
func (r *resolver) collectFromZones() {
if len(r.linkGroups) == 0 {
return
}
for _, zone := range r.snap.zones {
if !zone.Enabled || len(zone.Records) == 0 {
continue
}
if anyInSet(zone.DistributionGroups, r.linkGroups) {
log.WithContext(r.ctx).Tracef("collectFromZones: zone %s references a linked group -> folding its groups %v (outputGroups only)", zone.ID, zone.DistributionGroups)
r.foldOutputGroups(zone.DistributionGroups)
}
}
}
// collectFromSSHAuthorizedGroups folds the destinations of the enabled SSH rules that
// authorize a group whose user membership changed. Those destination peers carry the
// group -> user mapping for the groups they authorize, so they refresh even when no
+19 -1
View File
@@ -3,9 +3,11 @@ package auth
import (
"context"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"strings"
"time"
)
@@ -52,6 +54,22 @@ func (s *SessionStore) RegisterToken(ctx context.Context, token string, expiresA
}
func hashToken(token string) string {
sum := sha256.Sum256([]byte(token))
sum := sha256.Sum256([]byte(canonicalizeToken(token)))
return hex.EncodeToString(sum[:])
}
// canonicalizeToken re-encodes the JWT signature segment so noncanonical
// spellings of the same signature map to one stable cache key.
func canonicalizeToken(token string) string {
parts := strings.Split(token, ".")
if len(parts) != 3 {
return token
}
sig, err := base64.RawURLEncoding.DecodeString(parts[2])
if err != nil {
return token
}
return parts[0] + "." + parts[1] + "." + base64.RawURLEncoding.EncodeToString(sig)
}
+51
View File
@@ -2,10 +2,15 @@ package auth
import (
"context"
"crypto/rand"
"crypto/rsa"
"encoding/base64"
"errors"
"strings"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -131,3 +136,49 @@ func TestHashToken_StableAndDoesNotLeak(t *testing.T) {
assert.Len(t, a, 64, "sha256 hex must be 64 chars")
assert.NotContains(t, a, "tokenA", "raw token must not appear in hash")
}
func TestSessionStore_NoncanonicalSpellingIsRejectedAsReplay(t *testing.T) {
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)
token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims{
"sub": "user",
"exp": time.Now().Add(time.Hour).Unix(),
})
canonical, err := token.SignedString(privateKey)
require.NoError(t, err)
parts := strings.Split(canonical, ".")
require.Len(t, parts, 3)
// A 256-byte RSA signature (256 mod 3 == 1) leaves unused bits in the final
// base64url character; flip one without changing the decoded signature.
const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_"
last := strings.IndexByte(alphabet, parts[2][len(parts[2])-1])
require.GreaterOrEqual(t, last, 0)
require.Equal(t, 0, last&3, "unexpected canonical RSA signature encoding")
equivalentSig := parts[2][:len(parts[2])-1] + string(alphabet[last|1])
equivalent := parts[0] + "." + parts[1] + "." + equivalentSig
require.NotEqual(t, canonical, equivalent, "spellings must differ as strings")
// Same decoded signature bytes, so they verify as the same JWT.
canonicalSig, err := base64.RawURLEncoding.DecodeString(parts[2])
require.NoError(t, err)
altSig, err := base64.RawURLEncoding.DecodeString(equivalentSig)
require.NoError(t, err)
require.Equal(t, canonicalSig, altSig, "spellings must decode to identical signature bytes")
// The replay-cache key must be identical for both spellings.
assert.Equal(t, hashToken(canonical), hashToken(equivalent),
"noncanonical spelling must map to the same replay-cache key")
s := newTestSessionStore(t)
ctx := context.Background()
exp := time.Now().Add(time.Hour)
require.NoError(t, s.RegisterToken(ctx, canonical, exp), "first claim should succeed")
err = s.RegisterToken(ctx, equivalent, exp)
require.Error(t, err, "alternate spelling must be treated as a replay")
assert.ErrorIs(t, err, ErrTokenAlreadyUsed)
}
+33 -13
View File
@@ -166,9 +166,11 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use
return err
}
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}); err != nil {
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID})
if err != nil {
return err
}
change.ChangedPeerIDs = ipv6Changed
// A membership change does not alter which entities reference the group, so
// the dependency walk runs once against the post-change snapshot. The new
@@ -321,7 +323,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us
var globalErr error
for _, newGroup := range groups {
change := affectedpeers.Change{ChangedGroupIDs: []string{newGroup.ID}}
events, snap, err := am.updateSingleGroup(ctx, accountID, userID, newGroup, change)
events, snap, change, err := am.updateSingleGroup(ctx, accountID, userID, newGroup, change)
if err != nil {
log.WithContext(ctx).Errorf("failed to update group %s: %v", newGroup.ID, err)
if len(groups) == 1 {
@@ -344,7 +346,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us
return globalErr
}
func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountID, userID string, newGroup *types.Group, change affectedpeers.Change) ([]func(), *affectedpeers.Snapshot, error) {
func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountID, userID string, newGroup *types.Group, change affectedpeers.Change) ([]func(), *affectedpeers.Snapshot, affectedpeers.Change, error) {
var events []func()
var snap *affectedpeers.Snapshot
err := am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
@@ -364,9 +366,11 @@ func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountI
return err
}
if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}); err != nil {
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID})
if err != nil {
return err
}
change.ChangedPeerIDs = ipv6Changed
if err := transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
return err
@@ -377,7 +381,7 @@ func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountI
snap, err = affectedpeers.Load(ctx, transaction, accountID, change)
return err
})
return events, snap, err
return events, snap, change, err
}
// prepareGroupEvents prepares a list of event functions to be stored.
@@ -480,8 +484,8 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us
var allErrors error
var groupIDsToDelete []string
var deletedGroups []*types.Group
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
var snap, ipv6Snap *affectedpeers.Snapshot
var change, ipv6Change affectedpeers.Change
extraSettings, err := am.settingsManager.GetExtraSettings(ctx, accountID)
if err != nil {
@@ -510,10 +514,20 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us
return err
}
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, groupIDsToDelete); err != nil {
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, groupIDsToDelete)
if err != nil {
return err
}
// Members of a deleted IPv6-enabled group lose their address, which the
// pre-delete snapshot cannot see, so they are resolved post-delete.
if len(ipv6Changed) > 0 {
ipv6Change = affectedpeers.Change{ChangedPeerIDs: ipv6Changed}
if ipv6Snap, err = affectedpeers.Load(ctx, transaction, accountID, ipv6Change); err != nil {
return err
}
}
return transaction.IncrementNetworkSerial(ctx, accountID)
})
if err != nil {
@@ -524,7 +538,7 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us
am.StoreEvent(ctx, userID, group.ID, accountID, activity.GroupDeleted, group.EventMeta())
}
am.ExpandAndUpdateAffected(ctx, accountID, snap, change)
go am.dispatchAffected(ctx, accountID, []*affectedpeers.Snapshot{snap, ipv6Snap}, []affectedpeers.Change{change, ipv6Change})
return allErrors
}
@@ -564,11 +578,14 @@ func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, gr
return err
}
if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}); err != nil {
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID})
if err != nil {
return err
}
// A peer whose IPv6 address changed is visible to every peer that reaches it
// through any of its groups, not only through this one.
change.ChangedPeerIDs = ipv6Changed
var err error
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return err
}
@@ -634,11 +651,14 @@ func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID,
return err
}
if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}); err != nil {
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID})
if err != nil {
return err
}
// A peer whose IPv6 address changed is visible to every peer that reaches it
// through any of its groups, not only through this one.
change.ChangedPeerIDs = ipv6Changed
var err error
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return err
}
+4 -3
View File
@@ -58,10 +58,11 @@ import (
"github.com/netbirdio/netbird/management/server/networks/resources"
"github.com/netbirdio/netbird/management/server/networks/routers"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/shared/ratelimit"
)
// NewAPIHandler creates the Management service HTTP API handler registering all the available endpoints.
func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *middleware.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager) (http.Handler, error) {
func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *ratelimit.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager, proxyTokenRevocationGuard proxytoken.RevocationGuard) (http.Handler, error) {
// Register bypass paths for unauthenticated endpoints
if err := bypass.AddBypassPath("/api/instance"); err != nil {
@@ -84,7 +85,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou
if rateLimiter == nil {
log.Warn("NewAPIHandler: nil rate limiter, rate limiting disabled")
rateLimiter = middleware.NewAPIRateLimiter(nil)
rateLimiter = ratelimit.NewAPIRateLimiter(nil)
rateLimiter.SetEnabled(false)
}
@@ -135,7 +136,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou
reverseproxymanager.RegisterEndpoints(serviceManager, *reverseProxyDomainManager, reverseProxyAccessLogsManager, permissionsManager, router)
}
proxytoken.RegisterEndpoints(accountManager.GetStore(), permissionsManager, router)
proxytoken.RegisterEndpoints(accountManager.GetStore(), permissionsManager, proxyTokenRevocationGuard, router)
// Register OAuth callback handler for proxy authentication
if proxyGRPCServer != nil {
@@ -127,7 +127,7 @@ func (h *handler) validateNetworkRange(ctx context.Context, accountID, userID st
}
func (h *handler) validateCapacity(ctx context.Context, accountID, userID string, prefix netip.Prefix) error {
peers, err := h.accountManager.GetPeers(ctx, accountID, userID, "", "")
peers, err := h.accountManager.GetPeers(ctx, accountID, userID, "", "", "")
if err != nil {
return status.Errorf(status.Internal, "get peer count: %v", err)
}
@@ -58,7 +58,7 @@ func (h *handler) getAllGroups(w http.ResponseWriter, r *http.Request) {
return
}
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
if err != nil {
util.WriteError(r.Context(), err, w)
return
@@ -77,7 +77,7 @@ func (h *handler) getAllGroups(w http.ResponseWriter, r *http.Request) {
return
}
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
if err != nil {
util.WriteError(r.Context(), err, w)
return
@@ -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{
@@ -172,7 +169,7 @@ func (h *handler) updateGroup(w http.ResponseWriter, r *http.Request) {
return
}
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
if err != nil {
util.WriteError(r.Context(), err, w)
return
@@ -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{
@@ -232,7 +226,7 @@ func (h *handler) createGroup(w http.ResponseWriter, r *http.Request) {
return
}
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
if err != nil {
util.WriteError(r.Context(), err, w)
return
@@ -293,7 +287,7 @@ func (h *handler) getGroup(w http.ResponseWriter, r *http.Request) {
return
}
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
if err != nil {
util.WriteError(r.Context(), err, w)
return
@@ -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"
@@ -78,7 +78,7 @@ func initGroupTestData(initGroups ...*types.Group) *handler {
return nil, status.Errorf(status.NotFound, "unknown group name")
},
GetPeersFunc: func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
GetPeersFunc: func(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
return maps.Values(TestPeers), nil
},
DeleteGroupFunc: func(_ context.Context, accountID, userId, groupID string) error {
@@ -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
@@ -317,10 +317,11 @@ func (h *Handler) GetAllPeers(w http.ResponseWriter, r *http.Request) {
nameFilter := r.URL.Query().Get("name")
ipFilter := r.URL.Query().Get("ip")
macFilter := r.URL.Query().Get("mac")
accountID, userID := userAuth.AccountId, userAuth.UserId
peers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, nameFilter, ipFilter)
peers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, nameFilter, ipFilter, macFilter)
if err != nil {
util.WriteError(r.Context(), err, w)
return
@@ -571,6 +572,17 @@ func peerToAccessiblePeer(peer *nbpeer.Peer, dnsDomain string) api.AccessiblePee
}
}
func toNetworkAddresses(addrs []nbpeer.NetworkAddress) *[]api.NetworkAddress {
if len(addrs) == 0 {
return nil
}
out := make([]api.NetworkAddress, 0, len(addrs))
for _, a := range addrs {
out = append(out, api.NetworkAddress{NetIp: a.NetIP.String(), Mac: a.Mac})
}
return &out
}
func toSinglePeerResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dnsDomain string, approved bool, reason string) *api.Peer {
osVersion := peer.Meta.OSVersion
if osVersion == "" {
@@ -583,6 +595,7 @@ func toSinglePeerResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dnsD
Name: peer.Name,
Ip: peer.IP.String(),
Ipv6: peerIPv6String(peer),
NetworkAddresses: toNetworkAddresses(peer.Meta.NetworkAddresses),
ConnectionIp: peer.Location.ConnectionIP.String(),
Connected: peer.Status.Connected,
LastSeen: peer.Status.LastSeen,
@@ -639,6 +652,7 @@ func toPeerListItemResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dn
Name: peer.Name,
Ip: peer.IP.String(),
Ipv6: peerIPv6String(peer),
NetworkAddresses: toNetworkAddresses(peer.Meta.NetworkAddresses),
ConnectionIp: peer.Location.ConnectionIP.String(),
Connected: peer.Status.Connected,
LastSeen: peer.Status.LastSeen,
@@ -173,7 +173,7 @@ func initTestMetaData(t *testing.T, peers ...*nbpeer.Peer) *Handler {
return nil, fmt.Errorf("user not found")
}
},
GetPeersFunc: func(_ context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
GetPeersFunc: func(_ context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
return peers, nil
},
GetPeerGroupsFunc: func(ctx context.Context, accountID, peerID string) ([]*types.Group, error) {
@@ -364,6 +364,50 @@ func TestGetPeers(t *testing.T) {
}
}
func TestPeerResponseNetworkAddresses(t *testing.T) {
tests := []struct {
name string
addresses []nbpeer.NetworkAddress
wantJSON string
}{
{name: "not reported"},
{name: "empty", addresses: []nbpeer.NetworkAddress{}},
{
name: "multiple interfaces",
addresses: []nbpeer.NetworkAddress{
{NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"},
{NetIP: netip.MustParsePrefix("2001:db8::123/64"), Mac: "00:93:37:bd:83:10"},
},
wantJSON: `[{"net_ip":"192.168.0.11/24","mac":"00:93:37:bd:83:0f"},{"net_ip":"2001:db8::123/64","mac":"00:93:37:bd:83:10"}]`,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
peer := &nbpeer.Peer{
Status: &nbpeer.PeerStatus{},
Meta: nbpeer.PeerSystemMeta{NetworkAddresses: tt.addresses},
}
responses := map[string]any{
"single peer": toSinglePeerResponse(peer, nil, "example.com", true, ""),
"peer list": toPeerListItemResponse(peer, nil, "example.com", 0),
}
for name, response := range responses {
t.Run(name, func(t *testing.T) {
body, err := json.Marshal(response)
require.NoError(t, err)
var fields map[string]json.RawMessage
require.NoError(t, json.Unmarshal(body, &fields))
if tt.wantJSON == "" {
assert.NotContains(t, fields, "network_addresses", "unreported interfaces should be omitted")
return
}
assert.JSONEq(t, tt.wantJSON, string(fields["network_addresses"]), "response should preserve interface addresses and MACs")
})
}
})
}
}
func TestGetAccessiblePeers(t *testing.T) {
peer1 := &nbpeer.Peer{
ID: "peer1",
+16 -7
View File
@@ -16,21 +16,21 @@ import (
"golang.org/x/oauth2"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/http/middleware"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/shared/ratelimit"
)
// AuthCallbackHandler handles OAuth callbacks for proxy authentication.
type AuthCallbackHandler struct {
proxyService *nbgrpc.ProxyServiceServer
rateLimiter *middleware.APIRateLimiter
rateLimiter *ratelimit.APIRateLimiter
trustedProxies []netip.Prefix
}
// NewAuthCallbackHandler creates a new OAuth callback handler.
func NewAuthCallbackHandler(proxyService *nbgrpc.ProxyServiceServer, trustedProxies []netip.Prefix) *AuthCallbackHandler {
rateLimiterConfig := &middleware.RateLimiterConfig{
rateLimiterConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 10,
Burst: 15,
CleanupInterval: 5 * time.Minute,
@@ -39,7 +39,7 @@ func NewAuthCallbackHandler(proxyService *nbgrpc.ProxyServiceServer, trustedProx
return &AuthCallbackHandler{
proxyService: proxyService,
rateLimiter: middleware.NewAPIRateLimiter(rateLimiterConfig),
rateLimiter: ratelimit.NewAPIRateLimiter(rateLimiterConfig),
trustedProxies: trustedProxies,
}
}
@@ -59,7 +59,7 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
state := r.URL.Query().Get("state")
codeVerifier, originalURL, err := h.proxyService.ValidateState(state)
codeVerifier, originalURL, useSessionCode, err := h.proxyService.ValidateState(state)
if err != nil {
log.WithError(err).Error("OAuth callback state validation failed")
http.Error(w, "Invalid state parameter", http.StatusBadRequest)
@@ -119,10 +119,19 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
redirectURL.Scheme = "https"
query := redirectURL.Query()
query.Set("session_token", sessionToken)
if useSessionCode {
code, ok := h.proxyService.GenerateSessionCode(sessionToken)
if !ok {
http.Error(w, "Failed to create session", http.StatusInternalServerError)
return
}
query.Set(auth.SessionCodeQueryParam, code)
} else {
query.Set(auth.SessionTokenQueryParam, sessionToken)
}
redirectURL.RawQuery = query.Encode()
log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user with session token")
log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user to proxy")
http.Redirect(w, r, redirectURL.String(), http.StatusFound)
}
@@ -181,6 +181,10 @@ func (m *testAccessLogManager) GetAllAccessLogs(_ context.Context, _, _ string,
}
func setupAuthCallbackTest(t *testing.T) *testSetup {
return setupAuthCallbackTestWithProxyManager(t, testSessionCodeManager{})
}
func setupAuthCallbackTestWithProxyManager(t *testing.T, proxyManager nbproxy.Manager) *testSetup {
t.Helper()
ctx := context.Background()
@@ -197,7 +201,7 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
require.NoError(t, err)
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore)
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore)
usersManager := users.NewManager(testStore)
@@ -212,12 +216,12 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
proxyService := nbgrpc.NewProxyServiceServer(
&testAccessLogManager{},
tokenStore,
pkceStore,
singleUseStore,
oidcConfig,
nil,
usersManager,
nil,
nil,
proxyManager,
nil,
)
@@ -242,6 +246,15 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
}
}
type testSessionCodeManager struct {
nbproxy.Manager
supported bool
}
func (m testSessionCodeManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
return m.supported
}
func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store.Store) {
t.Helper()
@@ -252,10 +265,11 @@ func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store
privKey := base64.StdEncoding.EncodeToString(priv)
testProxy := &service.Service{
ID: "testProxyId",
AccountID: "testAccountId",
Name: "Test Proxy",
Domain: "test-proxy.example.com",
ID: "testProxyId",
AccountID: "testAccountId",
Name: "Test Proxy",
Domain: "test-proxy.example.com",
ProxyCluster: "cluster.example.com",
Targets: []*service.Target{{
Path: strPtr("/"),
Host: "localhost",
@@ -512,29 +526,56 @@ func createTestState(t *testing.T, ps *nbgrpc.ProxyServiceServer, redirectURL st
}
func TestAuthCallback_UserAllowedToLogin(t *testing.T) {
setup := setupAuthCallbackTest(t)
defer setup.cleanup()
tests := []struct {
name string
manager nbproxy.Manager
wantParam string
absentParam string
}{
{name: "legacy proxy", manager: testSessionCodeManager{}, wantParam: "session_token", absentParam: "nb_session_code"},
{name: "compatible proxy", manager: testSessionCodeManager{supported: true}, wantParam: "nb_session_code", absentParam: "session_token"},
}
setup.oidcServer.tokenSubject = "allowedUserId"
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
setup := setupAuthCallbackTestWithProxyManager(t, tt.manager)
defer setup.cleanup()
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
setup.oidcServer.tokenSubject = "allowedUserId"
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
rec := httptest.NewRecorder()
setup.router.ServeHTTP(rec, req)
require.Equal(t, http.StatusFound, rec.Code)
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
rec := httptest.NewRecorder()
location, err := url.Parse(rec.Header().Get("Location"))
require.NoError(t, err)
require.Equal(t, "test-proxy.example.com", location.Host)
require.NotEmpty(t, location.Query().Get(tt.wantParam))
require.Empty(t, location.Query().Get(tt.absentParam))
require.Empty(t, location.Query().Get("error"))
setup.router.ServeHTTP(rec, req)
if tt.wantParam == "nb_session_code" {
code := location.Query().Get("nb_session_code")
response, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: location.Hostname(),
SessionCode: code,
})
require.NoError(t, err)
require.True(t, response.GetValid())
require.NotEmpty(t, response.GetSessionToken())
require.NotEqual(t, code, response.GetSessionToken())
require.Equal(t, http.StatusFound, rec.Code)
location := rec.Header().Get("Location")
require.NotEmpty(t, location)
parsedLocation, err := url.Parse(location)
require.NoError(t, err)
require.Equal(t, "test-proxy.example.com", parsedLocation.Host)
require.NotEmpty(t, parsedLocation.Query().Get("session_token"), "Should include session token")
require.Empty(t, parsedLocation.Query().Get("error"), "Should not have error parameter")
replayed, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: location.Hostname(),
SessionCode: code,
})
require.NoError(t, err)
require.False(t, replayed.GetValid())
require.Empty(t, replayed.GetSessionToken())
}
})
}
}
// TestAuthCallback_UserDeniedByAccountStatus asserts that a user whose account
@@ -11,15 +11,15 @@ import (
"github.com/netbirdio/netbird/management/server/account"
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/management/server/http/middleware"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/http/api"
"github.com/netbirdio/netbird/shared/management/http/util"
"github.com/netbirdio/netbird/shared/management/status"
"github.com/netbirdio/netbird/shared/ratelimit"
)
// publicInviteRateLimiter limits public invite requests by IP address to prevent brute-force attacks
var publicInviteRateLimiter = middleware.NewAPIRateLimiter(&middleware.RateLimiterConfig{
var publicInviteRateLimiter = ratelimit.NewAPIRateLimiter(&ratelimit.RateLimiterConfig{
RequestsPerMinute: 10, // 10 attempts per minute per IP
Burst: 5, // Allow burst of 5 requests
CleanupInterval: 10 * time.Minute,
@@ -18,6 +18,7 @@ import (
"github.com/netbirdio/netbird/shared/auth"
"github.com/netbirdio/netbird/shared/management/http/util"
"github.com/netbirdio/netbird/shared/management/status"
"github.com/netbirdio/netbird/shared/ratelimit"
)
type EnsureAccountFunc func(ctx context.Context, userAuth auth.UserAuth) (string, string, error)
@@ -33,7 +34,7 @@ type AuthMiddleware struct {
ensureAccount EnsureAccountFunc
getUserFromUserAuth GetUserFromUserAuthFunc
syncUserJWTGroups SyncUserJWTGroupsFunc
rateLimiter *APIRateLimiter
rateLimiter *ratelimit.APIRateLimiter
patUsageTracker *PATUsageTracker
isValidChildAccount IsValidChildAccountFunc
}
@@ -44,7 +45,7 @@ func NewAuthMiddleware(
ensureAccount EnsureAccountFunc,
syncUserJWTGroups SyncUserJWTGroupsFunc,
getUserFromUserAuth GetUserFromUserAuthFunc,
rateLimiter *APIRateLimiter,
rateLimiter *ratelimit.APIRateLimiter,
meter metric.Meter,
isValidChildAccount IsValidChildAccountFunc,
) *AuthMiddleware {
@@ -18,6 +18,7 @@ import (
"github.com/netbirdio/netbird/management/server/util"
nbauth "github.com/netbirdio/netbird/shared/auth"
nbjwt "github.com/netbirdio/netbird/shared/auth/jwt"
"github.com/netbirdio/netbird/shared/ratelimit"
)
const (
@@ -196,7 +197,7 @@ func TestAuthMiddleware_Handler(t *testing.T) {
GetPATInfoFunc: mockGetAccountInfoFromPAT,
}
disabledLimiter := NewAPIRateLimiter(nil)
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
disabledLimiter.SetEnabled(false)
authMiddleware := NewAuthMiddleware(
mockAuth,
@@ -260,7 +261,7 @@ func TestAuthMiddleware_SyncUserJWTGroupsDetachedFromRequestCancellation(t *test
GetPATInfoFunc: mockGetAccountInfoFromPAT,
}
disabledLimiter := NewAPIRateLimiter(nil)
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
disabledLimiter.SetEnabled(false)
authMiddleware := NewAuthMiddleware(
@@ -311,7 +312,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("PAT Token Rate Limiting - Burst Works", func(t *testing.T) {
// Configure rate limiter: 10 requests per minute with burst of 5
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 10,
Burst: 5,
CleanupInterval: 5 * time.Minute,
@@ -329,7 +330,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -364,7 +365,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("PAT Token Rate Limiting - Rate Limit Enforced", func(t *testing.T) {
// Configure very low rate limit: 1 request per minute
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -382,7 +383,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -408,7 +409,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("Bearer Token Not Rate Limited", func(t *testing.T) {
// Configure strict rate limit
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -426,7 +427,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -453,7 +454,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("PAT Token Rate Limiting Per Token", func(t *testing.T) {
// Configure rate limiter
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -471,7 +472,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -518,7 +519,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("Rate Limiter Cleanup", func(t *testing.T) {
// Configure rate limiter with short cleanup interval and TTL for testing
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 1,
CleanupInterval: 100 * time.Millisecond,
@@ -536,7 +537,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -578,7 +579,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
})
t.Run("Terraform User Agent Not Rate Limited", func(t *testing.T) {
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -596,7 +597,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -634,7 +635,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
})
t.Run("Non-Terraform User Agent With PAT Is Rate Limited", func(t *testing.T) {
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -652,7 +653,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -740,7 +741,7 @@ func TestAuthMiddleware_Handler_Child(t *testing.T) {
GetPATInfoFunc: mockGetAccountInfoFromPAT,
}
disabledLimiter := NewAPIRateLimiter(nil)
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
disabledLimiter.SetEnabled(false)
authMiddleware := NewAuthMiddleware(
mockAuth,
@@ -1,269 +0,0 @@
package middleware
import (
"context"
"net"
"net/http"
"os"
"strconv"
"sync"
"sync/atomic"
"time"
log "github.com/sirupsen/logrus"
"golang.org/x/time/rate"
"github.com/netbirdio/netbird/shared/management/http/util"
)
const (
RateLimitingEnabledEnv = "NB_API_RATE_LIMITING_ENABLED"
RateLimitingBurstEnv = "NB_API_RATE_LIMITING_BURST"
RateLimitingRPMEnv = "NB_API_RATE_LIMITING_RPM"
defaultAPIRPM = 6
defaultAPIBurst = 500
)
// RateLimiterConfig holds configuration for the API rate limiter
type RateLimiterConfig struct {
// RequestsPerMinute defines the rate at which tokens are replenished
RequestsPerMinute float64
// Burst defines the maximum number of requests that can be made in a burst
Burst int
// CleanupInterval defines how often to clean up old limiters (how often garbage collection runs)
CleanupInterval time.Duration
// LimiterTTL defines how long a limiter should be kept after last use (age threshold for removal)
LimiterTTL time.Duration
}
// DefaultRateLimiterConfig returns a default configuration
func DefaultRateLimiterConfig() *RateLimiterConfig {
return &RateLimiterConfig{
RequestsPerMinute: 100,
Burst: 120,
CleanupInterval: 5 * time.Minute,
LimiterTTL: 10 * time.Minute,
}
}
func RateLimiterConfigFromEnv() (cfg *RateLimiterConfig, enabled bool) {
rpm := defaultAPIRPM
if v := os.Getenv(RateLimitingRPMEnv); v != "" {
value, err := strconv.Atoi(v)
if err != nil {
log.Warnf("parsing %s env var: %v, using default %d", RateLimitingRPMEnv, err, rpm)
} else {
rpm = value
}
}
if rpm <= 0 {
log.Warnf("%s=%d is non-positive, using default %d", RateLimitingRPMEnv, rpm, defaultAPIRPM)
rpm = defaultAPIRPM
}
burst := defaultAPIBurst
if v := os.Getenv(RateLimitingBurstEnv); v != "" {
value, err := strconv.Atoi(v)
if err != nil {
log.Warnf("parsing %s env var: %v, using default %d", RateLimitingBurstEnv, err, burst)
} else {
burst = value
}
}
if burst <= 0 {
log.Warnf("%s=%d is non-positive, using default %d", RateLimitingBurstEnv, burst, defaultAPIBurst)
burst = defaultAPIBurst
}
return &RateLimiterConfig{
RequestsPerMinute: float64(rpm),
Burst: burst,
CleanupInterval: 6 * time.Hour,
LimiterTTL: 24 * time.Hour,
}, os.Getenv(RateLimitingEnabledEnv) == "true"
}
// limiterEntry holds a rate limiter and its last access time
type limiterEntry struct {
limiter *rate.Limiter
lastAccess time.Time
}
// APIRateLimiter manages rate limiting for API tokens
type APIRateLimiter struct {
config *RateLimiterConfig
limiters map[string]*limiterEntry
mu sync.RWMutex
stopChan chan struct{}
enabled atomic.Bool
}
// NewAPIRateLimiter creates a new API rate limiter with the given configuration
func NewAPIRateLimiter(config *RateLimiterConfig) *APIRateLimiter {
if config == nil {
config = DefaultRateLimiterConfig()
}
rl := &APIRateLimiter{
config: config,
limiters: make(map[string]*limiterEntry),
stopChan: make(chan struct{}),
}
rl.enabled.Store(true)
go rl.cleanupLoop()
return rl
}
func (rl *APIRateLimiter) SetEnabled(enabled bool) {
rl.enabled.Store(enabled)
}
func (rl *APIRateLimiter) Enabled() bool {
return rl.enabled.Load()
}
func (rl *APIRateLimiter) UpdateConfig(config *RateLimiterConfig) {
if config == nil {
return
}
if config.RequestsPerMinute <= 0 || config.Burst <= 0 {
log.Warnf("UpdateConfig: ignoring invalid rpm=%v burst=%d", config.RequestsPerMinute, config.Burst)
return
}
newRPS := rate.Limit(config.RequestsPerMinute / 60.0)
newBurst := config.Burst
rl.mu.Lock()
rl.config.RequestsPerMinute = config.RequestsPerMinute
rl.config.Burst = newBurst
snapshot := make([]*rate.Limiter, 0, len(rl.limiters))
for _, entry := range rl.limiters {
snapshot = append(snapshot, entry.limiter)
}
rl.mu.Unlock()
for _, l := range snapshot {
l.SetLimit(newRPS)
l.SetBurst(newBurst)
}
}
// Allow checks if a request for the given key (token) is allowed
func (rl *APIRateLimiter) Allow(key string) bool {
if !rl.enabled.Load() {
return true
}
limiter := rl.getLimiter(key)
return limiter.Allow()
}
// Wait blocks until the rate limiter allows another request for the given key
// Returns an error if the context is canceled
func (rl *APIRateLimiter) Wait(ctx context.Context, key string) error {
if !rl.enabled.Load() {
return nil
}
limiter := rl.getLimiter(key)
return limiter.Wait(ctx)
}
// getLimiter retrieves or creates a rate limiter for the given key
func (rl *APIRateLimiter) getLimiter(key string) *rate.Limiter {
rl.mu.RLock()
entry, exists := rl.limiters[key]
rl.mu.RUnlock()
if exists {
rl.mu.Lock()
entry.lastAccess = time.Now()
rl.mu.Unlock()
return entry.limiter
}
rl.mu.Lock()
defer rl.mu.Unlock()
if entry, exists := rl.limiters[key]; exists {
entry.lastAccess = time.Now()
return entry.limiter
}
requestsPerSecond := rl.config.RequestsPerMinute / 60.0
limiter := rate.NewLimiter(rate.Limit(requestsPerSecond), rl.config.Burst)
rl.limiters[key] = &limiterEntry{
limiter: limiter,
lastAccess: time.Now(),
}
return limiter
}
// cleanupLoop periodically removes old limiters that haven't been used recently
func (rl *APIRateLimiter) cleanupLoop() {
ticker := time.NewTicker(rl.config.CleanupInterval)
defer ticker.Stop()
for {
select {
case <-ticker.C:
rl.cleanup()
case <-rl.stopChan:
return
}
}
}
// cleanup removes limiters that haven't been used within the TTL period
func (rl *APIRateLimiter) cleanup() {
rl.mu.Lock()
defer rl.mu.Unlock()
now := time.Now()
for key, entry := range rl.limiters {
if now.Sub(entry.lastAccess) > rl.config.LimiterTTL {
delete(rl.limiters, key)
}
}
}
// Stop stops the cleanup goroutine
func (rl *APIRateLimiter) Stop() {
close(rl.stopChan)
}
// Reset removes the rate limiter for a specific key
func (rl *APIRateLimiter) Reset(key string) {
rl.mu.Lock()
defer rl.mu.Unlock()
delete(rl.limiters, key)
}
// Middleware returns an HTTP middleware that rate limits requests by client IP.
// Returns 429 Too Many Requests if the rate limit is exceeded.
func (rl *APIRateLimiter) Middleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if !rl.enabled.Load() {
next.ServeHTTP(w, r)
return
}
clientIP := getClientIP(r)
if !rl.Allow(clientIP) {
util.WriteErrorResponse("rate limit exceeded, please try again later", http.StatusTooManyRequests, w)
return
}
next.ServeHTTP(w, r)
})
}
// getClientIP extracts the client IP address from the request.
func getClientIP(r *http.Request) string {
ip, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return r.RemoteAddr
}
return ip
}
@@ -1,329 +0,0 @@
package middleware
import (
"fmt"
"net/http"
"net/http/httptest"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
)
func TestAPIRateLimiter_Allow(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 60, // 1 per second
Burst: 2,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
// First two requests should be allowed (burst)
assert.True(t, rl.Allow("test-key"))
assert.True(t, rl.Allow("test-key"))
// Third request should be denied (exceeded burst)
assert.False(t, rl.Allow("test-key"))
// Different key should be allowed
assert.True(t, rl.Allow("different-key"))
}
func TestAPIRateLimiter_Middleware(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 60, // 1 per second
Burst: 2,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
// Create a simple handler that returns 200 OK
nextHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
})
// Wrap with rate limiter middleware
handler := rl.Middleware(nextHandler)
// First two requests should pass (burst)
for i := 0; i < 2; i++ {
req := httptest.NewRequest(http.MethodGet, "/test", nil)
req.RemoteAddr = "192.168.1.1:12345"
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
assert.Equal(t, http.StatusOK, rr.Code, "request %d should be allowed", i+1)
}
// Third request should be rate limited
req := httptest.NewRequest(http.MethodGet, "/test", nil)
req.RemoteAddr = "192.168.1.1:12345"
rr := httptest.NewRecorder()
handler.ServeHTTP(rr, req)
assert.Equal(t, http.StatusTooManyRequests, rr.Code)
}
func TestAPIRateLimiter_Middleware_DifferentIPs(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 1,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
nextHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
})
handler := rl.Middleware(nextHandler)
// Request from first IP
req1 := httptest.NewRequest(http.MethodGet, "/test", nil)
req1.RemoteAddr = "192.168.1.1:12345"
rr1 := httptest.NewRecorder()
handler.ServeHTTP(rr1, req1)
assert.Equal(t, http.StatusOK, rr1.Code)
// Second request from first IP should be rate limited
req2 := httptest.NewRequest(http.MethodGet, "/test", nil)
req2.RemoteAddr = "192.168.1.1:12345"
rr2 := httptest.NewRecorder()
handler.ServeHTTP(rr2, req2)
assert.Equal(t, http.StatusTooManyRequests, rr2.Code)
// Request from different IP should be allowed
req3 := httptest.NewRequest(http.MethodGet, "/test", nil)
req3.RemoteAddr = "192.168.1.2:12345"
rr3 := httptest.NewRecorder()
handler.ServeHTTP(rr3, req3)
assert.Equal(t, http.StatusOK, rr3.Code)
}
func TestGetClientIP(t *testing.T) {
tests := []struct {
name string
remoteAddr string
expected string
}{
{
name: "remote addr with port",
remoteAddr: "192.168.1.1:12345",
expected: "192.168.1.1",
},
{
name: "remote addr without port",
remoteAddr: "192.168.1.1",
expected: "192.168.1.1",
},
{
name: "IPv6 with port",
remoteAddr: "[::1]:12345",
expected: "::1",
},
{
name: "IPv6 without port",
remoteAddr: "::1",
expected: "::1",
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/test", nil)
req.RemoteAddr = tc.remoteAddr
assert.Equal(t, tc.expected, getClientIP(req))
})
}
}
func TestAPIRateLimiter_Reset(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 1,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
// Use up the burst
assert.True(t, rl.Allow("test-key"))
assert.False(t, rl.Allow("test-key"))
// Reset the limiter
rl.Reset("test-key")
// Should be allowed again
assert.True(t, rl.Allow("test-key"))
}
func TestAPIRateLimiter_SetEnabled(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 1,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
assert.True(t, rl.Allow("key"))
assert.False(t, rl.Allow("key"), "burst exhausted while enabled")
rl.SetEnabled(false)
assert.False(t, rl.Enabled())
for i := 0; i < 5; i++ {
assert.True(t, rl.Allow("key"), "disabled limiter must always allow")
}
rl.SetEnabled(true)
assert.True(t, rl.Enabled())
assert.False(t, rl.Allow("key"), "re-enabled limiter retains prior bucket state")
}
func TestAPIRateLimiter_UpdateConfig(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 2,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
assert.True(t, rl.Allow("k1"))
assert.True(t, rl.Allow("k1"))
assert.False(t, rl.Allow("k1"), "burst=2 exhausted")
rl.UpdateConfig(&RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 10,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
// New burst applies to existing keys in place; bucket refills up to new burst over time,
// but importantly newly-added keys use the updated config immediately.
assert.True(t, rl.Allow("k2"))
for i := 0; i < 9; i++ {
assert.True(t, rl.Allow("k2"))
}
assert.False(t, rl.Allow("k2"), "new burst=10 exhausted")
}
func TestAPIRateLimiter_UpdateConfig_NilIgnored(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 1,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
rl.UpdateConfig(nil) // must not panic or zero the config
assert.True(t, rl.Allow("k"))
assert.False(t, rl.Allow("k"))
}
func TestAPIRateLimiter_UpdateConfig_NonPositiveIgnored(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 1,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
assert.True(t, rl.Allow("k"))
assert.False(t, rl.Allow("k"))
rl.UpdateConfig(&RateLimiterConfig{RequestsPerMinute: 0, Burst: 0, CleanupInterval: time.Minute, LimiterTTL: time.Minute})
rl.UpdateConfig(&RateLimiterConfig{RequestsPerMinute: -1, Burst: 5, CleanupInterval: time.Minute, LimiterTTL: time.Minute})
rl.UpdateConfig(&RateLimiterConfig{RequestsPerMinute: 60, Burst: -1, CleanupInterval: time.Minute, LimiterTTL: time.Minute})
rl.Reset("k")
assert.True(t, rl.Allow("k"))
assert.False(t, rl.Allow("k"), "burst should still be 1 — invalid UpdateConfig calls were ignored")
}
func TestAPIRateLimiter_ConcurrentAllowAndUpdate(t *testing.T) {
rl := NewAPIRateLimiter(&RateLimiterConfig{
RequestsPerMinute: 600,
Burst: 10,
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
defer rl.Stop()
var wg sync.WaitGroup
stop := make(chan struct{})
for i := 0; i < 8; i++ {
wg.Add(1)
go func(id int) {
defer wg.Done()
key := fmt.Sprintf("k%d", id)
for {
select {
case <-stop:
return
default:
rl.Allow(key)
}
}
}(i)
}
wg.Add(1)
go func() {
defer wg.Done()
for i := 0; i < 200; i++ {
select {
case <-stop:
return
default:
rl.UpdateConfig(&RateLimiterConfig{
RequestsPerMinute: float64(30 + (i % 90)),
Burst: 1 + (i % 20),
CleanupInterval: time.Minute,
LimiterTTL: time.Minute,
})
rl.SetEnabled(i%2 == 0)
}
}
}()
time.Sleep(100 * time.Millisecond)
close(stop)
wg.Wait()
}
func TestRateLimiterConfigFromEnv(t *testing.T) {
t.Setenv(RateLimitingEnabledEnv, "true")
t.Setenv(RateLimitingRPMEnv, "42")
t.Setenv(RateLimitingBurstEnv, "7")
cfg, enabled := RateLimiterConfigFromEnv()
assert.True(t, enabled)
assert.Equal(t, float64(42), cfg.RequestsPerMinute)
assert.Equal(t, 7, cfg.Burst)
t.Setenv(RateLimitingEnabledEnv, "false")
_, enabled = RateLimiterConfigFromEnv()
assert.False(t, enabled)
t.Setenv(RateLimitingEnabledEnv, "")
t.Setenv(RateLimitingRPMEnv, "")
t.Setenv(RateLimitingBurstEnv, "")
cfg, enabled = RateLimiterConfigFromEnv()
assert.False(t, enabled)
assert.Equal(t, float64(defaultAPIRPM), cfg.RequestsPerMinute)
assert.Equal(t, defaultAPIBurst, cfg.Burst)
t.Setenv(RateLimitingRPMEnv, "0")
t.Setenv(RateLimitingBurstEnv, "-5")
cfg, _ = RateLimiterConfigFromEnv()
assert.Equal(t, float64(defaultAPIRPM), cfg.RequestsPerMinute, "non-positive rpm must fall back to default")
assert.Equal(t, defaultAPIBurst, cfg.Burst, "non-positive burst must fall back to default")
}
@@ -46,14 +46,14 @@ import (
"github.com/netbirdio/netbird/management/server/networks/routers"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/management/server/store"
nbstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/management/server/users"
"github.com/netbirdio/netbird/shared/auth"
)
func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPeerUpdate *network_map.UpdateMessage, validateUpdate bool) (http.Handler, account.Manager, chan struct{}) {
store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir())
store, cleanup, err := nbstore.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir())
if err != nil {
t.Fatalf("Failed to create test store: %v", err)
}
@@ -108,15 +108,15 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee
t.Fatalf("Failed to create manager: %v", err)
}
accessLogsManager := accesslogsmanager.NewManager(store, permissionsManager, nil)
accessLogsManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(store.(*nbstore.SqlStore).Conn()), store, permissionsManager, nil)
proxyTokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore)
pkceverifierStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore)
noopMeter := noop.NewMeterProvider().Meter("")
proxyMgr, err := proxymanager.NewManager(store, noopMeter)
if err != nil {
t.Fatalf("Failed to create proxy manager: %v", err)
}
proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, pkceverifierStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil)
proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil)
// NewProxyServiceServer starts cleanupStaleProxies on a context it derives
// from context.Background(), independent of the cancellable ctx above;
// Close() cancels it so the goroutine does not outlive the test.
@@ -147,7 +147,7 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter()
apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil)
apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil, nil)
if err != nil {
t.Fatalf("Failed to create API handler: %v", err)
}
@@ -204,7 +204,7 @@ func PeerShouldNotReceiveAnyUpdate(t testing_tools.TB, updateMessage <-chan *net
// BuildApiBlackBoxWithDBStateAndPeerChannel creates the API handler and returns
// the peer update channel directly so tests can verify updates inline.
func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile string) (http.Handler, account.Manager, <-chan *network_map.UpdateMessage) {
store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir())
store, cleanup, err := nbstore.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir())
if err != nil {
t.Fatalf("Failed to create test store: %v", err)
}
@@ -248,15 +248,15 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin
t.Fatalf("Failed to create manager: %v", err)
}
accessLogsManager := accesslogsmanager.NewManager(store, permissionsManager, nil)
accessLogsManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(store.(*nbstore.SqlStore).Conn()), store, permissionsManager, nil)
proxyTokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore)
pkceverifierStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore)
noopMeter := noop.NewMeterProvider().Meter("")
proxyMgr, err := proxymanager.NewManager(store, noopMeter)
if err != nil {
t.Fatalf("Failed to create proxy manager: %v", err)
}
proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, pkceverifierStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil)
proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil)
// NewProxyServiceServer starts cleanupStaleProxies on a context it derives
// from context.Background(), independent of the cancellable ctx above;
// Close() cancels it so the goroutine does not outlive the test.
@@ -287,7 +287,7 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter()
apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil)
apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil, nil)
if err != nil {
t.Fatalf("Failed to create API handler: %v", err)
}
+1 -7
View File
@@ -132,13 +132,7 @@ type ConnectionOptions struct {
// NewAuth0Manager creates a new instance of the Auth0Manager
func NewAuth0Manager(config Auth0ClientConfig, appMetrics telemetry.AppMetrics) (*Auth0Manager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
+1 -7
View File
@@ -49,13 +49,7 @@ type AuthentikCredentials struct {
// NewAuthentikManager creates a new instance of the AuthentikManager.
func NewAuthentikManager(config AuthentikClientConfig, appMetrics telemetry.AppMetrics) (*AuthentikManager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
+1 -7
View File
@@ -54,13 +54,7 @@ type azureProfile map[string]any
// NewAzureManager creates a new instance of the AzureManager.
func NewAzureManager(config AzureClientConfig, appMetrics telemetry.AppMetrics) (*AzureManager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
+1 -9
View File
@@ -4,10 +4,8 @@ import (
"context"
"encoding/base64"
"fmt"
"net/http"
"strings"
"sync"
"time"
"github.com/dexidp/dex/api/v2"
log "github.com/sirupsen/logrus"
@@ -44,13 +42,7 @@ func NewDexManager(config DexClientConfig, appMetrics telemetry.AppMetrics) (*De
return nil, fmt.Errorf("dex IdP configuration is incomplete, GRPCAddr is missing")
}
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: 10 * time.Second,
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
return &DexManager{
+1 -8
View File
@@ -4,7 +4,6 @@ import (
"context"
"encoding/base64"
"fmt"
"net/http"
log "github.com/sirupsen/logrus"
"golang.org/x/oauth2/google"
@@ -44,13 +43,7 @@ func (gc *GoogleWorkspaceCredentials) Authenticate(_ context.Context) (JWTToken,
// NewGoogleWorkspaceManager creates a new instance of the GoogleWorkspaceManager.
func NewGoogleWorkspaceManager(ctx context.Context, config GoogleWorkspaceClientConfig, appMetrics telemetry.AppMetrics) (*GoogleWorkspaceManager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
+1 -7
View File
@@ -58,13 +58,7 @@ type JumpCloudCredentials struct {
// NewJumpCloudManager creates a new instance of the JumpCloudManager.
func NewJumpCloudManager(config JumpCloudClientConfig, appMetrics telemetry.AppMetrics) (*JumpCloudManager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
+1 -7
View File
@@ -59,13 +59,7 @@ type keycloakProfile struct {
// NewKeycloakManager creates a new instance of the KeycloakManager.
func NewKeycloakManager(config KeycloakClientConfig, appMetrics telemetry.AppMetrics) (*KeycloakManager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
+1 -7
View File
@@ -40,13 +40,7 @@ type OktaCredentials struct {
// NewOktaManager creates a new instance of the OktaManager.
func NewOktaManager(config OktaClientConfig, appMetrics telemetry.AppMetrics) (*OktaManager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
config.Issuer = baseURL(config.Issuer)
+1 -7
View File
@@ -83,13 +83,7 @@ type pocketIdUserGroupDto struct {
}
func NewPocketIdManager(config PocketIdClientConfig, appMetrics telemetry.AppMetrics) (*PocketIdManager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
+19
View File
@@ -2,6 +2,8 @@ package idp
import (
"encoding/json"
"errors"
"net/http"
"net/url"
"os"
"strings"
@@ -81,6 +83,23 @@ const (
defaultTimeout = 10 * time.Second
)
// errRedirectRefused is returned instead of http.ErrUseLastResponse so the
// client closes the redirect response rather than handing it back unread.
var errRedirectRefused = errors.New("redirect refused")
func newHTTPClient() *http.Client {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
return &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
CheckRedirect: func(*http.Request, []*http.Request) error {
return errRedirectRefused
},
}
}
// idpTimeout returns a timeout value for the IDP
func idpTimeout() time.Duration {
timeoutStr, ok := os.LookupEnv(idpTimeoutEnv)
+1 -7
View File
@@ -160,13 +160,7 @@ func verifyJWTConfig(config ZitadelClientConfig) error {
// NewZitadelManager creates a new instance of the ZitadelManager.
func NewZitadelManager(config ZitadelClientConfig, appMetrics telemetry.AppMetrics) (*ZitadelManager, error) {
httpTransport := http.DefaultTransport.(*http.Transport).Clone()
httpTransport.MaxIdleConns = 5
httpClient := &http.Client{
Timeout: idpTimeout(),
Transport: httpTransport,
}
httpClient := newHTTPClient()
helper := JsonParser{}
+1 -1
View File
@@ -100,7 +100,7 @@ func (am *DefaultAccountManager) GetValidatedPeers(ctx context.Context, accountI
return nil, nil, err
}
peers, err = am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
peers, err = am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
if err != nil {
return nil, nil, err
}
@@ -39,7 +39,7 @@ type MockAccountManager struct {
GetAccountIDByUserIdFunc func(ctx context.Context, userAuth auth.UserAuth) (string, error)
GetUserFromUserAuthFunc func(ctx context.Context, userAuth auth.UserAuth) (*types.User, error)
ListUsersFunc func(ctx context.Context, accountID string) ([]*types.User, error)
GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error)
GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error)
MarkPeerConnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error
MarkPeerDisconnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error
SyncAndMarkPeerFunc func(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error)
@@ -807,9 +807,9 @@ func (am *MockAccountManager) GetAccountIDFromUserAuth(ctx context.Context, user
}
// GetPeers mocks GetPeers of the AccountManager interface
func (am *MockAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
func (am *MockAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
if am.GetPeersFunc != nil {
return am.GetPeersFunc(ctx, accountID, userID, nameFilter, ipFilter)
return am.GetPeersFunc(ctx, accountID, userID, nameFilter, ipFilter, macFilter)
}
return nil, status.Errorf(codes.Unimplemented, "method GetPeers is not implemented")
}
@@ -32,7 +32,7 @@ type NetworkResource struct {
ID string `gorm:"primaryKey"`
NetworkID string `gorm:"index"`
AccountID string `gorm:"index"`
PublicID string `json:"-"`
PublicID string `json:"-" gorm:"index"`
Name string
Description string
Type NetworkResourceType
+2 -2
View File
@@ -47,7 +47,7 @@ const (
// GetPeers returns peers visible to the user within an account.
// Users with "peers:read" see all peers. Otherwise, users see only their own peers, or none if restricted by account settings.
func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
user, err := am.Store.GetUserByUserID(ctx, store.LockingStrengthNone, userID)
if err != nil {
return nil, err
@@ -59,7 +59,7 @@ func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID
}
if allowed {
return am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, nameFilter, ipFilter)
return am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, nameFilter, ipFilter, macFilter)
}
settings, err := am.Store.GetAccountSettings(ctx, store.LockingStrengthNone, accountID)
+74 -2
View File
@@ -4,10 +4,14 @@ import (
"context"
"crypto/sha256"
b64 "encoding/base64"
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"os"
"runtime"
"strconv"
@@ -33,12 +37,15 @@ import (
"github.com/netbirdio/netbird/management/internals/server/config"
"github.com/netbirdio/netbird/management/internals/shared/grpc"
nbcache "github.com/netbirdio/netbird/management/server/cache"
nbcontext "github.com/netbirdio/netbird/management/server/context"
peershandler "github.com/netbirdio/netbird/management/server/http/handlers/peers"
"github.com/netbirdio/netbird/management/server/http/testing/testing_tools"
"github.com/netbirdio/netbird/management/server/integrations/port_forwarding"
"github.com/netbirdio/netbird/management/server/job"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/shared/auth"
"github.com/netbirdio/netbird/shared/management/http/api"
"github.com/netbirdio/netbird/shared/management/status"
"github.com/netbirdio/netbird/management/server/util"
@@ -718,7 +725,7 @@ func TestDefaultAccountManager_GetPeers(t *testing.T) {
return
}
peers, err := manager.GetPeers(context.Background(), accountID, someUser, "", "")
peers, err := manager.GetPeers(context.Background(), accountID, someUser, "", "", "")
if err != nil {
t.Fatal(err)
return
@@ -731,6 +738,71 @@ func TestDefaultAccountManager_GetPeers(t *testing.T) {
}
}
func TestDefaultAccountManager_GetPeers_FilterByMac(t *testing.T) {
ctx := context.Background()
manager, _, err := createManager(t)
require.NoError(t, err)
account := newAccountWithId(ctx, "mac-account", "mac-admin", "", "", "", false)
account.Peers["matching"] = &nbpeer.Peer{
ID: "matching", Key: "matching-key", Name: "laptop", DNSLabel: "laptop",
IP: netip.MustParseAddr("100.64.0.10"), Status: &nbpeer.PeerStatus{LastSeen: time.Now()},
Meta: nbpeer.PeerSystemMeta{NetworkAddresses: []nbpeer.NetworkAddress{
{NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"},
{NetIP: netip.MustParsePrefix("192.168.1.11/24"), Mac: "aa:bb:cc:dd:ee:ff"},
}},
}
account.Peers["other"] = &nbpeer.Peer{
ID: "other", Key: "other-key", Name: "desktop", DNSLabel: "desktop",
IP: netip.MustParseAddr("100.64.0.20"), Status: &nbpeer.PeerStatus{LastSeen: time.Now()},
}
require.NoError(t, manager.Store.SaveAccount(ctx, account))
otherAccount := newAccountWithId(ctx, "other-account", "other-admin", "", "", "", false)
otherPeer := account.Peers["matching"].Copy()
otherPeer.ID, otherPeer.Key = "outside-account", "outside-key"
otherAccount.Peers[otherPeer.ID] = otherPeer
require.NoError(t, manager.Store.SaveAccount(ctx, otherAccount))
handler := peershandler.NewHandler(manager, manager.networkMapController, manager.permissionsManager)
tests := []struct {
name, nameFilter, ipFilter, macFilter string
wantIDs []string
}{
{name: "no filter", wantIDs: []string{"matching", "other"}},
{name: "full MAC", macFilter: "00:93:37:bd:83:0f", wantIDs: []string{"matching"}},
{name: "partial MAC", macFilter: "93:37:bd", wantIDs: []string{"matching"}},
{name: "second interface", macFilter: "aa:bb:cc:dd:ee:ff", wantIDs: []string{"matching"}},
{name: "unknown MAC", macFilter: "11:22:33:44:55:66"},
{name: "combined filters", nameFilter: "laptop", ipFilter: "100.64.0.10", macFilter: "00:93:37", wantIDs: []string{"matching"}},
{name: "name mismatch", nameFilter: "desktop", macFilter: "00:93:37"},
{name: "IP mismatch", ipFilter: "100.64.0.20", macFilter: "00:93:37"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
peers, err := manager.GetPeers(ctx, account.Id, "mac-admin", tt.nameFilter, tt.ipFilter, tt.macFilter)
require.NoError(t, err)
ids := make([]string, 0, len(peers))
for _, peer := range peers {
ids = append(ids, peer.ID)
}
assert.ElementsMatch(t, tt.wantIDs, ids, "filters should return only matching peers in the account")
query := url.Values{"name": {tt.nameFilter}, "ip": {tt.ipFilter}, "mac": {tt.macFilter}}
req := httptest.NewRequest(http.MethodGet, "/api/peers?"+query.Encode(), nil)
req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: account.Id, UserId: "mac-admin"})
recorder := httptest.NewRecorder()
handler.GetAllPeers(recorder, req)
require.Equal(t, http.StatusOK, recorder.Code, "peer listing should succeed: %s", recorder.Body.String())
var response []api.PeerBatch
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response))
responseIDs := make([]string, 0, len(response))
for _, peer := range response {
responseIDs = append(responseIDs, peer.Id)
}
assert.ElementsMatch(t, tt.wantIDs, responseIDs, "HTTP query filters should reach the store")
})
}
}
func setupTestAccountManager(b testing.TB, peers int, groups int) (*DefaultAccountManager, *update_channel.PeersUpdateManager, string, string, error) {
b.Helper()
@@ -934,7 +1006,7 @@ func BenchmarkGetPeers(b *testing.B) {
b.ResetTimer()
for i := 0; i < b.N; i++ {
_, err := manager.GetPeers(context.Background(), accountID, userID, "", "")
_, err := manager.GetPeers(context.Background(), accountID, userID, "", "", "")
if err != nil {
b.Fatalf("GetPeers failed: %v", err)
}
@@ -62,11 +62,11 @@ func TestAgentNetworkAdminRole(t *testing.T) {
}
}
// TestUsageViewerRole pins the least-privilege cost role: read on the
// aggregated usage overview plus read-only on the resources its filters
// and display columns resolve against (users, groups, peers, the provider
// list) — no policies, no request-level logs (which can contain captured
// prompts), nothing else in the account.
// TestUsageViewerRole pins the read-only usage role: read on the aggregated
// usage overview and the account-wide request-level logs, plus read-only on
// the resources their filters and display columns resolve against (users,
// groups, peers, the provider list) — no policies, guardrails, budgets, or
// settings, nothing else in the account.
func TestUsageViewerRole(t *testing.T) {
manager := NewManager(nil)
ctx := context.Background()
@@ -76,6 +76,7 @@ func TestUsageViewerRole(t *testing.T) {
readOnly := []modules.Module{
modules.AgentNetworkUsage,
modules.AgentNetworkLogs,
modules.AgentNetworkProviders,
modules.Users,
modules.Groups,
@@ -83,7 +84,7 @@ func TestUsageViewerRole(t *testing.T) {
}
for _, m := range readOnly {
assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, m, operations.Read),
"usage_viewer must read %s for the usage view and its filters", m)
"usage_viewer must read %s for the usage and log views and their filters", m)
for _, op := range []operations.Operation{operations.Create, operations.Update, operations.Delete} {
assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, m, op),
"usage_viewer must not have %s on %s", op, m)
@@ -95,7 +96,6 @@ func TestUsageViewerRole(t *testing.T) {
modules.AgentNetworkPolicies,
modules.AgentNetworkGuardrails,
modules.AgentNetworkBudgets,
modules.AgentNetworkLogs,
modules.AgentNetworkSettings,
modules.Networks,
modules.SetupKeys,
@@ -7,16 +7,15 @@ import (
)
// UsageViewer is the regular User baseline plus read access to the
// aggregated Agent Network usage and cost overview, and read-only access
// to the resources the usage filters and display columns resolve against:
// users and groups (identity filters and name resolution), peers (agent
// principals in the caller column), and the provider list (provider and
// model filter options — the manager redacts connection config such as
// upstream URLs and operator-supplied header values for callers holding
// read without update). It sees no policies and no account-wide
// request-level access logs (which can contain captured prompts); its own
// requests remain readable through the self-scoped endpoints, like any
// caller's.
// aggregated Agent Network usage and cost overview and to the account-wide
// request-level access logs (which can contain captured prompts), and
// read-only access to the resources the usage and log filters and display
// columns resolve against: users and groups (identity filters and name
// resolution), peers (agent principals in the caller column), and the
// provider list (provider and model filter options — the manager redacts
// connection config such as upstream URLs and operator-supplied header
// values for callers holding read without update). It sees no policies,
// guardrails, budgets, or Agent Network settings.
var UsageViewer = RolePermissions{
Role: types.UserRoleUsageViewer,
AutoAllowNew: map[operations.Operation]bool{
@@ -32,6 +31,12 @@ var UsageViewer = RolePermissions{
operations.Update: false,
operations.Delete: false,
},
modules.AgentNetworkLogs: {
operations.Read: true,
operations.Create: false,
operations.Update: false,
operations.Delete: false,
},
modules.AgentNetworkProviders: {
operations.Read: true,
operations.Create: false,

Some files were not shown because too many files have changed in this diff Show More