Merge main into poc/certificate-posture

This commit is contained in:
Viktor Liu
2026-10-05 19:09:43 +02:00
390 changed files with 28443 additions and 14545 deletions
@@ -44,7 +44,7 @@ func (r *repository) GetAccountNetwork(ctx context.Context, accountID string) (*
}
func (r *repository) GetAccountPeers(ctx context.Context, accountID string) ([]*peer.Peer, error) {
return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
}
func (r *repository) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) {
@@ -245,7 +245,7 @@ func computeMode(t *testing.T, ctx context.Context, mode Mode, nmData *networkma
peerGroups := maps.Keys(nmData.GetPeerGroups(peerID))
resp := mgmtgrpc.ToComponentSyncResponse(ctx, nil, nil, nil, peer, nil, nil, components, nil,
dnsDomain, nil, nmData.AccountSettings, nil, peerGroups, dnsFwdPort)
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain)
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain, false)
require.NoError(t, err, "expand envelope")
return res.NetworkMap
default:
@@ -0,0 +1,76 @@
package agentnetwork
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/store"
nbtypes "github.com/netbirdio/netbird/management/server/types"
)
// TestCleanupAccessLogs_RealStore_DeletedAccount covers a deleted account's access logs.
// The sweep is driven by settings rows, which go with the account, so without a fallback
// those logs would never expire. They get the default retention instead. A live account
// can delete its own settings row, so "no settings" must not be mistaken for "deleted":
// that account's logs are left alone, as are those of an account that keeps logs forever.
func TestCleanupAccessLogs_RealStore_DeletedAccount(t *testing.T) {
ctx := context.Background()
s, cleanup, err := store.NewTestStoreFromSQL(ctx, "", t.TempDir())
require.NoError(t, err, "real sqlite test store must come up")
defer cleanup()
const (
deletedAccountID = "acc-deleted"
keepAccountID = "acc-keep-forever"
noSettingsAccountID = "acc-live-no-settings"
)
old := time.Now().UTC().AddDate(0, 0, -(types.DefaultAccessLogRetentionDays + 10))
recent := time.Now().UTC().AddDate(0, 0, -1)
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: keepAccountID}))
require.NoError(t, s.SaveAccount(ctx, &nbtypes.Account{Id: noSettingsAccountID}))
keepSettings := types.DefaultSettings(keepAccountID)
keepSettings.Domain = "keep.gw.example.com"
keepSettings.AccessLogRetentionDays = 0
require.NoError(t, s.SaveAgentNetworkSettings(ctx, keepSettings))
mkLog := func(id, accountID string, ts time.Time) {
t.Helper()
entry := &types.AgentNetworkAccessLog{
ID: id, AccountID: accountID, ServiceID: "svc", Timestamp: ts, StatusCode: 200, Model: "gpt-4o",
}
groups := []types.AgentNetworkAccessLogGroup{{LogID: id, GroupID: "grp-eng", AccountID: accountID}}
require.NoError(t, s.CreateAgentNetworkAccessLog(ctx, entry, groups))
}
mkLog("deleted-old", deletedAccountID, old)
mkLog("deleted-recent", deletedAccountID, recent)
mkLog("keep-old", keepAccountID, old)
mkLog("no-settings-old", noSettingsAccountID, old)
m := &managerImpl{store: s}
m.cleanupAccessLogsOnce(ctx)
logIDs := func(accountID string) []string {
t.Helper()
logs, _, err := s.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID,
types.AgentNetworkAccessLogFilter{Page: 1, PageSize: 50})
require.NoError(t, err)
ids := make([]string, 0, len(logs))
for _, l := range logs {
ids = append(ids, l.ID)
}
return ids
}
assert.Equal(t, []string{"deleted-recent"}, logIDs(deletedAccountID),
"a deleted account should have logs past the default retention swept")
assert.Equal(t, []string{"keep-old"}, logIDs(keepAccountID),
"an account with retention disabled should keep its old logs")
assert.Equal(t, []string{"no-settings-old"}, logIDs(noSettingsAccountID),
"a live account without a settings row should keep its old logs")
}
@@ -80,6 +80,9 @@ type Manager interface {
ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error)
GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error)
StartAccessLogCleanup(ctx context.Context, cleanupIntervalHours int)
// RemoveAccountGateway drops the account's gateway mappings from the
// proxies. It runs as an account deletion hook.
RemoveAccountGateway(ctx context.Context, accountID string) error
RecordConsumption(ctx context.Context, accountID string, kind types.ConsumptionDimension, dimID string, windowSeconds, tokensIn, tokensOut int64, costUSD float64) error
RecordAccountBudgetUsage(ctx context.Context, accountID, userID string, groupIDs []string, tokensIn, tokensOut int64, costUSD float64) error
RecordUsage(ctx context.Context, in RecordUsageInput) error
@@ -1350,8 +1353,8 @@ func (m *managerImpl) scopeFilterToCaller(ctx context.Context, accountID, userID
// StartAccessLogCleanup launches a background sweep that periodically deletes
// each account's agent-network access-log rows older than that account's
// AccessLogRetentionDays. Usage records are never swept. A non-positive
// interval defaults to 24h.
// AccessLogRetentionDays, and the consumption counters of deleted accounts.
// Usage records are never swept. A non-positive interval defaults to 24h.
func (m *managerImpl) StartAccessLogCleanup(ctx context.Context, cleanupIntervalHours int) {
if cleanupIntervalHours <= 0 {
cleanupIntervalHours = 24
@@ -1362,21 +1365,40 @@ func (m *managerImpl) StartAccessLogCleanup(ctx context.Context, cleanupInterval
ticker := time.NewTicker(interval)
defer ticker.Stop()
m.cleanupAccessLogsOnce(ctx) // run once on startup
m.cleanupOnce(ctx) // run once on startup
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
m.cleanupAccessLogsOnce(ctx)
m.cleanupOnce(ctx)
}
}
}()
}
func (m *managerImpl) cleanupOnce(ctx context.Context) {
m.cleanupAccessLogsOnce(ctx)
m.cleanupDeletedAccountConsumption(ctx)
}
// cleanupDeletedAccountConsumption deletes the consumption counters of accounts
// that no longer exist. Best-effort: a failure is logged and retried next sweep.
func (m *managerImpl) cleanupDeletedAccountConsumption(ctx context.Context) {
deleted, err := m.store.DeleteAgentNetworkConsumptionOfDeletedAccounts(ctx)
if err != nil {
log.WithContext(ctx).Warnf("agent-network consumption cleanup: %v", err)
return
}
if deleted > 0 {
log.WithContext(ctx).Infof("agent-network consumption cleanup: deleted %d counters of deleted accounts", deleted)
}
}
// cleanupAccessLogsOnce sweeps every account's expired access-log rows against
// its configured retention. Best-effort: a per-account failure is logged and
// the sweep continues.
// its configured retention. Deleted accounts, whose settings rows went with
// them, get the default retention. Best-effort: a per-account failure is
// logged and the sweep continues.
func (m *managerImpl) cleanupAccessLogsOnce(ctx context.Context) {
settings, err := m.store.GetAllAgentNetworkSettings(ctx, store.LockingStrengthNone)
if err != nil {
@@ -1384,18 +1406,31 @@ func (m *managerImpl) cleanupAccessLogsOnce(ctx context.Context) {
return
}
for _, s := range settings {
if s.AccessLogRetentionDays <= 0 {
continue // keep indefinitely
}
cutoff := time.Now().UTC().AddDate(0, 0, -s.AccessLogRetentionDays)
deleted, err := m.store.DeleteOldAgentNetworkAccessLogs(ctx, s.AccountID, cutoff)
if err != nil {
log.WithContext(ctx).Warnf("agent-network access-log cleanup for account %s: %v", s.AccountID, err)
continue
}
if deleted > 0 {
log.WithContext(ctx).Infof("agent-network access-log cleanup: deleted %d rows for account %s (retention %d days)", deleted, s.AccountID, s.AccessLogRetentionDays)
}
m.cleanupAccountAccessLogs(ctx, s.AccountID, s.AccessLogRetentionDays)
}
deleted, err := m.store.GetDeletedAccountIDsWithAgentNetworkAccessLogs(ctx)
if err != nil {
log.WithContext(ctx).Errorf("agent-network access-log cleanup: list deleted accounts: %v", err)
return
}
for _, accountID := range deleted {
m.cleanupAccountAccessLogs(ctx, accountID, types.DefaultAccessLogRetentionDays)
}
}
func (m *managerImpl) cleanupAccountAccessLogs(ctx context.Context, accountID string, retentionDays int) {
if retentionDays <= 0 {
return // keep indefinitely
}
cutoff := time.Now().UTC().AddDate(0, 0, -retentionDays)
deleted, err := m.store.DeleteOldAgentNetworkAccessLogs(ctx, accountID, cutoff)
if err != nil {
log.WithContext(ctx).Warnf("agent-network access-log cleanup for account %s: %v", accountID, err)
return
}
if deleted > 0 {
log.WithContext(ctx).Infof("agent-network access-log cleanup: deleted %d rows for account %s (retention %d days)", deleted, accountID, retentionDays)
}
}
@@ -1545,6 +1580,8 @@ func (*mockManager) GetUsageOverview(_ context.Context, _, _ string, _ types.Age
func (*mockManager) StartAccessLogCleanup(_ context.Context, _ int) {}
func (*mockManager) RemoveAccountGateway(_ context.Context, _ string) error { return nil }
func (*mockManager) RecordConsumption(_ context.Context, _ string, _ types.ConsumptionDimension, _ string, _, _, _ int64, _ float64) error {
return nil
}
@@ -2,8 +2,10 @@ package agentnetwork
import (
"context"
"fmt"
log "github.com/sirupsen/logrus"
goproto "google.golang.org/protobuf/proto"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/types"
@@ -81,18 +83,66 @@ func (m *managerImpl) reconcile(ctx context.Context, accountID string) {
}
m.reconcileMu.Unlock()
for _, entry := range creates {
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
m.sendMappings(ctx, accountID, creates, proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED)
m.sendMappings(ctx, accountID, updates, proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED)
m.sendMappings(ctx, accountID, deletes, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED)
}
// sendMappings sends each entry as updateType. It sends a copy: the entries'
// mappings are shared with reconcileCache, which another reconcile or
// RemoveAccountGateway may be reading, so they are never written.
func (m *managerImpl) sendMappings(ctx context.Context, accountID string, entries []syntheticMapping, updateType proto.ProxyMappingUpdateType) {
for _, entry := range entries {
update := goproto.Clone(entry.mapping).(*proto.ProxyMapping)
update.Type = updateType
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, update, entry.cluster)
}
for _, entry := range updates {
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
}
// RemoveAccountGateway tells the proxies to drop every mapping of the account's
// gateway, so a deleted account's proxy config, provider API keys included, does
// not linger in proxy memory until the next resync. It is an account deletion
// hook: it runs before the account's data is removed, the last point at which
// the mappings can be synthesised from the store. The cache alone would miss
// them, since it is per instance and empty after a restart. If the deletion
// then fails, the gateway stays down until the account's next change reconciles
// it back.
func (m *managerImpl) RemoveAccountGateway(ctx context.Context, accountID string) error {
if m.proxyController == nil {
return nil
}
for _, entry := range deletes {
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
services, err := SynthesizeServices(ctx, m.store, accountID)
if err != nil {
return fmt.Errorf("synthesise agent network services: %w", err)
}
oidcCfg := m.proxyController.GetOIDCValidationConfig()
removed := make(map[string]syntheticMapping, len(services))
for _, svc := range services {
if svc == nil || svc.ID == "" {
continue
}
removed[svc.ID] = syntheticMapping{
mapping: svc.ToProtoMapping(rpservice.Delete, "", oidcCfg),
cluster: svc.ProxyCluster,
}
}
m.reconcileMu.Lock()
for id, entry := range m.reconcileCache[accountID] {
if _, ok := removed[id]; !ok {
removed[id] = entry
}
}
delete(m.reconcileCache, accountID)
m.reconcileMu.Unlock()
entries := make([]syntheticMapping, 0, len(removed))
for _, entry := range removed {
entries = append(entries, entry)
}
m.sendMappings(ctx, accountID, entries, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED)
return nil
}
// diffMappings classifies the previous→current transition for a single
@@ -2,6 +2,8 @@ package agentnetwork
import (
"context"
"sync"
"sync/atomic"
"testing"
"go.uber.org/mock/gomock"
@@ -12,6 +14,7 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/management/status"
)
func newReconcileMgr(t *testing.T, ctrl *gomock.Controller) (*managerImpl, *store.MockStore, *proxy.MockController) {
@@ -287,3 +290,154 @@ func TestDiffMappings_RemovedServiceIsDeletedOnItsOwnCluster(t *testing.T) {
assert.Equal(t, "brave-otter.gateway.example.com", deletes[0].cluster)
}
}
// TestRemoveAccountGateway_EmitsRemovedFromStore — account deletion runs on an
// instance that may never have reconciled the account, so its cache is empty.
// The mappings are synthesised from the store, still intact before the delete,
// and each is sent as REMOVED to the cluster that serves it.
func TestRemoveAccountGateway_EmitsRemovedFromStore(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
provider := newReconcileTestProvider()
policy := newReconcileTestPolicy(provider.ID, "grp-eng")
expectReconcileSynthInputs(mockStore, ctx, []*types.Provider{provider}, []*types.Policy{policy}, []*types.Guardrail{})
mockProxy.EXPECT().GetOIDCValidationConfig().Return(proxy.OIDCValidationConfig{})
var sent []*proto.ProxyMapping
mockProxy.EXPECT().
SendServiceUpdateToCluster(ctx, "acct-1", gomock.Any(), "eu.proxy.netbird.io").
Do(func(_ context.Context, _ string, m *proto.ProxyMapping, _ string) {
sent = append(sent, m)
})
require.NoError(t, mgr.RemoveAccountGateway(ctx, "acct-1"))
require.Len(t, sent, 1, "the account's one gateway mapping must be removed")
assert.Equal(t, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED, sent[0].Type, "the update must be a removal")
assert.Equal(t, "agent-net-svc-acct-1", sent[0].Id, "the removal must name the account's gateway service")
}
// TestRemoveAccountGateway_AlsoRemovesCachedMappings — a mapping this instance
// last sent but the store no longer synthesises (here, one on another cluster)
// is removed too, and the account's cache entry is cleared.
func TestRemoveAccountGateway_AlsoRemovesCachedMappings(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
mgr.reconcileCache["acct-1"] = map[string]syntheticMapping{
"stale-svc": {mapping: &proto.ProxyMapping{Id: "stale-svc"}, cluster: "us.proxy.netbird.io"},
}
// Settings but no providers: the store synthesises nothing.
mockStore.EXPECT().
GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").
Return(newReconcileTestSettings(), nil)
mockStore.EXPECT().
GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "acct-1").
Return([]*types.Provider{}, nil)
mockProxy.EXPECT().GetOIDCValidationConfig().Return(proxy.OIDCValidationConfig{})
var sent []*proto.ProxyMapping
mockProxy.EXPECT().
SendServiceUpdateToCluster(ctx, "acct-1", gomock.Any(), "us.proxy.netbird.io").
Do(func(_ context.Context, _ string, m *proto.ProxyMapping, _ string) {
sent = append(sent, m)
})
require.NoError(t, mgr.RemoveAccountGateway(ctx, "acct-1"))
require.Len(t, sent, 1, "the cached mapping must be removed from its own cluster")
assert.Equal(t, "stale-svc", sent[0].Id)
assert.Equal(t, proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED, sent[0].Type)
mgr.reconcileMu.Lock()
_, present := mgr.reconcileCache["acct-1"]
mgr.reconcileMu.Unlock()
assert.False(t, present, "the deleted account's cache entry must be cleared")
}
// TestRemoveAccountGateway_SynthFailureAbortsDeletion — if the mappings cannot
// be read, nothing is sent and the error is returned, which as an account
// deletion hook keeps the account rather than leaving its gateway running.
func TestRemoveAccountGateway_SynthFailureAbortsDeletion(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mgr, mockStore, _ := newReconcileMgr(t, ctrl)
mockStore.EXPECT().
GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").
Return(nil, status.Errorf(status.Internal, "store unavailable"))
assert.Error(t, mgr.RemoveAccountGateway(ctx, "acct-1"), "a failed synthesis must fail the hook")
}
func TestRemoveAccountGateway_NilProxyController_NoOp(t *testing.T) {
mgr := &managerImpl{reconcileCache: make(map[string]map[string]syntheticMapping)}
// Must not panic and must not query the store.
assert.NoError(t, mgr.RemoveAccountGateway(context.Background(), "acct-1"))
}
// TestReconcile_ConcurrentWithGatewayChanges — while an account's gateway
// flaps (its policy is removed and re-added between reads), concurrent
// reconciles and RemoveAccountGateway share the cached mappings: one caches a
// mapping and sends it, another finds it gone and sends its removal. Run under
// -race: neither path may write a cached mapping, only copies of it.
func TestReconcile_ConcurrentWithGatewayChanges(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mgr, mockStore, mockProxy := newReconcileMgr(t, ctrl)
// gomock serialises every call on the controller's mutex, which would give
// the race detector the ordering the code under test lacks. The sends go
// through a fake that takes no lock.
mgr.proxyController = unsyncedSender{MockController: mockProxy}
provider := newReconcileTestProvider()
policy := newReconcileTestPolicy(provider.ID, "grp-eng")
var reads atomic.Int64
mockStore.EXPECT().GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "acct-1").Return(newReconcileTestSettings(), nil).AnyTimes()
mockStore.EXPECT().GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "acct-1").Return([]*types.Provider{provider}, nil).AnyTimes()
mockStore.EXPECT().GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, "acct-1").
DoAndReturn(func(context.Context, store.LockingStrength, string) ([]*types.Policy, error) {
if reads.Add(1)%2 == 0 {
return []*types.Policy{}, nil
}
return []*types.Policy{policy}, nil
}).AnyTimes()
mockStore.EXPECT().GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, "acct-1").Return([]*types.Guardrail{}, nil).AnyTimes()
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(1)
go func(remove bool) {
defer wg.Done()
for j := 0; j < 50; j++ {
if remove && j%10 == 0 {
_ = mgr.RemoveAccountGateway(ctx, "acct-1")
continue
}
mgr.reconcile(ctx, "acct-1")
}
}(i == 0)
}
wg.Wait()
}
// unsyncedSender answers the calls reconcile makes on every pass without any
// locking, so concurrent callers are not ordered by the fake itself.
type unsyncedSender struct {
*proxy.MockController
}
func (unsyncedSender) GetOIDCValidationConfig() proxy.OIDCValidationConfig {
return proxy.OIDCValidationConfig{}
}
func (unsyncedSender) SendServiceUpdateToCluster(context.Context, string, *proto.ProxyMapping, string) {}
@@ -97,7 +97,7 @@ func (m *managerImpl) GetAllPeers(ctx context.Context, accountID, userID string)
return m.store.GetUserPeers(ctx, store.LockingStrengthNone, accountID, userID)
}
return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
}
func (m *managerImpl) GetPeerAccountID(ctx context.Context, peerID string) (string, error) {
@@ -9,6 +9,7 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/management/server/geolocation"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
@@ -18,14 +19,16 @@ import (
)
type managerImpl struct {
repo accesslogs.Repository
store store.Store
permissionsManager permissions.Manager
geo geolocation.Geolocation
cleanupCancel context.CancelFunc
}
func NewManager(store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager {
func NewManager(repo accesslogs.Repository, store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager {
return &managerImpl{
repo: repo,
store: store,
permissionsManager: permissionsManager,
geo: geo,
@@ -54,7 +57,7 @@ func (m *managerImpl) SaveAccessLog(ctx context.Context, logEntry *accesslogs.Ac
}
}
if err := m.store.CreateAccessLog(ctx, logEntry); err != nil {
if err := m.repo.Create(ctx, logEntry); err != nil {
log.WithContext(ctx).WithFields(log.Fields{
"service_id": logEntry.ServiceID,
"method": logEntry.Method,
@@ -82,7 +85,7 @@ func (m *managerImpl) GetAllAccessLogs(ctx context.Context, accountID, userID st
log.WithContext(ctx).Warnf("failed to resolve user filters: %v", err)
}
logs, totalCount, err := m.store.GetAccountAccessLogs(ctx, store.LockingStrengthNone, accountID, *filter)
logs, totalCount, err := m.repo.ListByAccount(ctx, db.LockingStrengthNone, accountID, *filter)
if err != nil {
return nil, 0, err
}
@@ -98,7 +101,7 @@ func (m *managerImpl) CleanupOldAccessLogs(ctx context.Context, retentionDays in
}
cutoffTime := time.Now().AddDate(0, 0, -retentionDays)
deletedCount, err := m.store.DeleteOldAccessLogs(ctx, cutoffTime)
deletedCount, err := m.repo.DeleteOlderThan(ctx, cutoffTime)
if err != nil {
log.WithContext(ctx).Errorf("failed to cleanup old access logs: %v", err)
return 0, err
@@ -5,27 +5,27 @@ import (
"testing"
"time"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
)
func TestCleanupOldAccessLogs(t *testing.T) {
tests := []struct {
name string
retentionDays int
setupMock func(*store.MockStore)
setupMock func(*accesslogs.MockRepository)
expectedCount int64
expectedError bool
}{
{
name: "cleanup logs older than retention period",
retentionDays: 30,
setupMock: func(mockStore *store.MockStore) {
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
setupMock: func(mockRepo *accesslogs.MockRepository) {
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) {
expectedCutoff := time.Now().AddDate(0, 0, -30)
timeDiff := olderThan.Sub(expectedCutoff)
@@ -41,9 +41,9 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "no logs to cleanup",
retentionDays: 30,
setupMock: func(mockStore *store.MockStore) {
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
setupMock: func(mockRepo *accesslogs.MockRepository) {
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(0), nil)
},
expectedCount: 0,
@@ -52,8 +52,8 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "zero retention days skips cleanup",
retentionDays: 0,
setupMock: func(mockStore *store.MockStore) {
// No expectations - DeleteOldAccessLogs should not be called
setupMock: func(mockRepo *accesslogs.MockRepository) {
// No expectations - DeleteOlderThan should not be called
},
expectedCount: 0,
expectedError: false,
@@ -61,8 +61,8 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "negative retention days skips cleanup",
retentionDays: -10,
setupMock: func(mockStore *store.MockStore) {
// No expectations - DeleteOldAccessLogs should not be called
setupMock: func(mockRepo *accesslogs.MockRepository) {
// No expectations - DeleteOlderThan should not be called
},
expectedCount: 0,
expectedError: false,
@@ -74,11 +74,11 @@ func TestCleanupOldAccessLogs(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
tt.setupMock(mockStore)
mockRepo := accesslogs.NewMockRepository(ctrl)
tt.setupMock(mockRepo)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -98,10 +98,10 @@ func TestCleanupWithExactBoundary(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) {
expectedCutoff := time.Now().AddDate(0, 0, -30)
timeDiff := olderThan.Sub(expectedCutoff)
@@ -110,7 +110,7 @@ func TestCleanupWithExactBoundary(t *testing.T) {
})
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -125,11 +125,11 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
// No expectations - cleanup should not run
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -139,22 +139,22 @@ func TestStartPeriodicCleanup(t *testing.T) {
time.Sleep(100 * time.Millisecond)
// If DeleteOldAccessLogs was called, the test will fail due to unexpected call
// If DeleteOlderThan was called, the test will fail due to unexpected call
})
t.Run("periodic cleanup runs immediately on start", func(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(2), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -171,15 +171,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(1), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -198,15 +198,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(0), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -223,15 +223,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(3), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -249,15 +249,15 @@ func TestStopPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(1), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -0,0 +1,135 @@
package manager
import (
"context"
"strings"
"time"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/shared/management/status"
)
type sqlRepository struct {
conn *db.Conn
db *gorm.DB
}
// NewRepository returns the access log repository backed by conn.
func NewRepository(conn *db.Conn) accesslogs.Repository {
return &sqlRepository{conn: conn, db: conn.DB(nil)}
}
func (r *sqlRepository) WithTx(tx *db.Tx) accesslogs.Repository {
return &sqlRepository{conn: r.conn, db: r.conn.DB(tx)}
}
func (r *sqlRepository) Create(ctx context.Context, entry *accesslogs.AccessLogEntry) error {
if err := r.db.Create(entry).Error; err != nil {
log.WithContext(ctx).WithFields(log.Fields{
"service_id": entry.ServiceID,
"method": entry.Method,
"host": entry.Host,
"path": entry.Path,
}).Errorf("failed to create access log entry in store: %v", err)
return status.Errorf(status.Internal, "failed to create access log entry in store")
}
return nil
}
// ListByAccount returns one page of an account's access logs together with the
// total number of entries matching the filter.
func (r *sqlRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) {
var totalCount int64
countQuery := applyFilters(r.db.Model(&accesslogs.AccessLogEntry{}).Where("account_id = ?", accountID), filter)
if err := countQuery.Count(&totalCount).Error; err != nil {
log.WithContext(ctx).Errorf("failed to count access logs: %v", err)
return nil, 0, status.Errorf(status.Internal, "failed to count access logs")
}
query := applyFilters(r.db.Where("account_id = ?", accountID), filter)
sortOrder := strings.ToUpper(filter.GetSortOrder())
for _, column := range strings.Split(filter.GetSortColumn(), ",") {
if column = strings.TrimSpace(column); column != "" {
query = query.Order(column + " " + sortOrder)
}
}
query = query.Limit(filter.GetLimit()).Offset(filter.GetOffset())
if lockStrength != db.LockingStrengthNone {
query = query.Clauses(clause.Locking{Strength: string(lockStrength)})
}
var logs []*accesslogs.AccessLogEntry
if err := query.Find(&logs).Error; err != nil {
log.WithContext(ctx).Errorf("failed to get access logs from store: %v", err)
return nil, 0, status.Errorf(status.Internal, "failed to get access logs from store")
}
return logs, totalCount, nil
}
func (r *sqlRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) {
result := r.db.Where("timestamp < ?", olderThan).Delete(&accesslogs.AccessLogEntry{})
if result.Error != nil {
log.WithContext(ctx).Errorf("failed to delete old access logs: %v", result.Error)
return 0, status.Errorf(status.Internal, "failed to delete old access logs")
}
return result.RowsAffected, nil
}
func applyFilters(query *gorm.DB, filter accesslogs.AccessLogFilter) *gorm.DB {
if filter.Search != nil {
searchPattern := "%" + *filter.Search + "%"
query = query.Where(
"id LIKE ? OR location_connection_ip LIKE ? OR host LIKE ? OR path LIKE ? OR CONCAT(host, path) LIKE ? OR user_id IN (SELECT id FROM users WHERE email LIKE ? OR name LIKE ?)",
searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern,
)
}
if filter.SourceIP != nil {
query = query.Where("location_connection_ip = ?", *filter.SourceIP)
}
if filter.Host != nil {
query = query.Where("host = ?", *filter.Host)
}
if filter.Path != nil {
query = query.Where("path LIKE ?", "%"+*filter.Path+"%")
}
if filter.UserID != nil {
query = query.Where("user_id = ?", *filter.UserID)
}
if filter.Method != nil {
query = query.Where("method = ?", *filter.Method)
}
if filter.Status != nil {
switch *filter.Status {
case "success":
query = query.Where("(status_code >= ? AND status_code < ?)", 200, 400)
case "failed":
query = query.Where("((status_code >= ? AND status_code < ?) OR status_code >= ?)", 100, 200, 400)
}
}
if filter.StatusCode != nil {
query = query.Where("status_code = ?", *filter.StatusCode)
}
if filter.StartDate != nil {
query = query.Where("timestamp >= ?", *filter.StartDate)
}
if filter.EndDate != nil {
query = query.Where("timestamp <= ?", *filter.EndDate)
}
return query
}
@@ -0,0 +1,125 @@
package manager
import (
"context"
"errors"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/management/internals/shared/db/dbtest"
)
func newTestRepository(t *testing.T) (accesslogs.Repository, *db.Conn) {
conn := dbtest.NewConn(t, &accesslogs.AccessLogEntry{})
return NewRepository(conn), conn
}
func newEntry(id, accountID, method string, age time.Duration) *accesslogs.AccessLogEntry {
return &accesslogs.AccessLogEntry{
ID: id,
AccountID: accountID,
Method: method,
Host: "app.example.com",
Path: "/",
StatusCode: 200,
Timestamp: time.Now().Add(-age),
}
}
func TestSqlRepository_ListByAccount(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
for _, entry := range []*accesslogs.AccessLogEntry{
newEntry("a1", "acc-a", "GET", 3*time.Hour),
newEntry("a2", "acc-a", "POST", 2*time.Hour),
newEntry("a3", "acc-a", "GET", time.Hour),
newEntry("b1", "acc-b", "GET", time.Hour),
} {
require.NoError(t, repo.Create(ctx, entry))
}
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 2})
require.NoError(t, err)
assert.EqualValues(t, 3, total)
require.Len(t, logs, 2)
assert.Equal(t, "a3", logs[0].ID)
assert.Equal(t, "a2", logs[1].ID)
method := "GET"
logs, total, err = repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Method: &method, SortOrder: "asc"})
require.NoError(t, err)
assert.EqualValues(t, 2, total)
require.Len(t, logs, 2)
assert.Equal(t, "a1", logs[0].ID)
assert.Equal(t, "a3", logs[1].ID)
}
func TestSqlRepository_DeleteOlderThan(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
require.NoError(t, repo.Create(ctx, newEntry("old", "acc", "GET", 48*time.Hour)))
require.NoError(t, repo.Create(ctx, newEntry("new", "acc", "GET", time.Hour)))
deleted, err := repo.DeleteOlderThan(ctx, time.Now().Add(-24*time.Hour))
require.NoError(t, err)
assert.EqualValues(t, 1, deleted)
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.EqualValues(t, 1, total)
require.Len(t, logs, 1)
assert.Equal(t, "new", logs[0].ID)
}
func TestSqlRepository_CreateInsideTransactionRollsBack(t *testing.T) {
repo, conn := newTestRepository(t)
ctx := context.Background()
failure := errors.New("abort")
err := conn.RunInTx(ctx, func(tx *db.Tx) error {
txRepo := repo.WithTx(tx)
require.NoError(t, txRepo.Create(ctx, newEntry("tx", "acc", "GET", 0)))
_, total, err := txRepo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.EqualValues(t, 1, total)
return failure
})
require.ErrorIs(t, err, failure)
_, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Zero(t, total)
}
func TestSqlRepository_ListByAccount_StatusFilter(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
statusCodes := map[string]int{"l4": 0, "info": 101, "ok": 200, "notfound": 404}
for id, code := range statusCodes {
entry := newEntry(id, "acc", "GET", time.Hour)
entry.StatusCode = code
require.NoError(t, repo.Create(ctx, entry))
}
foreign := newEntry("foreign", "other", "GET", time.Hour)
foreign.StatusCode = 500
require.NoError(t, repo.Create(ctx, foreign))
listIDs := func(status string) []string {
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Status: &status, SortBy: "status_code", SortOrder: "asc"})
require.NoError(t, err)
require.EqualValues(t, len(logs), total)
ids := make([]string, 0, len(logs))
for _, entry := range logs {
ids = append(ids, entry.ID)
}
return ids
}
assert.Equal(t, []string{"info", "notfound"}, listIDs("failed"))
assert.Equal(t, []string{"ok"}, listIDs("success"))
}
@@ -0,0 +1,18 @@
package accesslogs
import (
"context"
"time"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
//go:generate go tool mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod
// Repository persists reverse proxy access log entries.
type Repository interface {
WithTx(tx *db.Tx) Repository
Create(ctx context.Context, entry *AccessLogEntry) error
ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error)
DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error)
}
@@ -0,0 +1,102 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./repository.go
//
// Generated by this command:
//
// mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod
//
// Package accesslogs is a generated GoMock package.
package accesslogs
import (
context "context"
reflect "reflect"
time "time"
db "github.com/netbirdio/netbird/management/internals/shared/db"
gomock "go.uber.org/mock/gomock"
)
// MockRepository is a mock of Repository interface.
type MockRepository struct {
ctrl *gomock.Controller
recorder *MockRepositoryMockRecorder
isgomock struct{}
}
// MockRepositoryMockRecorder is the mock recorder for MockRepository.
type MockRepositoryMockRecorder struct {
mock *MockRepository
}
// NewMockRepository creates a new mock instance.
func NewMockRepository(ctrl *gomock.Controller) *MockRepository {
mock := &MockRepository{ctrl: ctrl}
mock.recorder = &MockRepositoryMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockRepository) EXPECT() *MockRepositoryMockRecorder {
return m.recorder
}
// Create mocks base method.
func (m *MockRepository) Create(ctx context.Context, entry *AccessLogEntry) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Create", ctx, entry)
ret0, _ := ret[0].(error)
return ret0
}
// Create indicates an expected call of Create.
func (mr *MockRepositoryMockRecorder) Create(ctx, entry any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Create", reflect.TypeOf((*MockRepository)(nil).Create), ctx, entry)
}
// DeleteOlderThan mocks base method.
func (m *MockRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "DeleteOlderThan", ctx, olderThan)
ret0, _ := ret[0].(int64)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// DeleteOlderThan indicates an expected call of DeleteOlderThan.
func (mr *MockRepositoryMockRecorder) DeleteOlderThan(ctx, olderThan any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteOlderThan", reflect.TypeOf((*MockRepository)(nil).DeleteOlderThan), ctx, olderThan)
}
// ListByAccount mocks base method.
func (m *MockRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ListByAccount", ctx, lockStrength, accountID, filter)
ret0, _ := ret[0].([]*AccessLogEntry)
ret1, _ := ret[1].(int64)
ret2, _ := ret[2].(error)
return ret0, ret1, ret2
}
// ListByAccount indicates an expected call of ListByAccount.
func (mr *MockRepositoryMockRecorder) ListByAccount(ctx, lockStrength, accountID, filter any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListByAccount", reflect.TypeOf((*MockRepository)(nil).ListByAccount), ctx, lockStrength, accountID, filter)
}
// WithTx mocks base method.
func (m *MockRepository) WithTx(tx *db.Tx) Repository {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "WithTx", tx)
ret0, _ := ret[0].(Repository)
return ret0
}
// WithTx indicates an expected call of WithTx.
func (mr *MockRepositoryMockRecorder) WithTx(tx any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WithTx", reflect.TypeOf((*MockRepository)(nil).WithTx), tx)
}
@@ -0,0 +1,80 @@
package manager
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/gorilla/mux"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/activity"
nbcontext "github.com/netbirdio/netbird/management/server/context"
nbstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/auth"
)
func TestDeleteDomain_ServiceDependencies(t *testing.T) {
for _, tt := range []struct {
name string
domainName string
serviceHost string
accountID string
enabled bool
protected bool
}{
{"exact", "example.com", "example.com", accountA, true, true},
{"subdomain", "example.com", "deep.app.example.com", accountA, true, true},
{"disabled", "example.com", "app.example.com", accountA, false, true},
// A service is authorized by its own account's registration, so another
// account's service under this namespace is not a dependency of it.
{"other account", "example.com", "app.example.com", accountB, true, false},
{"case and trailing dot", "example.com", "APP.EXAMPLE.COM.", accountA, true, true},
{"suffix boundary", "example.com", "notexample.com", accountA, true, false},
{"literal underscore", "a_b.example.com", "app.a_b.example.com", accountA, true, true},
{"underscore wildcard", "a_b.example.com", "app.axb.example.com", accountA, true, false},
} {
t.Run(tt.name, func(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
events := captureDomainEvents(env)
d, err := env.store.CreateCustomDomain(ctx, accountA, tt.domainName, testCluster, true)
require.NoError(t, err)
svc := &rpservice.Service{
ID: "dependent", AccountID: tt.accountID, Domain: tt.serviceHost,
Enabled: tt.enabled, ProxyCluster: testCluster,
}
require.NoError(t, env.store.CreateService(ctx, svc))
router := mux.NewRouter()
RegisterEndpoints(router, env.manager)
deleteDomain := func() *httptest.ResponseRecorder {
req := httptest.NewRequest(http.MethodDelete, "/domains/"+d.ID, nil)
req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: accountA, UserId: accountAUser})
response := httptest.NewRecorder()
router.ServeHTTP(response, req)
return response
}
response := deleteDomain()
if tt.protected {
require.Equal(t, http.StatusPreconditionFailed, response.Code, "dependent services must block deletion: %s", response.Body.String())
assert.NotContains(t, response.Body.String(), tt.accountID, "the error must not reveal the service's account")
assert.NotNil(t, storedDomain(t, env.store, accountA, d.Domain), "the namespace must remain reserved")
assert.Empty(t, events.get(), "rejected deletion must not emit DomainDeleted")
stored, err := env.store.GetServiceByID(ctx, nbstore.LockingStrengthNone, tt.accountID, svc.ID)
require.NoError(t, err)
assert.Equal(t, svc.Enabled, stored.Enabled, "rejected deletion must preserve the service")
require.NoError(t, env.store.DeleteService(ctx, tt.accountID, svc.ID))
response = deleteDomain()
}
require.Equal(t, http.StatusNoContent, response.Code, "deletion must succeed without dependencies: %s", response.Body.String())
assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "the registration must be deleted")
captured := events.get()
require.Len(t, captured, 1, "only successful deletion may emit an event")
assert.Equal(t, activity.DomainDeleted, captured[0].Activity, "the event must describe the successful deletion")
})
}
}
@@ -357,6 +357,26 @@ func (m Manager) DeriveClusterFromDomain(ctx context.Context, accountID, domain
return "", fmt.Errorf("domain %s does not match any available proxy cluster", domain)
}
// ValidateServiceDomain holds custom domain authorization through a service write transaction.
func (m Manager) ValidateServiceDomain(ctx context.Context, tx nbstore.Store, accountID, serviceDomain, cluster string) error {
if _, ok := ExtractClusterFromFreeDomain(serviceDomain, []string{cluster}); ok {
return nil
}
name, err := nbdomain.FromString(serviceDomain)
if err != nil {
return status.Errorf(status.InvalidArgument, "invalid service domain: %v", err)
}
customDomains, err := tx.LockCustomDomains(ctx, accountID, name)
if err != nil {
return err
}
target, match := extractClusterFromCustomDomains(serviceDomain, customDomains)
if match != customDomainValidated || target != cluster {
return status.Errorf(status.PreconditionFailed, "custom domain authorization changed; retry the service operation")
}
return nil
}
func (m Manager) getClusterAllowList(ctx context.Context, accountID string) ([]string, error) {
byopAddresses, err := m.proxyManager.GetActiveClusterAddressesForAccount(ctx, accountID)
if err != nil {
@@ -99,7 +99,7 @@ func setupDomainTest(t *testing.T) *domainTestEnv {
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", nil, nil)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", "", nil, nil)
require.NoError(t, err)
resolver := &stubResolver{cnames: make(map[string]string)}
@@ -11,7 +11,7 @@ import (
// Manager defines the interface for proxy operations
type Manager interface {
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error)
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error)
Disconnect(ctx context.Context, proxyID, sessionID string) error
Heartbeat(ctx context.Context, p *Proxy) error
GetActiveClusterAddresses(ctx context.Context) ([]string, error)
@@ -20,6 +20,8 @@ type Manager interface {
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool
CleanupStale(ctx context.Context, inactivityDuration time.Duration) error
GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error)
CountAccountProxies(ctx context.Context, accountID string) (int64, error)
@@ -8,6 +8,7 @@ import (
"go.opentelemetry.io/otel/metric"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
nbversion "github.com/netbirdio/netbird/version"
)
// store defines the interface for proxy persistence operations
@@ -22,6 +23,8 @@ type store interface {
GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error)
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error)
CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error)
@@ -29,6 +32,8 @@ type store interface {
DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error
}
const minSessionCodeVersion = "0.81.0"
// Manager handles all proxy operations
type Manager struct {
store store
@@ -50,7 +55,7 @@ func NewManager(store store, meter metric.Meter) (*Manager, error) {
// Connect registers a new proxy connection in the database.
// capabilities may be nil for old proxies that do not report them.
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
now := time.Now()
var caps proxy.Capabilities
if capabilities != nil {
@@ -61,6 +66,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
SessionID: sessionID,
ClusterAddress: clusterAddress,
IPAddress: ipAddress,
Version: truncateVersion(version),
AccountID: accountID,
LastSeen: now,
ConnectedAt: &now,
@@ -78,6 +84,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
"sessionID": sessionID,
"clusterAddress": clusterAddress,
"ipAddress": ipAddress,
"version": p.Version,
}).Info("proxy connected")
return p, nil
@@ -143,6 +150,27 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string)
return m.store.GetClusterSupportsPrivate(ctx, clusterAddr)
}
// ClusterAllProxiesPrivate reports whether every active proxy claims the private capability (nil = unreported).
func (m Manager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
return m.store.GetClusterAllProxiesPrivate(ctx, clusterAddr)
}
// ClusterSupportsSessionCode reports whether all active proxies support session codes.
func (m Manager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
versions, err := m.store.GetActiveProxyVersions(ctx, clusterAddr)
if err != nil || len(versions) == 0 {
return false
}
for _, version := range versions {
if supported, err := nbversion.MeetsMinVersion(minSessionCodeVersion, version); err != nil || !supported {
return false
}
}
return true
}
// CleanupStale removes proxies that haven't sent heartbeat in the specified duration
func (m *Manager) CleanupStale(ctx context.Context, inactivityDuration time.Duration) error {
if err := m.store.CleanupStaleProxies(ctx, inactivityDuration); err != nil {
@@ -184,3 +212,13 @@ func (m *Manager) DeleteAccountCluster(ctx context.Context, clusterAddress, acco
}
return nil
}
// truncateVersion cuts a proxy-reported version to the column width so an
// oversized value cannot fail the save and block the connect.
func truncateVersion(version string) string {
runes := []rune(version)
if len(runes) <= proxy.MaxVersionLength {
return version
}
return string(runes[:proxy.MaxVersionLength])
}
@@ -4,8 +4,10 @@ import (
"context"
"errors"
"fmt"
"strings"
"testing"
"time"
"unicode/utf8"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -20,6 +22,7 @@ type mockStore struct {
updateProxyHeartbeatFunc func(ctx context.Context, p *proxy.Proxy) error
getActiveProxyClusterAddressesFunc func(ctx context.Context) ([]string, error)
getActiveProxyClusterAddressesForAccFunc func(ctx context.Context, accountID string) ([]string, error)
getActiveProxyVersionsFunc func(ctx context.Context, clusterAddress string) ([]string, error)
cleanupStaleProxiesFunc func(ctx context.Context, d time.Duration) error
getProxyByAccountIDFunc func(ctx context.Context, accountID string) (*proxy.Proxy, error)
countProxiesByAccountIDFunc func(ctx context.Context, accountID string) (int64, error)
@@ -102,6 +105,15 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo
func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool {
return nil
}
func (m *mockStore) GetClusterAllProxiesPrivate(_ context.Context, _ string) *bool {
return nil
}
func (m *mockStore) GetActiveProxyVersions(ctx context.Context, clusterAddress string) ([]string, error) {
if m.getActiveProxyVersionsFunc != nil {
return m.getActiveProxyVersionsFunc(ctx, clusterAddress)
}
return nil, nil
}
func newTestManager(s store) *Manager {
meter := noop.NewMeterProvider().Meter("test")
@@ -112,6 +124,34 @@ func newTestManager(s store) *Manager {
return m
}
func TestClusterSupportsSessionCode(t *testing.T) {
tests := []struct {
name string
versions []string
storeErr error
want bool
}{
{name: "all supported", versions: []string{"0.81.0", "0.81.2"}, want: true},
{name: "one old proxy", versions: []string{"0.81.0", "0.80.0"}},
{name: "missing version", versions: []string{"0.81.0", ""}},
{name: "no active proxies"},
{name: "store error", storeErr: errors.New("db error")},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
s := &mockStore{
getActiveProxyVersionsFunc: func(_ context.Context, _ string) ([]string, error) {
return tt.versions, tt.storeErr
},
}
got := newTestManager(s).ClusterSupportsSessionCode(context.Background(), "cluster.example.com")
assert.Equal(t, tt.want, got)
})
}
}
func TestConnect_WithAccountID(t *testing.T) {
accountID := "acc-123"
@@ -124,7 +164,7 @@ func TestConnect_WithAccountID(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", &accountID, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "0.60.0", &accountID, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
@@ -132,6 +172,7 @@ func TestConnect_WithAccountID(t *testing.T) {
assert.Equal(t, "session-1", savedProxy.SessionID)
assert.Equal(t, "cluster.example.com", savedProxy.ClusterAddress)
assert.Equal(t, "10.0.0.1", savedProxy.IPAddress)
assert.Equal(t, "0.60.0", savedProxy.Version, "reported proxy version should be stored")
assert.Equal(t, &accountID, savedProxy.AccountID)
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
assert.NotNil(t, savedProxy.ConnectedAt)
@@ -147,7 +188,7 @@ func TestConnect_WithoutAccountID(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", nil, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", "", nil, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
@@ -155,6 +196,29 @@ func TestConnect_WithoutAccountID(t *testing.T) {
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
}
func TestConnect_TruncatesOversizedVersion(t *testing.T) {
var savedProxy *proxy.Proxy
s := &mockStore{
saveProxyFunc: func(_ context.Context, p *proxy.Proxy) error {
savedProxy = p
return nil
},
}
// Multi-byte runes make sure the cut counts characters, as varchar does,
// and never splits a rune into invalid UTF-8.
version := strings.Repeat("ü", proxy.MaxVersionLength+10)
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", version, nil, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
assert.Equal(t, proxy.MaxVersionLength, utf8.RuneCountInString(savedProxy.Version), "stored version should be cut to the column width")
assert.True(t, utf8.ValidString(savedProxy.Version), "stored version should remain valid UTF-8")
assert.True(t, strings.HasPrefix(version, savedProxy.Version), "stored version should be a prefix of the reported one")
}
func TestConnect_StoreError(t *testing.T) {
s := &mockStore{
saveProxyFunc: func(_ context.Context, _ *proxy.Proxy) error {
@@ -163,7 +227,7 @@ func TestConnect_StoreError(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", nil, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "", nil, nil)
assert.Error(t, err)
}
@@ -56,6 +56,20 @@ func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *go
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration)
}
// ClusterAllProxiesPrivate mocks base method.
func (m *MockManager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ClusterAllProxiesPrivate", ctx, clusterAddr)
ret0, _ := ret[0].(*bool)
return ret0
}
// ClusterAllProxiesPrivate indicates an expected call of ClusterAllProxiesPrivate.
func (mr *MockManagerMockRecorder) ClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterAllProxiesPrivate", reflect.TypeOf((*MockManager)(nil).ClusterAllProxiesPrivate), ctx, clusterAddr)
}
// ClusterRequireSubdomain mocks base method.
func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool {
m.ctrl.T.Helper()
@@ -112,19 +126,33 @@ func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
}
// Connect mocks base method.
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
// ClusterSupportsSessionCode mocks base method.
func (m *MockManager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
ret := m.ctrl.Call(m, "ClusterSupportsSessionCode", ctx, clusterAddr)
ret0, _ := ret[0].(bool)
return ret0
}
// ClusterSupportsSessionCode indicates an expected call of ClusterSupportsSessionCode.
func (mr *MockManagerMockRecorder) ClusterSupportsSessionCode(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsSessionCode", reflect.TypeOf((*MockManager)(nil).ClusterSupportsSessionCode), ctx, clusterAddr)
}
// Connect mocks base method.
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
ret0, _ := ret[0].(*Proxy)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// Connect indicates an expected call of Connect.
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call {
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
}
// CountAccountProxies mocks base method.
@@ -9,6 +9,9 @@ const (
StatusDisconnected = "disconnected"
)
// MaxVersionLength is the width of the Version column, in characters.
const MaxVersionLength = 255
// Capabilities describes what a proxy can handle, as reported via gRPC.
// Nil fields mean the proxy never reported this capability.
type Capabilities struct {
@@ -31,6 +34,7 @@ type Proxy struct {
SessionID string `gorm:"type:varchar(36)"`
ClusterAddress string `gorm:"type:varchar(255);not null;index:idx_proxy_cluster_status"`
IPAddress string `gorm:"type:varchar(45)"`
Version string `gorm:"type:varchar(255)"`
AccountID *string `gorm:"type:varchar(255);index:idx_proxy_account_id"`
LastSeen time.Time `gorm:"not null;index:idx_proxy_last_seen"`
ConnectedAt *time.Time
@@ -1,6 +1,7 @@
package proxytoken
import (
"context"
"encoding/json"
"net/http"
"time"
@@ -18,13 +19,29 @@ import (
"github.com/netbirdio/netbird/shared/management/status"
)
// RevocationGuard vetoes the tenant-facing revocation of a proxy access
// token. Implementations are supplied by integrations; none is installed by
// default, so every token the caller's account owns may be revoked. It is
// consulted after the ownership check and before the token is revoked. A
// returned status error is written with util.WriteError: its type selects the
// HTTP status and its message is shown to the caller, so it must not carry
// internal detail. Any other error is reported as a generic internal error.
type RevocationGuard interface {
CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error
}
type handler struct {
store store.Store
permissionsManager permissions.Manager
// revocationGuard vetoes revocations. Optional — when nil every owned
// token may be revoked.
revocationGuard RevocationGuard
}
func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, router *mux.Router) {
h := &handler{store: s, permissionsManager: permissionsManager}
// RegisterEndpoints registers the proxy token endpoints. revocationGuard is
// optional; pass nil for no revocation policy.
func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, revocationGuard RevocationGuard, router *mux.Router) {
h := &handler{store: s, permissionsManager: permissionsManager, revocationGuard: revocationGuard}
router.HandleFunc("/reverse-proxies/proxy-tokens", h.listTokens).Methods("GET", "OPTIONS")
router.HandleFunc("/reverse-proxies/proxy-tokens", h.createToken).Methods("POST", "OPTIONS")
router.HandleFunc("/reverse-proxies/proxy-tokens/{tokenId}", h.revokeToken).Methods("DELETE", "OPTIONS")
@@ -154,6 +171,13 @@ func (h *handler) revokeToken(w http.ResponseWriter, r *http.Request) {
return
}
if h.revocationGuard != nil {
if err := h.revocationGuard.CheckProxyAccessTokenRevocation(ctx, token); err != nil {
util.WriteError(ctx, err, w)
return
}
}
if err := h.store.RevokeProxyAccessToken(ctx, tokenID); err != nil {
util.WriteErrorResponse("failed to revoke token", http.StatusInternalServerError, w)
return
@@ -4,6 +4,7 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"testing"
@@ -22,6 +23,7 @@ import (
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/auth"
"github.com/netbirdio/netbird/shared/management/http/api"
"github.com/netbirdio/netbird/shared/management/status"
)
func authContext(accountID, userID string) context.Context {
@@ -273,3 +275,152 @@ func TestRevokeToken_ManagementWideToken(t *testing.T) {
h.revokeToken(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
}
type revocationGuardFunc func(ctx context.Context, token *types.ProxyAccessToken) error
func (f revocationGuardFunc) CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error {
return f(ctx, token)
}
func TestRevokeToken_GuardRefuses(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
// No RevokeProxyAccessToken expectation: a refused revocation must not
// reach the store.
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
var checked *types.ProxyAccessToken
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(_ context.Context, token *types.ProxyAccessToken) error {
checked = token
return status.Errorf(status.PreconditionFailed, "token is in use")
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusPreconditionFailed, w.Code)
assert.Contains(t, w.Body.String(), "token is in use")
require.NotNil(t, checked)
assert.Equal(t, "tok-1", checked.ID)
}
func TestRevokeToken_GuardAllows(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
mockStore.EXPECT().RevokeProxyAccessToken(gomock.Any(), "tok-1").Return(nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
return nil
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusOK, w.Code)
}
func TestRevokeToken_GuardFailure(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
// No RevokeProxyAccessToken expectation: a guard that cannot decide must
// not let the revocation through.
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
return errors.New("connection refused")
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusInternalServerError, w.Code)
assert.Contains(t, w.Body.String(), "internal server error")
assert.NotContains(t, w.Body.String(), "connection refused")
}
func TestRevokeToken_GuardNotConsultedForForeignToken(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
otherAccount := "acc-other"
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &otherAccount,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), "acc-123", "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
// A foreign token must read as not found, not reveal through the guard's
// answer that it belongs to some account's managed proxy.
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
t.Fatal("guard consulted for a token the caller does not own")
return nil
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext("acc-123", "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
}
@@ -30,7 +30,7 @@ func withRealDomainManager(t *testing.T, mgr *Manager, testStore store.Store) {
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", nil, nil)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", "", nil, nil)
require.NoError(t, err)
accountMgr := &mock_server.MockAccountManager{
@@ -125,3 +125,53 @@ func TestUpdateService_RefusesMoveToUnvalidatedDomain(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, "app.proven.example.com", stored.Domain, "the service must keep its original domain")
}
func TestCreateService_DomainDeletedBeforeWrite(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
d, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true)
require.NoError(t, err)
svc := newTestService("app.proven.example.com")
require.NoError(t, mgr.initializeServiceForCreate(ctx, testAccountID, svc))
// Delete after the initial authorization check, before the service transaction starts.
require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID))
err = mgr.persistNewService(ctx, testAccountID, svc)
require.Error(t, err, "an earlier validation result must not authorize a deleted registration")
sErr, ok := status.FromError(err)
require.True(t, ok, "the caller must receive a typed precondition error")
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the service must require current domain authorization")
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, err)
assert.Empty(t, services, "the failed write must not leave a service")
}
func TestUpdateService_DomainDeletedBeforeWrite(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "original.example.com", validationTestCluster, true)
require.NoError(t, err)
d, err := testStore.CreateCustomDomain(ctx, testAccountID, "destination.example.com", validationTestCluster, true)
require.NoError(t, err)
svc, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.original.example.com"))
require.NoError(t, err)
moved := svc.Copy()
moved.Domain = "app.destination.example.com"
cluster, err := mgr.resolveEffectiveCluster(ctx, testAccountID, moved)
require.NoError(t, err)
require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID))
err = testStore.ExecuteInTransaction(ctx, func(tx store.Store) error {
return mgr.executeServiceUpdate(ctx, tx, testAccountID, moved, &serviceUpdateInfo{}, nil, cluster)
})
require.Error(t, err, "a domain deleted after cluster resolution must reject the update")
sErr, ok := status.FromError(err)
require.True(t, ok, "the caller must receive a typed precondition error")
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the move must require current domain authorization")
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, svc.ID)
require.NoError(t, err)
assert.Equal(t, svc.Domain, stored.Domain, "the service must retain its authorized domain")
}
@@ -74,6 +74,7 @@ const unknownHostPlaceholder = "unknown"
// ClusterDeriver derives the proxy cluster from a domain.
type ClusterDeriver interface {
DeriveClusterFromDomain(ctx context.Context, accountID, domain string) (string, error)
ValidateServiceDomain(ctx context.Context, tx store.Store, accountID, domain, cluster string) error
GetClusterDomains() []string
}
@@ -83,6 +84,7 @@ type CapabilityProvider interface {
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
}
type Manager struct {
@@ -331,7 +333,14 @@ func (m *Manager) persistNewService(ctx context.Context, accountID string, svc *
return err
}
if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil {
return err
}
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
return err
}
if svc.Domain != "" {
if err := m.checkDomainAvailable(ctx, transaction, svc.Domain, ""); err != nil {
return err
@@ -365,6 +374,43 @@ func (m *Manager) clusterCustomPorts(ctx context.Context, svc *service.Service)
return m.capabilities.ClusterSupportsCustomPorts(ctx, svc.ProxyCluster)
}
// validatePrivateClusterTargets rejects cluster and direct upstream targets unless
// every active proxy in the service's cluster reports the private capability. The
// mapping reaches all proxies in the cluster, so one non-private proxy would serve
// these targets too. An unreported capability is treated as unsupported. Must be
// called outside a transaction, like clusterCustomPorts.
func (m *Manager) validatePrivateClusterTargets(ctx context.Context, targets []*service.Target, cluster string) error {
target := firstPrivateClusterTarget(targets)
if target == nil {
return nil
}
if private := m.capabilities.ClusterAllProxiesPrivate(ctx, cluster); private != nil && *private {
return nil
}
if target.TargetType == service.TargetTypeCluster {
return status.Errorf(status.InvalidArgument,
"target_type %q requires a proxy cluster with private mode enabled, cluster %s does not support it",
service.TargetTypeCluster, cluster)
}
return status.Errorf(status.InvalidArgument,
"direct_upstream requires a proxy cluster with private mode enabled, cluster %s does not support it", cluster)
}
// firstPrivateClusterTarget returns the first target that only a private cluster may serve.
func firstPrivateClusterTarget(targets []*service.Target) *service.Target {
for _, target := range targets {
if target == nil {
continue
}
if target.TargetType == service.TargetTypeCluster || target.Options.DirectUpstream {
return target
}
}
return nil
}
// ensureL4Port auto-assigns a listen port when needed and validates cluster support.
// customPorts must be pre-computed via clusterCustomPorts before entering a transaction.
func (m *Manager) ensureL4Port(ctx context.Context, tx store.Store, svc *service.Service, customPorts *bool, serviceUpdate bool) error {
@@ -460,7 +506,14 @@ func (m *Manager) persistNewEphemeralService(ctx context.Context, accountID, pee
return err
}
if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil {
return err
}
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
return err
}
if err := m.validateEphemeralPreconditions(ctx, transaction, accountID, peerID, svc); err != nil {
return err
}
@@ -577,6 +630,10 @@ func (m *Manager) persistServiceUpdate(ctx context.Context, accountID string, se
return nil, err
}
if err := m.validatePrivateClusterTargets(ctx, service.Targets, effectiveCluster); err != nil {
return nil, err
}
// Validate subdomain requirement *before* the transaction: the underlying
// capability lookup talks to the main DB pool, and SQLite's single-connection
// pool would self-deadlock if this ran while the tx already held the only
@@ -622,6 +679,9 @@ func (m *Manager) resolveEffectiveCluster(ctx context.Context, accountID string,
}
func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.Store, accountID string, service *service.Service, updateInfo *serviceUpdateInfo, customPorts *bool, effectiveCluster string) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, service, effectiveCluster); err != nil {
return err
}
existingService, err := transaction.GetServiceByID(ctx, store.LockingStrengthUpdate, accountID, service.ID)
if err != nil {
return err
@@ -677,6 +737,13 @@ func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.St
return nil
}
func (m *Manager) validateServiceDomain(ctx context.Context, tx store.Store, accountID string, svc *service.Service, cluster string) error {
if m.clusterDeriver == nil {
return nil
}
return m.clusterDeriver.ValidateServiceDomain(ctx, tx, accountID, svc.Domain, cluster)
}
// validateL4PortDiffOnClusterDiff checks if custom L4 ports are configured and validates port changes across clusters.
// It ensures no port changes if custom ports are unsupported for a given cluster and protocol mode.
// Returns an error if validation fails, otherwise returns nil.
@@ -433,8 +433,8 @@ func TestDeletePeerService_SourcePeerValidation(t *testing.T) {
newProxyServer := func(t *testing.T) *nbgrpc.ProxyServiceServer {
t.Helper()
tokenStore := nbgrpc.NewOneTimeTokenStore(context.Background(), testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(context.Background(), testCacheStore(t))
srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(context.Background(), testCacheStore(t))
srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
return srv
}
@@ -655,6 +655,10 @@ func (d *testClusterDeriver) GetClusterDomains() []string {
return d.domains
}
func (d *testClusterDeriver) ValidateServiceDomain(context.Context, store.Store, string, string, string) error {
return nil
}
const (
testAccountID = "test-account"
testPeerID = "test-peer-1"
@@ -722,8 +726,8 @@ func setupIntegrationTest(t *testing.T) (*Manager, store.Store) {
}
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
@@ -1146,8 +1150,8 @@ func TestDeleteService_DeletesTargets(t *testing.T) {
mockAcct := account.NewMockManager(ctrl)
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
@@ -0,0 +1,218 @@
package manager
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel/metric/noop"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/management/status"
)
// setupPrivateClusterTest wires the real proxy manager as the capability
// provider and connects one proxy to testCluster reporting the given private
// capability. A nil private connects no proxy, so the capability is unreported.
func setupPrivateClusterTest(t *testing.T, private *bool) (*Manager, store.Store) {
t.Helper()
mgr, testStore := setupIntegrationTest(t)
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
mgr.capabilities = proxyMgr
if private != nil {
connectTestProxy(t, proxyMgr, "proxy-1", &proxy.Capabilities{Private: private})
}
return mgr, testStore
}
func connectTestProxy(t *testing.T, proxyMgr *proxymanager.Manager, proxyID string, caps *proxy.Capabilities) {
t.Helper()
_, err := proxyMgr.Connect(context.Background(), proxyID, "session-"+proxyID, testCluster, "127.0.0.1", "", nil, caps)
require.NoError(t, err)
}
func clusterTarget() *rpservice.Target {
return &rpservice.Target{
TargetId: testCluster,
TargetType: rpservice.TargetTypeCluster,
Host: "backend.lan",
Port: 8080,
Protocol: "http",
Enabled: true,
Options: rpservice.TargetOptions{DirectUpstream: true},
}
}
func directUpstreamPeerTarget() *rpservice.Target {
return &rpservice.Target{
TargetId: testPeerID,
TargetType: rpservice.TargetTypePeer,
Host: "backend.lan",
Port: 8080,
Protocol: "http",
Enabled: true,
Options: rpservice.TargetOptions{DirectUpstream: true},
}
}
func TestCreateService_PrivateClusterTargets(t *testing.T) {
tests := []struct {
name string
private *bool
target *rpservice.Target
wantErr string
}{
{name: "cluster target on private cluster", private: boolPtr(true), target: clusterTarget()},
{name: "direct upstream on private cluster", private: boolPtr(true), target: directUpstreamPeerTarget()},
{name: "cluster target on non-private cluster", private: boolPtr(false), target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
{name: "direct upstream on non-private cluster", private: boolPtr(false), target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
{name: "cluster target with unreported capability", private: nil, target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
{name: "direct upstream with unreported capability", private: nil, target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupPrivateClusterTest(t, tc.private)
svc := newTestService("app.test.netbird.io")
svc.Targets = []*rpservice.Target{tc.target}
_, err := mgr.CreateService(ctx, testAccountID, testUserID, svc)
services, listErr := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, listErr)
if tc.wantErr == "" {
require.NoError(t, err)
assert.Len(t, services, 1, "the service should be persisted")
return
}
require.Error(t, err)
assert.Contains(t, err.Error(), tc.wantErr)
sErr, ok := status.FromError(err)
require.True(t, ok, "the caller must receive a typed error")
assert.Equal(t, status.InvalidArgument, sErr.Type(), "the rejection should be an invalid argument")
assert.Empty(t, services, "a rejected service must not be persisted")
})
}
}
// A cluster where only some proxies run in private mode must not accept these
// targets: the mapping is delivered to every proxy in the cluster, so the
// non-private ones would serve the target from their host network as well.
func TestCreateService_MixedClusterRejectsPrivateTargets(t *testing.T) {
tests := []struct {
name string
secondCaps *proxy.Capabilities
}{
{name: "second proxy reports not private", secondCaps: &proxy.Capabilities{Private: boolPtr(false)}},
{name: "second proxy predates capability reporting", secondCaps: nil},
}
for _, tc := range tests {
for _, target := range []*rpservice.Target{clusterTarget(), directUpstreamPeerTarget()} {
t.Run(tc.name+"/"+string(target.TargetType), func(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupPrivateClusterTest(t, boolPtr(true))
connectTestProxy(t, mgr.capabilities.(*proxymanager.Manager), "proxy-2", tc.secondCaps)
svc := newTestService("app.test.netbird.io")
svc.Targets = []*rpservice.Target{target}
_, err := mgr.CreateService(ctx, testAccountID, testUserID, svc)
require.Error(t, err, "a cluster with a non-private proxy must not accept the target")
assert.Contains(t, err.Error(), "requires a proxy cluster with private mode enabled")
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, err)
assert.Empty(t, services, "a rejected service must not be persisted")
})
}
}
}
func TestCreateService_RegularTargetIgnoresPrivateCapability(t *testing.T) {
ctx := context.Background()
mgr, _ := setupPrivateClusterTest(t, boolPtr(false))
_, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
require.NoError(t, err, "a peer target without direct upstream must not need a private cluster")
}
func TestUpdateService_PrivateClusterTargets(t *testing.T) {
tests := []struct {
name string
target *rpservice.Target
wantErr string
}{
{name: "switch to cluster target", target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
{name: "enable direct upstream", target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupPrivateClusterTest(t, boolPtr(false))
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
require.NoError(t, err)
updated := newTestService("app.test.netbird.io")
updated.ID = created.ID
updated.AccountID = testAccountID
updated.Targets = []*rpservice.Target{tc.target}
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated)
require.Error(t, err)
assert.Contains(t, err.Error(), tc.wantErr)
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
require.NoError(t, err)
require.Len(t, stored.Targets, 1)
assert.Equal(t, rpservice.TargetTypePeer, stored.Targets[0].TargetType, "the stored target must be unchanged")
assert.False(t, stored.Targets[0].Options.DirectUpstream, "the stored target must keep direct upstream disabled")
})
}
}
func TestUpdateService_PrivateClusterAllowsClusterTarget(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupPrivateClusterTest(t, boolPtr(true))
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
require.NoError(t, err)
updated := newTestService("app.test.netbird.io")
updated.ID = created.ID
updated.AccountID = testAccountID
updated.Targets = []*rpservice.Target{clusterTarget()}
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated)
require.NoError(t, err)
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
require.NoError(t, err)
require.Len(t, stored.Targets, 1)
assert.Equal(t, rpservice.TargetTypeCluster, stored.Targets[0].TargetType, "the cluster target should be stored")
}
func TestValidatePrivateClusterTargets_NoLookupWithoutPrivateTargets(t *testing.T) {
ctrl := gomock.NewController(t)
// No ClusterAllProxiesPrivate expectation: a lookup would fail the test.
mgr := &Manager{capabilities: proxy.NewMockManager(ctrl)}
targets := []*rpservice.Target{{TargetId: testPeerID, TargetType: rpservice.TargetTypePeer}}
require.NoError(t, mgr.validatePrivateClusterTargets(context.Background(), targets, testCluster))
}
@@ -3,15 +3,16 @@ package manager
import (
"context"
"fmt"
"slices"
"github.com/netbirdio/netbird/management/internals/modules/zones"
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/affectedpeers"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/status"
)
@@ -69,6 +70,9 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
}
zone = zones.NewZone(accountID, zone.Name, zone.Domain, zone.Enabled, zone.EnableSearchDomain, zone.DistributionGroups)
var snap *affectedpeers.Snapshot
change := affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
existingZone, err := transaction.GetZoneByDomain(ctx, accountID, zone.Domain)
if err != nil {
@@ -88,7 +92,15 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
}
if err = transaction.CreateZone(ctx, zone); err != nil {
return fmt.Errorf("failed to create zone: %w", err)
return fmt.Errorf("create zone: %w", err)
}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
return fmt.Errorf("increment network serial: %w", err)
}
return nil
@@ -99,6 +111,8 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneCreated, zone.EventMeta())
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return zone, nil
}
@@ -111,21 +125,26 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
return nil, status.NewPermissionDeniedError()
}
zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID)
if err != nil {
return nil, fmt.Errorf("failed to get zone: %w", err)
}
if zone.Domain != updatedZone.Domain {
return nil, status.Errorf(status.InvalidArgument, "zone domain cannot be updated")
}
zone.Name = updatedZone.Name
zone.Enabled = updatedZone.Enabled
zone.EnableSearchDomain = updatedZone.EnableSearchDomain
zone.DistributionGroups = updatedZone.DistributionGroups
var zone *zones.Zone
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID)
if err != nil {
return fmt.Errorf("get zone: %w", err)
}
if zone.Domain != updatedZone.Domain {
return status.Errorf(status.InvalidArgument, "zone domain cannot be updated")
}
oldGroups := zone.DistributionGroups
zone.Name = updatedZone.Name
zone.Enabled = updatedZone.Enabled
zone.EnableSearchDomain = updatedZone.EnableSearchDomain
zone.DistributionGroups = updatedZone.DistributionGroups
for _, groupID := range zone.DistributionGroups {
_, err = transaction.GetGroupByID(ctx, store.LockingStrengthNone, accountID, groupID)
if err != nil {
@@ -134,7 +153,16 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
}
if err = transaction.UpdateZone(ctx, zone); err != nil {
return fmt.Errorf("failed to update zone: %w", err)
return fmt.Errorf("update zone: %w", err)
}
change = affectedpeers.Change{DistributionGroupIDs: slices.Concat(zone.DistributionGroups, oldGroups)}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
return fmt.Errorf("increment network serial: %w", err)
}
return nil
@@ -145,7 +173,7 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneUpdated, zone.EventMeta())
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationUpdate})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return zone, nil
}
@@ -159,13 +187,23 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID
return status.NewPermissionDeniedError()
}
zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
if err != nil {
return fmt.Errorf("failed to get zone: %w", err)
}
var zone *zones.Zone
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
var eventsToStore []func()
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
if err != nil {
return fmt.Errorf("get zone: %w", err)
}
// Load before delete: the post-delete state no longer references the groups.
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
records, err := transaction.GetZoneDNSRecords(ctx, store.LockingStrengthNone, accountID, zoneID)
if err != nil {
return fmt.Errorf("failed to get records: %w", err)
@@ -207,7 +245,7 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID
event()
}
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationDelete})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return nil
}
@@ -9,11 +9,11 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/zones/records"
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/affectedpeers"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/status"
)
@@ -65,6 +65,8 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
}
var zone *zones.Zone
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
record = records.NewRecord(accountID, zoneID, record.Name, record.Type, record.Content, record.TTL)
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
@@ -82,6 +84,11 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
return fmt.Errorf("failed to create dns record: %w", err)
}
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
err = transaction.IncrementNetworkSerial(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to increment network serial: %w", err)
@@ -96,7 +103,7 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
meta := record.EventMeta(zone.ID, zone.Name)
m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordCreated, meta)
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationCreate})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return record, nil
}
@@ -112,6 +119,8 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
var zone *zones.Zone
var record *records.Record
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
@@ -141,6 +150,11 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
return fmt.Errorf("failed to update dns record: %w", err)
}
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
err = transaction.IncrementNetworkSerial(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to increment network serial: %w", err)
@@ -155,7 +169,7 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
meta := record.EventMeta(zone.ID, zone.Name)
m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordUpdated, meta)
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationUpdate})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return record, nil
}
@@ -171,6 +185,8 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
var record *records.Record
var zone *zones.Zone
var snap *affectedpeers.Snapshot
var change affectedpeers.Change
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
@@ -188,6 +204,11 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
return fmt.Errorf("failed to delete dns record: %w", err)
}
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
return fmt.Errorf("load affected peers: %w", err)
}
err = transaction.IncrementNetworkSerial(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to increment network serial: %w", err)
@@ -202,7 +223,7 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
meta := record.EventMeta(zone.ID, zone.Name)
m.accountManager.StoreEvent(ctx, userID, recordID, accountID, activity.DNSRecordDeleted, meta)
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationDelete})
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
return nil
}
+24 -12
View File
@@ -32,17 +32,18 @@ import (
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
networkmapdbfactory "github.com/netbirdio/netbird/management/internals/network_map_db/factory"
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
"github.com/netbirdio/netbird/management/internals/shared/db"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/activity"
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
nbcache "github.com/netbirdio/netbird/management/server/cache"
nbContext "github.com/netbirdio/netbird/management/server/context"
nbhttp "github.com/netbirdio/netbird/management/server/http"
"github.com/netbirdio/netbird/management/server/http/middleware"
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/telemetry"
mgmtProto "github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/ratelimit"
"github.com/netbirdio/netbird/util/crypt"
)
@@ -84,9 +85,20 @@ func (s *BaseServer) CacheStore() nbcache.Store {
})
}
// DBConn opens the database connection shared by the store and the domain repositories.
func (s *BaseServer) DBConn() *db.Conn {
return Create(s, func() *db.Conn {
conn, err := store.OpenConn(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir)
if err != nil {
log.Fatalf("failed to open database connection: %v", err)
}
return conn
})
}
func (s *BaseServer) Store() store.Store {
return Create(s, func() store.Store {
store, err := store.NewStore(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir, s.Metrics(), false)
store, err := store.NewSqlStore(context.Background(), s.DBConn(), s.Metrics(), false)
if err != nil {
log.Fatalf("failed to create store: %v", err)
}
@@ -147,7 +159,7 @@ func (s *BaseServer) EventStore() activity.Store {
func (s *BaseServer) APIHandler() http.Handler {
return Create(s, func() http.Handler {
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager())
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager(), nil)
if err != nil {
log.Fatalf("failed to create API handler: %v", err)
}
@@ -171,10 +183,10 @@ func (s *BaseServer) Router() *mux.Router {
})
}
func (s *BaseServer) RateLimiter() *middleware.APIRateLimiter {
return Create(s, func() *middleware.APIRateLimiter {
cfg, enabled := middleware.RateLimiterConfigFromEnv()
limiter := middleware.NewAPIRateLimiter(cfg)
func (s *BaseServer) RateLimiter() *ratelimit.APIRateLimiter {
return Create(s, func() *ratelimit.APIRateLimiter {
cfg, enabled := ratelimit.RateLimiterConfigFromEnv()
limiter := ratelimit.NewAPIRateLimiter(cfg)
limiter.SetEnabled(enabled)
return limiter
})
@@ -236,7 +248,7 @@ func (s *BaseServer) GRPCServer() *grpc.Server {
func (s *BaseServer) ReverseProxyGRPCServer() *nbgrpc.ProxyServiceServer {
return Create(s, func() *nbgrpc.ProxyServiceServer {
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.PKCEVerifierStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.SingleUseStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
s.AfterInit(func(s *BaseServer) {
proxyService.SetServiceManager(s.ServiceManager())
proxyService.SetActivityManager(s.ProxyActivityManager())
@@ -293,9 +305,9 @@ func (s *BaseServer) ProxyTokenStore() *nbgrpc.OneTimeTokenStore {
})
}
func (s *BaseServer) PKCEVerifierStore() *nbgrpc.PKCEVerifierStore {
return Create(s, func() *nbgrpc.PKCEVerifierStore {
return nbgrpc.NewPKCEVerifierStore(context.Background(), s.CacheStore())
func (s *BaseServer) SingleUseStore() *nbgrpc.SingleUseStore {
return Create(s, func() *nbgrpc.SingleUseStore {
return nbgrpc.NewSingleUseStore(context.Background(), s.CacheStore())
})
}
@@ -308,7 +320,7 @@ func (s *BaseServer) ProxyActivityManager() proxyactivity.Manager {
func (s *BaseServer) AccessLogsManager() accesslogs.Manager {
return Create(s, func() accesslogs.Manager {
accessLogManager := accesslogsmanager.NewManager(s.Store(), s.PermissionsManager(), s.GeoLocationManager())
accessLogManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(s.DBConn()), s.Store(), s.PermissionsManager(), s.GeoLocationManager())
accessLogManager.StartPeriodicCleanup(
context.Background(),
s.Config.ReverseProxy.AccessLogRetentionDays,
+1
View File
@@ -103,6 +103,7 @@ func (s *BaseServer) AccountManager() account.Manager {
s.AfterInit(func(s *BaseServer) {
accountManager.SetServiceManager(s.ServiceManager())
accountManager.AddAccountDeletionHook(s.AgentNetworkManager().RemoveAccountGateway)
})
return accountManager
+18
View File
@@ -23,6 +23,8 @@ import (
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/metrics"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/lifecycle"
"github.com/netbirdio/netbird/shared/profiling"
"github.com/netbirdio/netbird/util/wsproxy"
wsproxyserver "github.com/netbirdio/netbird/util/wsproxy/server"
"github.com/netbirdio/netbird/version"
@@ -36,6 +38,8 @@ const (
DefaultSelfHostedDomain = "netbird.selfhosted"
ContainerKeyBaseServer = "baseServer"
applicationName = "management"
)
type Server interface {
@@ -82,6 +86,8 @@ type BaseServer struct {
errCh chan error
wg sync.WaitGroup
cancel context.CancelFunc
lifecycle.StopHandlers
}
// Config holds the configuration parameters for creating a new server
@@ -117,6 +123,9 @@ func NewServer(cfg *Config) *BaseServer {
}
s.container[ContainerKeyBaseServer] = s
stopProfiling := profiling.Start(applicationName)
s.OnStop(stopProfiling)
return s
}
@@ -126,6 +135,14 @@ func (s *BaseServer) AfterInit(fn func(s *BaseServer)) {
// Start begins listening for HTTP requests on the configured address
func (s *BaseServer) Start(ctx context.Context) error {
if err := s.start(ctx); err != nil {
s.RunStopHandlers()
return err
}
return nil
}
func (s *BaseServer) start(ctx context.Context) error {
srvCtx, cancel := context.WithCancel(ctx)
s.cancel = cancel
s.errCh = make(chan error, 4)
@@ -278,6 +295,7 @@ func (s *BaseServer) setupTLS(ctx context.Context) (bool, error) {
func (s *BaseServer) Stop() error {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
defer s.RunStopHandlers()
if s.domainCleanupStop != nil {
s.domainCleanupStop()
}
+121
View File
@@ -0,0 +1,121 @@
package db
import (
"context"
"fmt"
"os"
"runtime"
"strconv"
"time"
"github.com/jackc/pgx/v5/pgxpool"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
)
const (
defaultTransactionTimeout = 5 * time.Minute
connMaxLifetime = time.Hour
connMaxIdleTime = 3 * time.Minute
)
// TxMetrics receives the duration of every committed top-level transaction.
type TxMetrics interface {
CountTransactionDuration(duration time.Duration)
}
// Conn is the database connection shared by all repositories: one gorm handle,
// the pgx pool of a Postgres deployment and the engine they talk to.
type Conn struct {
db *gorm.DB
pool *pgxpool.Pool
engine Engine
txTimeout time.Duration
metrics TxMetrics
}
// NewConn takes ownership of an open gorm handle and pool once it returns
// without error, applying the connection limits and transaction timeout
// configured through the environment.
func NewConn(ctx context.Context, gormDB *gorm.DB, engine Engine, pool *pgxpool.Pool) (*Conn, error) {
sqlDB, err := gormDB.DB()
if err != nil {
return nil, err
}
txTimeout := defaultTransactionTimeout
if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" {
if parsed, err := time.ParseDuration(v); err == nil {
txTimeout = parsed
}
}
log.WithContext(ctx).Infof("Setting transaction timeout to %v", txTimeout)
conns := runtime.NumCPU()
configuredConns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS"))
connsConfigured := err == nil
if connsConfigured {
conns = configuredConns
}
if engine == SqliteStoreEngine {
if connsConfigured {
log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1")
}
conns = 1
}
sqlDB.SetMaxOpenConns(conns)
sqlDB.SetMaxIdleConns(conns)
sqlDB.SetConnMaxLifetime(connMaxLifetime)
sqlDB.SetConnMaxIdleTime(connMaxIdleTime)
log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v",
conns, conns, connMaxLifetime, connMaxIdleTime)
return &Conn{db: gormDB, pool: pool, engine: engine, txTimeout: txTimeout}, nil
}
// DB returns the handle a query must run on: the transaction when tx is set,
// otherwise the shared connection.
func (c *Conn) DB(tx *Tx) *gorm.DB {
if tx != nil {
return tx.db
}
return c.db
}
// Pool returns the pgx pool for read paths that bypass gorm. It is nil on
// engines other than Postgres and inside a transaction, where the pool would
// not see the uncommitted writes.
func (c *Conn) Pool(tx *Tx) *pgxpool.Pool {
if tx != nil {
return nil
}
return c.pool
}
func (c *Conn) Engine() Engine {
return c.engine
}
// SetTxMetrics registers the sink that receives transaction durations.
func (c *Conn) SetTxMetrics(metrics TxMetrics) {
c.metrics = metrics
}
// AutoMigrate creates or updates the tables of the given models.
func (c *Conn) AutoMigrate(models ...any) error {
return c.db.AutoMigrate(models...)
}
// Close releases the gorm connection and the pgx pool.
func (c *Conn) Close() error {
if c.pool != nil {
c.pool.Close()
}
sqlDB, err := c.db.DB()
if err != nil {
return fmt.Errorf("get db: %w", err)
}
return sqlDB.Close()
}
+148
View File
@@ -0,0 +1,148 @@
package db
import (
"context"
"errors"
"path/filepath"
"testing"
"time"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
type testRow struct {
ID uint `gorm:"primaryKey"`
Name string
}
func openTestConn(t *testing.T) *Conn {
t.Helper()
conn, err := OpenSqliteFile(context.Background(), t.TempDir(), SqliteFileName)
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, conn.Close()) })
require.NoError(t, conn.AutoMigrate(&testRow{}))
return conn
}
func countRows(t *testing.T, conn *Conn) int64 {
t.Helper()
var count int64
require.NoError(t, conn.DB(nil).Model(&testRow{}).Count(&count).Error)
return count
}
func TestNewConn_ReadsTransactionTimeoutFromEnv(t *testing.T) {
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "1s")
conn := openTestConn(t)
assert.Equal(t, time.Second, conn.txTimeout)
assert.Equal(t, SqliteStoreEngine, conn.Engine())
}
func TestRunInTx_CommitsOnSuccess(t *testing.T) {
conn := openTestConn(t)
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
})
require.NoError(t, err)
assert.EqualValues(t, 1, countRows(t, conn))
}
func TestRunInTx_RollsBackOnError(t *testing.T) {
conn := openTestConn(t)
failure := errors.New("boom")
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
return failure
})
require.ErrorIs(t, err, failure)
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_RollsBackOnPanic(t *testing.T) {
conn := openTestConn(t)
require.Panics(t, func() {
_ = conn.RunInTx(context.Background(), func(tx *Tx) error {
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
panic("boom")
})
})
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_FailsWhenTimeoutExceeded(t *testing.T) {
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "50ms")
conn := openTestConn(t)
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
time.Sleep(100 * time.Millisecond)
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
})
require.ErrorIs(t, err, context.DeadlineExceeded)
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_ReportsDurationToMetrics(t *testing.T) {
conn := openTestConn(t)
metrics := &recordingMetrics{}
conn.SetTxMetrics(metrics)
require.NoError(t, conn.RunInTx(context.Background(), func(*Tx) error { return nil }))
assert.Equal(t, 1, metrics.calls)
}
func TestConn_DBSelectsTransactionHandle(t *testing.T) {
conn := openTestConn(t)
assert.Same(t, conn.db, conn.DB(nil))
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
assert.Same(t, tx.db, conn.DB(tx))
assert.NotSame(t, conn.db, conn.DB(tx))
return nil
})
require.NoError(t, err)
}
func TestConn_PoolIsUnavailableInsideTransaction(t *testing.T) {
conn := openTestConn(t)
conn.pool = &pgxpool.Pool{}
defer func() { conn.pool = nil }()
assert.Same(t, conn.pool, conn.Pool(nil))
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
assert.Nil(t, conn.Pool(tx))
return nil
})
require.NoError(t, err)
}
type recordingMetrics struct {
calls int
}
func (m *recordingMetrics) CountTransactionDuration(time.Duration) {
m.calls++
}
func TestNewConn_MaxOpenConnsFromEnv(t *testing.T) {
t.Setenv("NB_SQL_MAX_OPEN_CONNS", "7")
gormDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "store.db")), GormConfig())
require.NoError(t, err)
conn, err := NewConn(context.Background(), gormDB, PostgresStoreEngine, nil)
require.NoError(t, err)
t.Cleanup(func() { _ = conn.Close() })
sqlDB, err := conn.DB(nil).DB()
require.NoError(t, err)
assert.Equal(t, 7, sqlDB.Stats().MaxOpenConnections)
sqliteDB, err := openTestConn(t).DB(nil).DB()
require.NoError(t, err)
assert.Equal(t, 1, sqliteDB.Stats().MaxOpenConnections)
}
@@ -0,0 +1,23 @@
package dbtest
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
// NewConn opens a fresh SQLite database in a temporary directory, migrates the
// given models and closes the connection when the test ends. It ignores
// NB_STORE_ENGINE_SQLITE_FILE, so a developer's configured database is never
// touched, and is safe to call from parallel tests.
func NewConn(t testing.TB, models ...any) *db.Conn {
t.Helper()
conn, err := db.OpenSqliteFile(context.Background(), t.TempDir(), db.SqliteFileName)
require.NoError(t, err)
t.Cleanup(func() { _ = conn.Close() })
require.NoError(t, conn.AutoMigrate(models...))
return conn
}
@@ -0,0 +1,31 @@
package dbtest
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
func TestNewConn_IgnoresSqliteFileOverride(t *testing.T) {
override := filepath.Join(t.TempDir(), "configured.db")
t.Setenv("NB_STORE_ENGINE_SQLITE_FILE", override)
conn := NewConn(t)
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
_, err := os.Stat(override)
require.ErrorIs(t, err, os.ErrNotExist)
}
func TestNewConn_Parallel(t *testing.T) {
t.Parallel()
conn := NewConn(t)
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
}
+10
View File
@@ -0,0 +1,10 @@
package db
// Engine identifies the SQL engine behind a Conn.
type Engine string
const (
SqliteStoreEngine Engine = "sqlite"
PostgresStoreEngine Engine = "postgres"
MysqlStoreEngine Engine = "mysql"
)
+12
View File
@@ -0,0 +1,12 @@
package db
// LockingStrength is the row lock a query holds until its transaction ends.
type LockingStrength string
const (
LockingStrengthUpdate LockingStrength = "UPDATE"
LockingStrengthShare LockingStrength = "SHARE"
LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE"
LockingStrengthKeyShare LockingStrength = "KEY SHARE"
LockingStrengthNone LockingStrength = "NONE"
)
+176
View File
@@ -0,0 +1,176 @@
package db
import (
"context"
"fmt"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"time"
"github.com/jackc/pgx/v5/pgxpool"
"gorm.io/driver/mysql"
"gorm.io/driver/postgres"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// SqliteFileName is the default SQLite database file inside the data directory.
const SqliteFileName = "store.db"
// PoolConfig sizes the pgx pool a Postgres deployment uses for the read paths
// that bypass gorm.
type PoolConfig struct {
MaxConns int32
MinConns int32
MaxConnLifetime time.Duration
HealthCheckPeriod time.Duration
}
var DefaultPoolConfig = PoolConfig{
MaxConns: 30,
MinConns: 1,
MaxConnLifetime: 60 * time.Minute,
HealthCheckPeriod: time.Minute,
}
// GormConfig is the configuration every engine is opened with.
func GormConfig() *gorm.Config {
return &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
CreateBatchSize: 400,
}
}
// OpenSqlite opens the SQLite database in dataDir, or the file named by
// NB_STORE_ENGINE_SQLITE_FILE.
func OpenSqlite(ctx context.Context, dataDir string) (*Conn, error) {
storeFile := SqliteFileName
if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" {
storeFile = envFile
}
return OpenSqliteFile(ctx, dataDir, storeFile)
}
// OpenSqliteFile opens the SQLite database storeFile, resolved against dataDir
// when relative. storeFile may carry SQLite URI query parameters.
func OpenSqliteFile(ctx context.Context, dataDir, storeFile string) (*Conn, error) {
// Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc")
filePath, query, hasQuery := strings.Cut(storeFile, "?")
connStr := filePath
if !filepath.IsAbs(filePath) {
connStr = filepath.Join(dataDir, filePath)
}
// Compose query parameters. User-provided ?_busy_timeout (or its mattn alias
// ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at
// most that long on a lock instead of blocking the only Go-side connection.
// mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so
// the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared
// stays the default on non-Windows for the same reason as before.
parsed, _ := url.ParseQuery(query)
var defaults []string
if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" {
defaults = append(defaults, "_busy_timeout=30000")
}
if !hasQuery && runtime.GOOS != "windows" {
// To avoid `The process cannot access the file because it is being used by another process` on Windows
defaults = append(defaults, "cache=shared")
}
parts := defaults
if hasQuery {
parts = append(parts, query)
}
if len(parts) > 0 {
connStr += "?" + strings.Join(parts, "&")
}
gormDB, err := gorm.Open(sqlite.Open(connStr), GormConfig())
if err != nil {
return nil, err
}
conn, err := NewConn(ctx, gormDB, SqliteStoreEngine, nil)
if err != nil {
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
// OpenPostgres opens a Postgres database through gorm and a pgx pool sized by pool.
func OpenPostgres(ctx context.Context, dsn string, pool PoolConfig) (*Conn, error) {
gormDB, err := gorm.Open(postgres.Open(dsn), GormConfig())
if err != nil {
return nil, err
}
pgxPool, err := newPgxPool(ctx, dsn, pool)
if err != nil {
closeGorm(gormDB)
return nil, err
}
conn, err := NewConn(ctx, gormDB, PostgresStoreEngine, pgxPool)
if err != nil {
pgxPool.Close()
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
// MysqlDSN adds the connection parameters every MySQL handle needs, keeping
// the options already present in dsn.
func MysqlDSN(dsn string) string {
separator := "?"
if strings.Contains(dsn, "?") {
separator = "&"
}
return dsn + separator + "charset=utf8&parseTime=True&loc=Local"
}
// OpenMysql opens a MySQL database through gorm.
func OpenMysql(ctx context.Context, dsn string) (*Conn, error) {
gormDB, err := gorm.Open(mysql.Open(MysqlDSN(dsn)), GormConfig())
if err != nil {
return nil, err
}
conn, err := NewConn(ctx, gormDB, MysqlStoreEngine, nil)
if err != nil {
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
func newPgxPool(ctx context.Context, dsn string, cfg PoolConfig) (*pgxpool.Pool, error) {
config, err := pgxpool.ParseConfig(dsn)
if err != nil {
return nil, fmt.Errorf("unable to parse database config: %w", err)
}
config.MaxConns = cfg.MaxConns
config.MinConns = cfg.MinConns
config.MaxConnLifetime = cfg.MaxConnLifetime
config.HealthCheckPeriod = cfg.HealthCheckPeriod
pool, err := pgxpool.NewWithConfig(ctx, config)
if err != nil {
return nil, fmt.Errorf("unable to create connection pool: %w", err)
}
if err := pool.Ping(ctx); err != nil {
pool.Close()
return nil, fmt.Errorf("unable to ping database: %w", err)
}
return pool, nil
}
func closeGorm(gormDB *gorm.DB) {
if sqlDB, err := gormDB.DB(); err == nil {
_ = sqlDB.Close()
}
}
@@ -0,0 +1,12 @@
package db
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestMysqlDSN(t *testing.T) {
assert.Equal(t, "user:pw@tcp(host:3306)/db?charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db"))
assert.Equal(t, "user:pw@tcp(host:3306)/db?tls=true&charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db?tls=true"))
}
@@ -0,0 +1,105 @@
package db
import (
"context"
"errors"
"fmt"
"runtime/debug"
"time"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
)
// Tx is an open transaction handed to repository calls; nil means autocommit.
type Tx struct {
db *gorm.DB
}
// RunInTx runs fn in one transaction that commits when fn returns nil and rolls
// back otherwise, bounded by the configured transaction timeout.
func (c *Conn) RunInTx(ctx context.Context, fn func(tx *Tx) error) error {
timeoutCtx, cancel := context.WithTimeout(ctx, c.txTimeout)
defer cancel()
startTime := time.Now()
tx := c.db.WithContext(timeoutCtx).Begin()
if tx.Error != nil {
return tx.Error
}
defer func() {
if r := recover(); r != nil {
tx.Rollback()
panic(r)
}
}()
if err := c.applyStatementTimeouts(tx); err != nil {
tx.Rollback()
return err
}
err := c.withForeignKeyChecksDisabled(tx, func() error {
return fn(&Tx{db: tx})
})
if err != nil {
tx.Rollback()
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction", startTime)
return err
}
if err := tx.Commit().Error; err != nil {
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction commit", startTime)
return err
}
log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime))
if c.metrics != nil {
c.metrics.CountTransactionDuration(time.Since(startTime))
}
return nil
}
func (c *Conn) applyStatementTimeouts(tx *gorm.DB) error {
if c.engine != PostgresStoreEngine {
return nil
}
if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil {
return fmt.Errorf("failed to set statement timeout: %w", err)
}
if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil {
return fmt.Errorf("failed to set lock timeout: %w", err)
}
return nil
}
// withForeignKeyChecksDisabled runs fn with MySQL's FK checks off, which avoids
// deadlocks on MySQL and Aurora without needing SUPER privilege. The setting is
// session-scoped and survives a rollback, so it is turned back on whenever fn
// returns or panics; otherwise the pooled connection would keep it disabled.
func (c *Conn) withForeignKeyChecksDisabled(tx *gorm.DB, fn func() error) (err error) {
if c.engine != MysqlStoreEngine {
return fn()
}
if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil {
return fmt.Errorf("failed to disable FK checks: %w", err)
}
defer func() {
restoreErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error
if restoreErr == nil {
return
}
if err == nil {
err = fmt.Errorf("failed to re-enable FK checks: %w", restoreErr)
return
}
log.WithContext(tx.Statement.Context).Warnf("failed to re-enable FK checks after failed transaction: %v", restoreErr)
}()
return fn()
}
func (c *Conn) logIfTimedOut(ctx, timeoutCtx context.Context, err error, phase string, startTime time.Time) {
if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) {
log.WithContext(ctx).Warnf("%s exceeded %s timeout after %v, stack: %s", phase, c.txTimeout, time.Since(startTime), debug.Stack())
}
}
@@ -14,6 +14,7 @@ const (
baseBlockDuration = 10 * time.Minute // Duration for which a peer is banned after exceeding the reconnection limit
reconnLimitForBan = 30 // Number of reconnections within the reconnTreshold that triggers a ban
metaChangeLimit = 5 // Number of reconnections with different metadata that triggers a ban of one peer
maxBanLevel = 6 // Highest ban level; the ban duration doubles per level up to this one
)
type lfConfig struct {
@@ -21,6 +22,7 @@ type lfConfig struct {
baseBlockDuration time.Duration
reconnLimitForBan int
metaChangeLimit int
maxBanLevel int
}
func initCfg() *lfConfig {
@@ -29,6 +31,7 @@ func initCfg() *lfConfig {
baseBlockDuration: baseBlockDuration,
reconnLimitForBan: reconnLimitForBan,
metaChangeLimit: metaChangeLimit,
maxBanLevel: maxBanLevel,
}
}
@@ -102,11 +105,18 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
return
}
if state.isBanned && now.After(state.banExpiresAt) {
if state.isBanned {
if now.Before(state.banExpiresAt) {
return
}
state.isBanned = false
}
if state.banLevel > 0 && now.Sub(state.lastSeen) > (2*l.cfg.baseBlockDuration) {
quietSince := state.lastSeen
if state.banExpiresAt.After(quietSince) {
quietSince = state.banExpiresAt
}
if state.banLevel > 0 && now.Sub(quietSince) > (2*l.cfg.baseBlockDuration) {
state.banLevel = 0
}
@@ -124,10 +134,17 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
return
}
if now.Sub(state.sessionStart) >= l.cfg.reconnThreshold {
state.sessionStart = now
state.sessionCounter = 0
}
state.sessionCounter++
if state.sessionCounter > l.cfg.reconnLimitForBan && now.Sub(state.sessionStart) < l.cfg.reconnThreshold {
if state.sessionCounter > l.cfg.reconnLimitForBan {
state.isBanned = true
state.banLevel++
if state.banLevel < l.cfg.maxBanLevel {
state.banLevel++
}
backoffFactor := math.Pow(2, float64(state.banLevel-1))
duration := time.Duration(float64(l.cfg.baseBlockDuration) * backoffFactor)
@@ -20,6 +20,7 @@ func testAdvancedCfg() *lfConfig {
baseBlockDuration: 100 * time.Millisecond,
reconnLimitForBan: 3,
metaChangeLimit: 2,
maxBanLevel: 3,
}
}
@@ -157,6 +158,187 @@ func (s *LoginFilterTestSuite) TestMetaChangeIsAllowedAfterWindowResets() {
s.Equal(1, s.filter.logged[pubKey].metaChangeCounter, "meta change counter should reset")
}
func (s *LoginFilterTestSuite) TestReconnectStormAfterQuietPeriodTriggersBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
s.filter.addLogin(pubKey, meta)
s.Require().Contains(s.filter.logged, pubKey)
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
s.filter.addLogin(pubKey, meta)
s.Equal(1, s.filter.logged[pubKey].sessionCounter, "expired window should restart the count")
for i := 1; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
s.filter.addLogin(pubKey, meta)
s.False(s.filter.allowLogin(pubKey, meta))
s.True(s.filter.logged[pubKey].isBanned)
}
func (s *LoginFilterTestSuite) TestReconnectStormAfterBanExpiresTriggersBanAgain() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.Require().True(s.filter.logged[pubKey].isBanned)
expired := time.Now().Add(-(s.filter.cfg.baseBlockDuration + time.Second))
s.filter.logged[pubKey].banExpiresAt = expired
s.filter.logged[pubKey].sessionStart = expired
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(2, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestSlowReconnectsAcrossWindowsDoNotBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
for i := 0; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
}
func (s *LoginFilterTestSuite) TestBanLevelEscalatesWhenStormResumesRightAfterBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
banTime := time.Now().Add(-3 * s.filter.cfg.baseBlockDuration)
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
isBanned: true,
banLevel: 1,
banExpiresAt: time.Now().Add(-time.Millisecond),
sessionStart: banTime,
lastSeen: banTime,
}
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(2, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestBanLevelResetsAfterQuietPeriodFollowingBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
quiet := 2*s.filter.cfg.baseBlockDuration + time.Second
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
banLevel: 2,
banExpiresAt: time.Now().Add(-s.filter.cfg.baseBlockDuration),
lastSeen: time.Now().Add(-2 * quiet),
}
s.filter.addLogin(pubKey, meta)
s.Equal(2, s.filter.logged[pubKey].banLevel, "ban ended more recently than the quiet period")
s.filter.logged[pubKey].banExpiresAt = time.Now().Add(-quiet)
s.filter.logged[pubKey].lastSeen = time.Now().Add(-2 * quiet)
s.filter.addLogin(pubKey, meta)
s.Equal(0, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestBanDurationIsCappedAtMaxLevel() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
maxLevel := s.filter.cfg.maxBanLevel
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
banLevel: maxLevel,
sessionStart: time.Now(),
lastSeen: time.Now(),
}
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(maxLevel, s.filter.logged[pubKey].banLevel)
expected := s.filter.cfg.baseBlockDuration << (maxLevel - 1)
s.InDelta(expected, s.filter.logged[pubKey].banExpiresAt.Sub(s.filter.logged[pubKey].lastSeen), float64(time.Millisecond))
}
func (s *LoginFilterTestSuite) TestEstablishedPeerReconnectingOnceIsAllowed() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
longAgo := time.Now().Add(-time.Hour)
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
sessionCounter: 1,
sessionStart: longAgo,
lastSeen: longAgo,
metaChangeWindowStart: longAgo,
metaChangeCounter: 1,
}
s.True(s.filter.allowLogin(pubKey, meta))
s.filter.addLogin(pubKey, meta)
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
s.Equal(1, s.filter.logged[pubKey].sessionCounter)
}
func (s *LoginFilterTestSuite) TestLoginsDuringActiveBanDoNotExtendIt() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.Require().True(s.filter.logged[pubKey].isBanned)
expiresAt := time.Now().Add(time.Hour)
s.filter.logged[pubKey].banExpiresAt = expiresAt
lastSeen := s.filter.logged[pubKey].lastSeen
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(1, s.filter.logged[pubKey].banLevel)
s.Equal(expiresAt, s.filter.logged[pubKey].banExpiresAt)
s.Equal(lastSeen, s.filter.logged[pubKey].lastSeen)
s.Equal(0, s.filter.logged[pubKey].sessionCounter)
}
func BenchmarkHashingMethods(b *testing.B) {
meta := nbpeer.PeerSystemMeta{
WtVersion: "1.25.1",
@@ -1,55 +0,0 @@
package grpc
import (
"context"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/store"
log "github.com/sirupsen/logrus"
nbcache "github.com/netbirdio/netbird/management/server/cache"
)
// PKCEVerifierStore manages PKCE verifiers for OAuth flows.
// Supports both in-memory and Redis storage via NB_IDP_CACHE_REDIS_ADDRESS env var.
type PKCEVerifierStore struct {
cache nbcache.Store
ctx context.Context
}
// NewPKCEVerifierStore creates a PKCE verifier store using the provided shared cache store.
func NewPKCEVerifierStore(ctx context.Context, cacheStore nbcache.Store) *PKCEVerifierStore {
return &PKCEVerifierStore{
cache: cacheStore,
ctx: ctx,
}
}
// Store saves a PKCE verifier associated with an OAuth state parameter.
// The verifier is stored with the specified TTL and will be automatically deleted after expiration.
func (s *PKCEVerifierStore) Store(state, verifier string, ttl time.Duration) error {
if err := s.cache.Set(s.ctx, state, verifier, store.WithExpiration(ttl)); err != nil {
return fmt.Errorf("failed to store PKCE verifier: %w", err)
}
log.Debugf("Stored PKCE verifier for state (expires in %s)", ttl)
return nil
}
// LoadAndDelete retrieves and removes a PKCE verifier for the given state.
// Returns the verifier and true if found, or empty string and false if not found.
// This enforces single-use semantics for PKCE verifiers.
func (s *PKCEVerifierStore) LoadAndDelete(state string) (string, bool) {
verifier, found, err := s.cache.GetDel(s.ctx, state)
if err != nil {
log.Warnf("Failed to consume PKCE verifier: %v", err)
return "", false
}
if !found {
log.Debug("PKCE verifier not found for state")
return "", false
}
return verifier, true
}
+74 -27
View File
@@ -27,8 +27,6 @@ import (
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/peers"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
@@ -42,6 +40,7 @@ import (
"github.com/netbirdio/netbird/management/server/users"
proxyauth "github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/shared/hash/argon2id"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/proto"
nbstatus "github.com/netbirdio/netbird/shared/management/status"
)
@@ -102,7 +101,8 @@ type ProxyServiceServer struct {
mu sync.RWMutex
// Manager for reverse proxy operations
serviceManager rpservice.Manager
serviceManager rpservice.Manager
credentialLimits credentialVerificationLimiter
// agentNetworkSynth produces synthesised reverse-proxy services from
// Agent Network state. Optional — when nil the snapshot path only ships
// persisted services.
@@ -141,8 +141,8 @@ type ProxyServiceServer struct {
// OIDC configuration for proxy authentication
oidcConfig ProxyOIDCConfig
// Store for PKCE verifiers
pkceVerifierStore *PKCEVerifierStore
// singleUseStore backs both PKCE verifiers and OIDC session exchange codes.
singleUseStore *SingleUseStore
// tokenTTL is the lifetime of one-time tokens generated for proxy
// authentication. Defaults to defaultProxyTokenTTL when zero.
@@ -157,6 +157,13 @@ type ProxyServiceServer struct {
const pkceVerifierTTL = 10 * time.Minute
const sessionCodeTTL = 60 * time.Second
const sessionCodeCacheNamespace = "proxy:session"
// The signed nonce binds the handoff mode without changing the state format.
const sessionCodeNoncePrefix = "code."
const defaultProxyTokenTTL = 5 * time.Minute
const defaultSnapshotBatchSize = 500
@@ -207,13 +214,13 @@ func enforceAccountScope(ctx context.Context, requestAccountID string) error {
}
// NewProxyServiceServer creates a new proxy service server.
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, pkceStore *PKCEVerifierStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, singleUseStore *SingleUseStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
ctx, cancel := context.WithCancel(context.Background())
s := &ProxyServiceServer{
accessLogManager: accessLogMgr,
oidcConfig: oidcConfig,
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
peersManager: peersManager,
usersManager: usersManager,
idpManager: idpManager,
@@ -242,9 +249,10 @@ func (s *ProxyServiceServer) cleanupStaleProxies(ctx context.Context) {
}
}
// Close stops background goroutines.
// Close stops background goroutines and releases credential verification state.
func (s *ProxyServiceServer) Close() {
s.cancel()
s.credentialLimits.close()
}
// SetServiceManager sets the service manager. Must be called before serving.
@@ -304,6 +312,16 @@ func (s *ProxyServiceServer) proxyConnectAuthorizer() ProxyConnectAuthorizer {
return s.connectAuthorizer
}
// GenerateSessionCode creates a single-use code for the given session token.
func (s *ProxyServiceServer) GenerateSessionCode(sessionToken string) (code string, ok bool) {
code, err := s.singleUseStore.Generate(sessionCodeCacheNamespace, sessionToken, sessionCodeTTL)
if err != nil {
log.WithError(err).Error("failed to generate proxy session code")
return "", false
}
return code, true
}
// CheckLLMPolicyLimits is the pre-flight policy gate the proxy calls before
// forwarding an LLM request upstream. Delegates to the agent-network selector,
// which scores applicable policies by remaining headroom and returns the
@@ -412,6 +430,7 @@ func (s *ProxyServiceServer) SetProxyController(proxyController proxy.Controller
type proxyConnectParams struct {
proxyID string
address string
version string
capabilities *proto.ProxyCapabilities
}
@@ -422,6 +441,7 @@ func (s *ProxyServiceServer) GetMappingUpdate(req *proto.GetMappingUpdateRequest
return err
}
params.capabilities = req.GetCapabilities()
params.version = req.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
stream: stream,
@@ -455,6 +475,7 @@ func (s *ProxyServiceServer) SyncMappings(stream proto.ProxyService_SyncMappings
return err
}
params.capabilities = init.GetCapabilities()
params.version = init.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
syncStream: stream,
@@ -566,7 +587,7 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params
}
}
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, accountID, caps)
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, params.version, accountID, caps)
if err != nil {
cancel()
if accountID != nil {
@@ -1223,6 +1244,7 @@ func shallowCloneMapping(m *proto.ProxyMapping) *proto.ProxyMapping {
}
}
// Authenticate verifies service credentials and issues a session token.
func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
if err := enforceAccountScope(ctx, req.GetAccountId()); err != nil {
return nil, err
@@ -1234,6 +1256,14 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
return nil, status.Errorf(codes.FailedPrecondition, "get service from store: %v", err)
}
switch req.GetRequest().(type) {
case *proto.AuthenticateRequest_Pin, *proto.AuthenticateRequest_Password:
key := credentialVerificationKey{accountID: credentialAccountID(service.AccountID), serviceID: credentialServiceID(service.ID)}
if err := s.credentialLimits.allow(key); err != nil {
return nil, err
}
}
authenticated, userId, method := s.authenticateRequest(ctx, req, service)
// Non-OIDC schemes (PIN/Password/Header) authenticate against per-service
@@ -1522,18 +1552,20 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
log.WithContext(ctx).Errorf("failed to get account services: %v", err)
return nil, status.Errorf(codes.FailedPrecondition, "get account services: %v", err)
}
var found bool
var matchedService *rpservice.Service
for _, service := range services {
if service.Domain == redirectURL.Hostname() {
found = true
matchedService = service
break
}
}
if !found {
if matchedService == nil {
log.WithContext(ctx).Debugf("OIDC redirect URL %q does not match any service domain", redirectURL.Hostname())
return nil, status.Errorf(codes.FailedPrecondition, "service not found in store")
}
useSessionCode := s.proxyManager.ClusterSupportsSessionCode(ctx, matchedService.ProxyCluster)
provider, err := oidc.NewProvider(ctx, s.oidcConfig.Issuer)
if err != nil {
log.WithContext(ctx).Errorf("failed to create OIDC provider: %v", err)
@@ -1553,15 +1585,18 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
return nil, status.Errorf(codes.Internal, "generate nonce: %v", err)
}
nonceB64 := base64.URLEncoding.EncodeToString(nonce)
if useSessionCode {
nonceB64 = sessionCodeNoncePrefix + nonceB64
}
// Using an HMAC here to avoid redirection state being modified.
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
payload := redirectURL.String() + "|" + nonceB64
hmacSum := s.generateHMAC(payload)
state := fmt.Sprintf("%s|%s|%s", base64.URLEncoding.EncodeToString([]byte(redirectURL.String())), nonceB64, hmacSum)
codeVerifier := oauth2.GenerateVerifier()
if err := s.pkceVerifierStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil {
if err := s.singleUseStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil {
log.WithContext(ctx).Errorf("failed to store PKCE verifier: %v", err)
return nil, status.Errorf(codes.Internal, "store PKCE verifier: %v", err)
}
@@ -1598,15 +1633,12 @@ func (s *ProxyServiceServer) generateHMAC(input string) string {
return hex.EncodeToString(mac.Sum(nil))
}
// ValidateState validates the state parameter from an OAuth callback.
// Returns the original redirect URL if valid, or an error if invalid.
// The HMAC is verified before consuming the PKCE verifier to prevent
// an attacker from invalidating a legitimate user's auth flow.
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, err error) {
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
// ValidateState validates and consumes an OIDC state.
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, useSessionCode bool, err error) {
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
parts := strings.Split(state, "|")
if len(parts) != 3 {
return "", "", errors.New("invalid state format")
return "", "", false, errors.New("invalid state format")
}
encodedURL := parts[0]
@@ -1615,7 +1647,7 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
redirectURLBytes, err := base64.URLEncoding.DecodeString(encodedURL)
if err != nil {
return "", "", fmt.Errorf("invalid state encoding: %w", err)
return "", "", false, fmt.Errorf("invalid state encoding: %w", err)
}
redirectURL = string(redirectURLBytes)
@@ -1623,16 +1655,17 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
expectedHMAC := s.generateHMAC(payload)
if !hmac.Equal([]byte(providedHMAC), []byte(expectedHMAC)) {
return "", "", errors.New("invalid state signature")
return "", "", false, errors.New("invalid state signature")
}
useSessionCode = strings.HasPrefix(nonce, sessionCodeNoncePrefix)
// Consume the PKCE verifier only after HMAC validation passes.
verifier, ok := s.pkceVerifierStore.LoadAndDelete(state)
verifier, ok := s.singleUseStore.LoadAndDelete(state)
if !ok {
return "", "", errors.New("no verifier for state")
return "", "", false, errors.New("no verifier for state")
}
return verifier, redirectURL, nil
return verifier, redirectURL, useSessionCode, nil
}
// Denied reasons reported to the proxy when access is refused because of the
@@ -1837,7 +1870,20 @@ func (s *ProxyServiceServer) getAccountServiceByDomain(ctx context.Context, acco
// ValidateSession validates a session token and checks if the user has access to the domain.
func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.ValidateSessionRequest) (*proto.ValidateSessionResponse, error) {
domain := req.GetDomain()
sessionToken := req.GetSessionToken()
sessionToken := req.GetSessionToken() //nolint:staticcheck
// A one-time code from the OIDC callback is redeemed here for the durable
// token, so the token never travels in a redirect URL. The redeemed token
// is returned to the proxy (mintedToken) to install as the session cookie.
mintedToken := ""
if code := req.GetSessionCode(); code != "" {
redeemed, found := s.singleUseStore.LoadAndDelete(singleUseCacheKey(sessionCodeCacheNamespace, code))
if !found {
return deniedSessionResponse("invalid or expired session code"), nil
}
sessionToken = redeemed
mintedToken = redeemed
}
if domain == "" || sessionToken == "" {
return deniedSessionResponse("missing domain or session_token"), nil
@@ -1910,6 +1956,7 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
UserEmail: user.Email,
PeerGroupIds: groupIDs,
PeerGroupNames: groupNames,
SessionToken: mintedToken,
}, nil
}
@@ -0,0 +1,93 @@
package grpc
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/shared/management/proto"
)
const (
versionTestProxyID = "proxy-a"
versionTestCluster = "cluster.example.com"
versionTestVersion = "0.60.0"
)
// hangupStream cancels its context on the first Send, emulating a proxy that
// disconnects right after receiving the initial snapshot. The legacy stream
// carries no proxy-to-management messages, so this is the only way for
// GetMappingUpdate to return.
type hangupStream struct {
recordingStream
ctx context.Context
cancel context.CancelFunc
}
func (s *hangupStream) Send(m *proto.GetMappingUpdateResponse) error {
s.cancel()
return s.recordingStream.Send(m)
}
func (s *hangupStream) Context() context.Context { return s.ctx }
// newVersionTestServer wires a server whose proxy manager only accepts a
// Connect carrying versionTestVersion, so a dropped or mangled version fails
// the test as an unexpected call.
func newVersionTestServer(t *testing.T) *ProxyServiceServer {
t.Helper()
ctrl := gomock.NewController(t)
svcMgr := rpservice.NewMockManager(ctrl)
svcMgr.EXPECT().GetGlobalServices(gomock.Any()).Return(nil, nil)
proxyMgr := proxy.NewMockManager(ctrl)
proxyMgr.EXPECT().
Connect(gomock.Any(), versionTestProxyID, gomock.Any(), versionTestCluster, gomock.Any(), versionTestVersion, gomock.Any(), gomock.Any()).
Return(&proxy.Proxy{ID: versionTestProxyID, Version: versionTestVersion}, nil)
proxyMgr.EXPECT().Disconnect(gomock.Any(), versionTestProxyID, gomock.Any()).Return(nil)
s := newSnapshotTestServer(t, 10)
s.serviceManager = svcMgr
s.proxyManager = proxyMgr
return s
}
func TestSyncMappings_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
// The init carries the version, the ack acknowledges the empty snapshot,
// and the exhausted fake stream then ends the RPC.
stream := &syncRecordingStream{
recvMsgs: []*proto.SyncMappingsRequest{
{Msg: &proto.SyncMappingsRequest_Init{Init: &proto.SyncMappingsInit{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}}},
ackMsg(),
},
}
err := s.SyncMappings(stream)
require.ErrorContains(t, err, "no more recv messages")
}
func TestGetMappingUpdate_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
stream := &hangupStream{ctx: ctx, cancel: cancel}
err := s.GetMappingUpdate(&proto.GetMappingUpdateRequest{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}, stream)
require.ErrorIs(t, err, context.Canceled)
}
@@ -0,0 +1,101 @@
package grpc
import (
"sync"
"time"
"golang.org/x/time/rate"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/types/known/durationpb"
)
const (
credentialVerificationInterval = 6 * time.Second
credentialVerificationBurst = 5
credentialVerificationMaxServices = 4096
credentialVerificationIdleTimeout = 15 * time.Minute
credentialVerificationCleanupInterval = time.Minute
)
type credentialAccountID string
type credentialServiceID string
type credentialVerificationKey struct {
accountID credentialAccountID
serviceID credentialServiceID
}
type credentialVerificationBudget struct {
limiter *rate.Limiter
lastUsed time.Time
}
// The zero value is ready to use. Budgets are local to this Management process;
// proxy replicas reaching this process share a service's verification budget.
type credentialVerificationLimiter struct {
mu sync.Mutex
now func() time.Time
services map[credentialVerificationKey]*credentialVerificationBudget
nextCleanup time.Time
closed bool
}
func (l *credentialVerificationLimiter) allow(key credentialVerificationKey) error {
l.mu.Lock()
defer l.mu.Unlock()
if l.closed {
return status.Error(codes.Unavailable, "credential verification is closed")
}
now := time.Now()
if l.now != nil {
now = l.now()
}
l.cleanup(now)
budget := l.services[key]
if budget == nil {
if len(l.services) >= credentialVerificationMaxServices {
return credentialVerificationThrottled(credentialVerificationCleanupInterval)
}
if l.services == nil {
l.services = make(map[credentialVerificationKey]*credentialVerificationBudget)
}
budget = &credentialVerificationBudget{limiter: rate.NewLimiter(rate.Every(credentialVerificationInterval), credentialVerificationBurst)}
l.services[key] = budget
}
budget.lastUsed = now
if budget.limiter.AllowN(now, 1) {
return nil
}
delay := max(time.Nanosecond, time.Duration((1-budget.limiter.TokensAt(now))*float64(credentialVerificationInterval)))
return credentialVerificationThrottled(delay)
}
func (l *credentialVerificationLimiter) cleanup(now time.Time) {
if now.Before(l.nextCleanup) {
return
}
l.nextCleanup = now.Add(credentialVerificationCleanupInterval)
for key, budget := range l.services {
if now.Sub(budget.lastUsed) >= credentialVerificationIdleTimeout {
delete(l.services, key)
}
}
}
func (l *credentialVerificationLimiter) close() {
l.mu.Lock()
defer l.mu.Unlock()
l.closed = true
l.services = nil
}
func credentialVerificationThrottled(delay time.Duration) error {
s := status.New(codes.ResourceExhausted, "too many credential verification attempts")
withRetry, err := s.WithDetails(&errdetails.RetryInfo{RetryDelay: durationpb.New(delay)})
if err != nil {
return s.Err()
}
return withRetry.Err()
}
@@ -0,0 +1,79 @@
package grpc
import (
"strconv"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
func TestCredentialVerificationRefillAndIsolation(t *testing.T) {
now := time.Now()
l := credentialVerificationLimiter{now: func() time.Time { return now }}
key := credentialVerificationKey{accountID: "account", serviceID: "service"}
for range credentialVerificationBurst {
require.NoError(t, l.allow(key))
}
err := l.allow(key)
require.Equal(t, codes.ResourceExhausted, status.Code(err), "the burst must be bounded")
now = now.Add(3 * time.Second)
err = l.allow(key)
require.Equal(t, codes.ResourceExhausted, status.Code(err), "a partially refilled token must not permit a check")
details := status.Convert(err).Details()
require.Len(t, details, 1, "throttling must provide RetryInfo")
retry, ok := details[0].(*errdetails.RetryInfo)
require.True(t, ok, "retry details must use the standard message")
assert.Equal(t, 3*time.Second, retry.RetryDelay.AsDuration(), "retry hint must reflect time until the next check")
now = now.Add(3 * time.Second)
require.NoError(t, l.allow(key))
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "only one check must refill every six seconds")
require.NoError(t, l.allow(credentialVerificationKey{accountID: "other-account", serviceID: key.serviceID}))
require.NoError(t, l.allow(credentialVerificationKey{accountID: key.accountID, serviceID: "other-service"}))
}
func TestCredentialVerificationCapacityAndExpiry(t *testing.T) {
now := time.Now()
l := credentialVerificationLimiter{now: func() time.Time { return now }}
for i := range credentialVerificationMaxServices {
require.NoError(t, l.allow(credentialVerificationKey{accountID: "account", serviceID: credentialServiceID(strconv.Itoa(i))}))
}
key := credentialVerificationKey{accountID: "account", serviceID: "new-service"}
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "capacity exhaustion must deny new checks")
now = now.Add(credentialVerificationIdleTimeout)
for range credentialVerificationBurst {
require.NoError(t, l.allow(key))
}
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "expiry must retain the normal burst bound")
}
func TestCredentialVerificationConcurrentChecksAndClose(t *testing.T) {
var l credentialVerificationLimiter
key := credentialVerificationKey{accountID: "account", serviceID: "service"}
var admitted atomic.Int32
var wg sync.WaitGroup
for range 100 {
wg.Go(func() {
if err := l.allow(key); err == nil {
admitted.Add(1)
} else {
assert.Equal(t, codes.ResourceExhausted, status.Code(err), "excess checks must be throttled")
}
})
}
wg.Wait()
assert.EqualValues(t, credentialVerificationBurst, admitted.Load(), "concurrent checks must share the burst")
for range 10 {
wg.Go(l.close)
wg.Go(func() { assert.Error(t, l.allow(key)) })
}
wg.Wait()
assert.Empty(t, l.services, "closing must release retained budgets")
assert.Equal(t, codes.Unavailable, status.Code(l.allow(key)), "checks after close must fail closed")
}
@@ -0,0 +1,18 @@
# Reverse proxy credential verification
The `ProxyService.Authenticate` RPC limits PIN and password checks before
verifying their Argon2 hashes. Both methods share one budget per account and
service: a burst of five checks, replenishing one check every six seconds
(ten per minute). Successful and failed checks consume the budget. Account
scope and service lookup run before the limiter.
Excess checks receive gRPC `ResourceExhausted` with a standard `RetryInfo` delay.
Updated proxies translate it to HTTP 429 and `Retry-After`. Older proxies show
an authentication-service error but cannot bypass the Management limit.
Budgets are held in memory per Management process and reset on restart. Proxy
replicas reaching the same Management process share its budgets. Multiple
Management processes have independent budgets; this is not a cluster-wide
limit. At most 4,096 service budgets are retained, with idle entries expiring
after fifteen minutes. Capacity exhaustion denies new checks until entries
expire. Closing the server releases the retained state.
@@ -0,0 +1,131 @@
package grpc_test
import (
"context"
"net"
"net/netip"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/peer"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
servicemanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service/manager"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/proto"
)
func credentialServer(t *testing.T) (*nbgrpc.ProxyServiceServer, context.Context, grpc.UnaryServerInterceptor) {
t.Helper()
ctx := context.Background()
s, err := store.NewStore(ctx, types.SqliteStoreEngine, t.TempDir(), nil, false)
require.NoError(t, err)
t.Cleanup(func() { assert.NoError(t, s.Close(ctx)) })
require.NoError(t, s.SaveAccount(ctx, &types.Account{Id: "account"}))
keys, err := sessionkey.GenerateKeyPair()
require.NoError(t, err)
for _, id := range []string{"service", "other-service"} {
svc := &service.Service{
ID: id, AccountID: "account", Name: id, Domain: id + ".example.com",
Enabled: true, SessionPrivateKey: keys.PrivateKey, SessionPublicKey: keys.PublicKey,
Auth: service.AuthConfig{
PinAuth: &service.PINAuthConfig{Enabled: true, Pin: "842716"},
PasswordAuth: &service.PasswordAuthConfig{Enabled: true, Password: "test-password"},
},
}
require.NoError(t, svc.Auth.HashSecrets())
require.NoError(t, s.CreateService(ctx, svc))
}
account := "account"
token, err := types.CreateNewProxyAccessToken("test proxy", time.Hour, &account, "admin")
require.NoError(t, err)
require.NoError(t, s.SaveProxyAccessToken(ctx, &token.ProxyAccessToken))
ctx = metadata.NewIncomingContext(ctx, metadata.Pairs("authorization", "Bearer "+string(token.PlainToken)))
ctx = peer.NewContext(ctx, &peer.Peer{Addr: net.TCPAddrFromAddrPort(netip.MustParseAddrPort("192.0.2.1:443"))})
server := nbgrpc.NewProxyServiceServer(nil, nil, nil, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
t.Cleanup(server.Close)
server.SetServiceManager(servicemanager.NewManager(s, nil, nil, nil, nil, nil))
interceptor, _, closeInterceptor := nbgrpc.NewProxyAuthInterceptors(s)
t.Cleanup(closeInterceptor)
return server, ctx, interceptor
}
func TestAuthenticateCredentialRateLimit(t *testing.T) {
server, ctx, interceptor := credentialServer(t)
authenticate := func(req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
response, err := interceptor(ctx, req, &grpc.UnaryServerInfo{FullMethod: "/management.ProxyService/Authenticate"}, func(ctx context.Context, req any) (any, error) {
return server.Authenticate(ctx, req.(*proto.AuthenticateRequest))
})
if err != nil {
return nil, err
}
return response.(*proto.AuthenticateResponse), nil
}
for i := range 5 {
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service"}
if i%2 == 0 {
req.Request = &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}
} else {
req.Request = &proto.AuthenticateRequest_Password{Password: &proto.PasswordRequest{Password: "wrong-password"}}
}
resp, err := authenticate(req)
require.NoError(t, err)
assert.False(t, resp.GetSuccess(), "incorrect PINs and passwords must be denied")
assert.Empty(t, resp.GetSessionToken(), "incorrect credentials must not issue a token")
}
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "842716"}}}
resp, err := authenticate(req)
assert.Nil(t, resp, "a throttled verification must not return a session")
require.Equal(t, codes.ResourceExhausted, status.Code(err), "PIN and password checks must share a service budget even with a valid proxy token")
details := status.Convert(err).Details()
require.Len(t, details, 1, "throttled responses must include a retry hint")
retry, ok := details[0].(*errdetails.RetryInfo)
require.True(t, ok, "the hint must use the standard RetryInfo message")
assert.Positive(t, retry.RetryDelay.AsDuration(), "the retry delay must be positive")
assert.LessOrEqual(t, retry.RetryDelay.AsDuration(), 6*time.Second, "the service must replenish one verification every six seconds")
req.AccountId = "another-account"
_, err = authenticate(req)
assert.Equal(t, codes.PermissionDenied, status.Code(err), "account scope must still be enforced before throttling")
req.AccountId = "account"
req.Id = "other-service"
resp, err = authenticate(req)
require.NoError(t, err)
assert.True(t, resp.GetSuccess(), "one service's throttle must not block another service")
assert.NotEmpty(t, resp.GetSessionToken(), "valid credentials on another service must issue a session")
}
func TestAuthenticateCredentialConcurrentLimit(t *testing.T) {
server, _, _ := credentialServer(t)
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}}
var checked, throttled atomic.Int32
var wg sync.WaitGroup
for range 20 {
wg.Go(func() {
resp, err := server.Authenticate(context.Background(), req)
switch status.Code(err) {
case codes.OK:
checked.Add(1)
assert.False(t, resp.GetSuccess(), "incorrect credentials must be denied")
case codes.ResourceExhausted:
throttled.Add(1)
default:
assert.NoError(t, err)
}
})
}
wg.Wait()
assert.EqualValues(t, 5, checked.Load(), "only the burst budget may reach concurrent credential verification")
assert.EqualValues(t, 15, throttled.Load(), "excess concurrent checks must be throttled")
}
+43 -19
View File
@@ -129,11 +129,11 @@ func drainEmpty(ch chan *proto.GetMappingUpdateResponse) bool {
func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -186,11 +186,11 @@ func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) {
func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -220,11 +220,11 @@ func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
func TestSendServiceUpdate_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -272,13 +272,13 @@ func generateState(s *ProxyServiceServer, redirectURL string) string {
func TestOAuthState_NeverTheSame(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
redirectURL := "https://app.example.com/callback"
@@ -300,20 +300,20 @@ func TestOAuthState_NeverTheSame(t *testing.T) {
func TestValidateState_RejectsOldTwoPartFormat(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
// Old format had only 2 parts: base64(url)|hmac
err := s.pkceVerifierStore.Store("base64url|hmac", "test", 10*time.Minute)
err := s.singleUseStore.Store("base64url|hmac", "test", 10*time.Minute)
require.NoError(t, err)
_, _, err = s.ValidateState("base64url|hmac")
_, _, _, err = s.ValidateState("base64url|hmac")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid state format")
}
@@ -372,24 +372,48 @@ func TestEnforceAccountScope_AllowsNoTokenInContext(t *testing.T) {
func TestValidateState_RejectsInvalidHMAC(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
// Store with tampered HMAC
err := s.pkceVerifierStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
err := s.singleUseStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
require.NoError(t, err)
_, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
_, _, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid state signature")
}
func TestSessionCodeCannotConsumeOIDCState(t *testing.T) {
const verifier = "pkce-verifier"
store := NewSingleUseStore(context.Background(), testCacheStore(t))
server := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
singleUseStore: store,
}
state := generateState(server, "https://service.example.com/callback")
require.NoError(t, store.Store(state, verifier, time.Minute))
response, err := server.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
SessionCode: state,
})
require.NoError(t, err)
assert.False(t, response.GetValid())
gotVerifier, _, _, err := server.ValidateState(state)
require.NoError(t, err)
assert.Equal(t, verifier, gotVerifier)
}
func TestSendServiceUpdateToCluster_FiltersOnCapability(t *testing.T) {
tokenStore := NewOneTimeTokenStore(context.Background(), testCacheStore(t))
@@ -0,0 +1,67 @@
package grpc
import (
"context"
"crypto/rand"
"encoding/base64"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/store"
log "github.com/sirupsen/logrus"
nbcache "github.com/netbirdio/netbird/management/server/cache"
)
// SingleUseStore stores short-lived values that can be retrieved only once.
type SingleUseStore struct {
cache nbcache.Store
ctx context.Context
}
// NewSingleUseStore creates a single-use value store over the shared cache.
func NewSingleUseStore(ctx context.Context, cacheStore nbcache.Store) *SingleUseStore {
return &SingleUseStore{
cache: cacheStore,
ctx: ctx,
}
}
// Store saves value under key with the given TTL, after which it is evicted.
func (s *SingleUseStore) Store(key, value string, ttl time.Duration) error {
if err := s.cache.Set(s.ctx, key, value, store.WithExpiration(ttl)); err != nil {
return fmt.Errorf("store single-use value: %w", err)
}
return nil
}
// Generate stores a value under a namespaced random key and returns the random key.
func (s *SingleUseStore) Generate(namespace, value string, ttl time.Duration) (string, error) {
buf := make([]byte, 32)
if _, err := rand.Read(buf); err != nil {
return "", fmt.Errorf("generate single-use key: %w", err)
}
key := base64.RawURLEncoding.EncodeToString(buf)
if err := s.Store(singleUseCacheKey(namespace, key), value, ttl); err != nil {
return "", err
}
return key, nil
}
func singleUseCacheKey(namespace, key string) string {
return namespace + ":" + key
}
// LoadAndDelete retrieves and removes the value for a key.
func (s *SingleUseStore) LoadAndDelete(key string) (string, bool) {
value, found, err := s.cache.GetDel(s.ctx, key)
if err != nil {
log.Warnf("failed to consume single-use value: %v", err)
return "", false
}
if !found {
return "", false
}
return value, true
}
@@ -6,7 +6,7 @@ import (
"time"
)
func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
func TestSingleUseStoreLoadAndDelete(t *testing.T) {
const (
state = "state"
verifier = "verifier"
@@ -14,7 +14,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
)
t.Run("exactly one concurrent caller consumes the verifier", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, time.Minute); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -50,7 +50,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("replayed state is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, time.Minute); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -64,7 +64,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("unknown state is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if got, found := store.LoadAndDelete("never-stored"); found {
t.Fatalf("unknown state should not resolve, got %q", got)
@@ -72,7 +72,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("expired verifier is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, 50*time.Millisecond); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -83,3 +83,40 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
}
})
}
func TestSingleUseStore_GenerateAndConsumeOnce(t *testing.T) {
const namespace = "test"
s := NewSingleUseStore(context.Background(), testCacheStore(t))
key, err := s.Generate(namespace, "the-value", time.Minute)
if err != nil {
t.Fatalf("generate: %v", err)
}
if key == "" || key == "the-value" {
t.Fatalf("unexpected key %q", key)
}
value, found := s.LoadAndDelete(singleUseCacheKey(namespace, key))
if !found || value != "the-value" {
t.Fatalf("expected to load the stored value, got %q found=%v", value, found)
}
if _, found := s.LoadAndDelete(singleUseCacheKey(namespace, key)); found {
t.Fatal("value must be consumed on first LoadAndDelete")
}
}
func TestSingleUseStore_GenerateUniqueKeys(t *testing.T) {
s := NewSingleUseStore(context.Background(), testCacheStore(t))
a, err := s.Generate("test", "v", time.Minute)
if err != nil {
t.Fatalf("generate a: %v", err)
}
b, err := s.Generate("test", "v", time.Minute)
if err != nil {
t.Fatalf("generate b: %v", err)
}
if a == b {
t.Fatal("generated keys must be distinct")
}
}
@@ -40,9 +40,9 @@ func setupValidateSessionTest(t *testing.T) *validateSessionTestSetup {
proxyManager := &testValidateSessionProxyManager{}
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
proxyService := NewProxyServiceServer(nil, tokenStore, pkceStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil)
proxyService := NewProxyServiceServer(nil, tokenStore, singleUseStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil)
proxyService.SetServiceManager(serviceManager)
createTestProxies(t, ctx, testStore)
@@ -570,7 +570,7 @@ func (m *testValidateSessionServiceManager) DeleteAccountCluster(_ context.Conte
type testValidateSessionProxyManager struct{}
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
return nil, nil
}
@@ -634,6 +634,10 @@ func (m *testValidateSessionProxyManager) ClusterSupportsPrivate(_ context.Conte
return nil
}
func (m *testValidateSessionProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
return false
}
type testValidateSessionUsersManager struct {
store store.Store
}
@@ -662,3 +666,47 @@ func (m *testValidateSessionUsersManager) GetUserWithGroups(ctx context.Context,
}
return user, groups, nil
}
func TestValidateSession_RedeemsSessionCode(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "testProxyId")
require.NoError(t, err)
token := createSessionToken(t, proxy.SessionPrivateKey, "allowedUserId", "test-proxy.example.com")
code, ok := setup.proxyService.GenerateSessionCode(token)
require.True(t, ok)
require.NotEqual(t, token, code, "code must not be the token itself")
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: code,
})
require.NoError(t, err)
assert.True(t, resp.Valid, "redeemed code should authorize the user")
assert.Equal(t, "allowedUserId", resp.UserId)
assert.Equal(t, token, resp.GetSessionToken(), "response must carry the durable token for the cookie")
// Single-use: the same code must not redeem again.
resp2, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: code,
})
require.NoError(t, err)
assert.False(t, resp2.Valid, "a consumed code must be rejected")
assert.Empty(t, resp2.GetSessionToken())
}
func TestValidateSession_InvalidSessionCode(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: "does-not-exist",
})
require.NoError(t, err)
assert.False(t, resp.Valid)
assert.Empty(t, resp.GetSessionToken())
}