mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-06 05:29:07 +02:00
Merge main into poc/certificate-posture
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
+51
-1
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+42
-5
@@ -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
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,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)
|
||||
}
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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{}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user