Merge branch 'main' into client-local-metrics

# Conflicts:
#	client/internal/debug/debug.go
This commit is contained in:
Viktor Liu
2026-07-31 21:55:23 +02:00
439 changed files with 32494 additions and 7767 deletions
+1
View File
@@ -351,6 +351,7 @@ func (a *Auth) setSystemInfoFlags(info *system.Info) {
a.config.BlockLANAccess,
a.config.BlockInbound,
a.config.DisableIPv6,
a.config.SyncMessageVersion,
a.config.EnableSSHRoot,
a.config.EnableSSHSFTP,
a.config.EnableSSHLocalPortForwarding,
+1 -1
View File
@@ -299,7 +299,7 @@ func (d *DeviceAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowIn
UseIDToken: d.providerConfig.UseIDToken,
}
err = isValidAccessToken(tokenInfo.GetTokenToUse(), d.providerConfig.Audience)
err = validateTokenAudience(tokenInfo.GetTokenToUse(), d.providerConfig.Audience)
if err != nil {
return TokenInfo{}, fmt.Errorf("validate access token failed with error: %v", err)
}
+6 -1
View File
@@ -306,7 +306,7 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo,
audience = p.providerConfig.ClientID
}
if err := isValidAccessToken(tokenInfo.GetTokenToUse(), audience); err != nil {
if err := validateTokenAudience(tokenInfo.GetTokenToUse(), audience); err != nil {
return TokenInfo{}, fmt.Errorf("authentication failed: invalid access token - %w", err)
}
@@ -320,6 +320,11 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo,
return tokenInfo, nil
}
// parseEmailFromIDToken extracts the email (or name) claim from an ID token
// without verifying its signature. The value is best-effort and used only as a
// UX convenience (login hint prefill and display); it never drives an
// authorization decision. The authoritative identity is established server-side
// from the signature-verified token.
func parseEmailFromIDToken(token string) (string, error) {
parts := strings.Split(token, ".")
if len(parts) < 2 {
+27 -32
View File
@@ -24,11 +24,7 @@ import (
)
const (
// Skew tolerates a small clock difference between the management
// server and this peer before treating a deadline as "in the past".
// Slightly above typical NTP drift; tight enough that the UI doesn't
// paint a stale expiry as if it were valid.
Skew = 30 * time.Second
maxPastHorizon = 30 * 24 * time.Hour
// maxDeadlineHorizon caps how far in the future an accepted deadline
// can sit. A timestamp beyond this is almost certainly a protocol
@@ -57,7 +53,7 @@ var (
ErrDeadlineTooFarFuture = errors.New("session deadline too far in the future")
// ErrDeadlineInPast is returned by Update when the supplied deadline
// is more than Skew in the past.
// is more than maxPastHorizon in the past.
ErrDeadlineInPast = errors.New("session deadline in the past")
)
@@ -66,15 +62,14 @@ var (
// for deadline change/clear, PublishEvent for the two warnings); tests pass
// a fake recorder so the same surface is observable without an engine.
//
// The watcher is the single owner of the deadline propagated to the
// recorder: every set, clear, sanity-check rejection and Close routes the
// value through SetSessionExpiresAt, so the SubscribeStatus snapshot the UI
// reads can never drift from the watcher's timer state. (SetSessionExpiresAt
// fans out its own state-change notification, so no separate notify is
// needed.) The recorder is server-scoped and outlives this engine-scoped
// watcher — without the Close-time clear a teardown (Down, or the Down+Up of
// a profile switch) would leave the next session showing the previous one's
// stale "expires in" value.
// While the watcher runs, it owns the deadline propagated to the recorder:
// every set, clear and sanity-check rejection routes the value through
// SetSessionExpiresAt, so the SubscribeStatus snapshot the UI reads can
// never drift from the watcher's timer state. (SetSessionExpiresAt fans
// out its own state-change notification, so no separate notify is needed.)
// The recorder is server-scoped and outlives this engine-scoped watcher;
// Close deliberately leaves the recorder value in place so transient engine
// restarts don't blank it — the client run loop clears it on real teardown.
//
// PublishEvent's signature mirrors peer.Status.PublishEvent: the watcher
// composes the metadata internally so the wire format (MetaSession*) is
@@ -135,10 +130,13 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher {
// was disabled).
//
// Same-value updates are no-ops. A different non-zero value cancels any
// pending timer, resets the "already fired" guard, and arms a new one.
// pending timer, resets the "already fired" guards, and — when the
// deadline lies in the future — arms fresh warning timers. A deadline
// already in the past (within maxPastHorizon) is recorded as-is with no
// timers: the session has expired and consumers render it that way.
//
// Returns one of the sentinel Err* values when the deadline fails the
// sanity checks (pre-epoch, far future, or in the past beyond Skew).
// sanity checks (pre-epoch, far future, or past beyond maxPastHorizon).
// In every error case the watcher first clears its state so it stays
// consistent with what the caller will push into its other sinks (e.g.
// applySessionDeadline forces a zero deadline into the status recorder
@@ -163,7 +161,7 @@ func (w *Watcher) Update(deadline time.Time) error {
case deadline.After(now.Add(maxDeadlineHorizon)):
w.clearLocked()
return fmt.Errorf("%w: %v", ErrDeadlineTooFarFuture, deadline)
case deadline.Before(now.Add(-Skew)):
case deadline.Before(now.Add(-maxPastHorizon)):
w.clearLocked()
return fmt.Errorf("%w: %v (now=%v)", ErrDeadlineInPast, deadline, now)
}
@@ -183,7 +181,9 @@ func (w *Watcher) Update(deadline time.Time) error {
w.finalFiredAt = time.Time{}
w.dismissedAt = time.Time{}
w.armTimerLocked(deadline)
if deadline.After(now) {
w.armTimerLocked(deadline)
}
recorder := w.recorder
w.mu.Unlock()
if recorder != nil {
@@ -227,30 +227,25 @@ func (w *Watcher) Dismiss() {
log.Infof("auth session final-warning dismissed for deadline %s", w.current.Format(time.RFC3339))
}
// Close stops any pending timer and drops the deadline on the status
// recorder. Update calls after Close are ignored. Clearing the recorder
// here is what keeps a teardown (Down, or the Down+Up of a profile switch)
// from leaving the next session showing this one's stale "expires in"
// value — the recorder is server-scoped and outlives this engine-scoped
// watcher, so nothing else drops the anchor on teardown.
// Close stops any pending timer. Update calls after Close are ignored.
// The recorder keeps its deadline: the watcher is engine-scoped and closes
// on every engine restart (network change, sleep/wake, stream errors)
// while the SSO deadline stays valid across those, so clearing here would
// blank the UI's "expires in" row on every transient reconnect. The
// client run loop clears the server-scoped recorder when it exits for
// real (Down, profile switch, permanent login failure).
func (w *Watcher) Close() {
w.mu.Lock()
defer w.mu.Unlock()
if w.closed {
w.mu.Unlock()
return
}
w.closed = true
w.stopTimerLocked()
hadDeadline := !w.current.IsZero()
w.current = time.Time{}
w.firedAt = time.Time{}
w.finalFiredAt = time.Time{}
w.dismissedAt = time.Time{}
recorder := w.recorder
w.mu.Unlock()
if recorder != nil && hadDeadline {
recorder.SetSessionExpiresAt(time.Time{})
}
}
// clearLocked drops the tracked deadline and notifies the recorder so
@@ -224,11 +224,13 @@ func TestNewDeadlineCancelsPriorTimer(t *testing.T) {
func TestRefreshAfterFireArmsNewWarning(t *testing.T) {
r := &fakeRecorder{}
lead := 30 * time.Millisecond
lead := 150 * time.Millisecond
w := newWatcher(lead, r)
defer w.Close()
first := time.Now().Add(50 * time.Millisecond)
// Warning fires ~20ms in; the deadline itself stays 150ms away so the
// replacement below lands well before it.
first := time.Now().Add(170 * time.Millisecond)
_ = w.Update(first)
// Wait for stateChange + warning of the first cycle.
@@ -306,7 +308,29 @@ func TestUpdateRejectsTooFarFuture(t *testing.T) {
}
}
func TestUpdateInPastClearsDeadline(t *testing.T) {
func TestUpdateRecentPastRecordedAsExpired(t *testing.T) {
r := &fakeRecorder{}
w := newWatcher(50*time.Millisecond, r)
defer w.Close()
d := time.Now().Add(-1 * time.Hour)
if err := w.Update(d); err != nil {
t.Fatalf("recent-past Update should succeed, got %v", err)
}
if !w.Deadline().Equal(d) {
t.Fatalf("expected deadline to be recorded, got %v want %v", w.Deadline(), d)
}
if got := r.deadline(); !got.Equal(d) {
t.Fatalf("recorder deadline = %v, want %v", got, d)
}
time.Sleep(80 * time.Millisecond)
if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 {
t.Fatalf("no warning events may fire for an already-past deadline, got %+v", r.snapshot())
}
}
func TestUpdateAncientPastRejected(t *testing.T) {
r := &fakeRecorder{}
w := newWatcher(50*time.Millisecond, r)
defer w.Close()
@@ -318,12 +342,12 @@ func TestUpdateInPastClearsDeadline(t *testing.T) {
// Drain the stateChange from the seed.
waitForEvents(t, r, 1)
err := w.Update(time.Now().Add(-1 * time.Hour))
err := w.Update(time.Now().Add(-31 * 24 * time.Hour))
if !errors.Is(err, ErrDeadlineInPast) {
t.Fatalf("want ErrDeadlineInPast, got %v", err)
}
if !w.Deadline().IsZero() {
t.Fatalf("in-past update must clear the deadline, got %v", w.Deadline())
t.Fatalf("rejected ancient-past update must clear the deadline, got %v", w.Deadline())
}
events := waitForEvents(t, r, 2)
if events[1].kind != stateChange {
@@ -331,39 +355,25 @@ func TestUpdateInPastClearsDeadline(t *testing.T) {
}
}
func TestUpdateWithinSkewAccepted(t *testing.T) {
r := &fakeRecorder{}
w := newWatcher(50*time.Millisecond, r)
defer w.Close()
// 5 seconds in the past is within the 30s Skew tolerance — accept it.
d := time.Now().Add(-5 * time.Second)
if err := w.Update(d); err != nil {
t.Fatalf("within-skew Update should succeed, got %v", err)
}
if !w.Deadline().Equal(d) {
t.Fatalf("expected deadline to be applied, got %v want %v", w.Deadline(), d)
}
}
func TestCloseSilencesUpdates(t *testing.T) {
r := &fakeRecorder{}
w := newWatcher(50*time.Millisecond, r)
w.Close()
_ = w.Update(time.Now().Add(time.Hour))
time.Sleep(20 * time.Millisecond)
if err := w.Update(time.Now().Add(time.Hour)); err != nil {
t.Fatalf("Update after Close: want nil, got %v", err)
}
if got := r.snapshot(); len(got) != 0 {
t.Fatalf("expected no events after Close, got %+v", got)
}
}
// TestCloseClearsRecorderDeadline pins the profile-switch fix: a watcher
// holding a live deadline must zero the recorder on Close so the next
// engine's watcher (and the UI reading the shared server-scoped recorder)
// doesn't start out showing the previous session's stale "expires in".
func TestCloseClearsRecorderDeadline(t *testing.T) {
// TestCloseKeepsRecorderDeadline pins the reconnect-flap fix: the watcher
// closes on every engine restart (network change, sleep/wake) while the
// SSO deadline stays valid across those, so Close must leave the
// server-scoped recorder's value in place. The client run loop clears the
// recorder when it exits for real.
func TestCloseKeepsRecorderDeadline(t *testing.T) {
r := &fakeRecorder{}
w := newWatcher(time.Hour, r)
@@ -377,8 +387,8 @@ func TestCloseClearsRecorderDeadline(t *testing.T) {
w.Close()
if got := r.deadline(); !got.IsZero() {
t.Fatalf("recorder deadline after Close = %v, want zero", got)
if got := r.deadline(); !got.Equal(d) {
t.Fatalf("recorder deadline after Close = %v, want %v", got, d)
}
}
+16 -4
View File
@@ -20,14 +20,26 @@ func randomBytesInHex(count int) (string, error) {
return hex.EncodeToString(buf), nil
}
// isValidAccessToken is a simple validation of the access token
func isValidAccessToken(token string, audience string) error {
// validateTokenAudience checks that the token is a well-formed JWT whose
// audience claim matches the expected audience.
//
// It does NOT verify the token's cryptographic signature and therefore must not
// be treated as an authenticity check. The token is obtained by the client
// directly from the IdP token endpoint over TLS, and its signature is verified
// server-side by the management server against the IdP's JWKS
// (see shared/auth/jwt/validator.go). This function is only a client-side
// sanity check that the returned token targets the expected audience.
func validateTokenAudience(token string, audience string) error {
if token == "" {
return fmt.Errorf("token received is empty")
}
encodedClaims := strings.Split(token, ".")[1]
claimsString, err := base64.RawURLEncoding.DecodeString(encodedClaims)
parts := strings.Split(token, ".")
if len(parts) != 3 {
return fmt.Errorf("token is not a well-formed JWT")
}
claimsString, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return err
}
+108
View File
@@ -0,0 +1,108 @@
package auth
import (
"encoding/base64"
"encoding/json"
"testing"
)
// makeJWT builds an unsigned JWT-shaped string (header.payload.signature) with
// the given claims payload. The signature part is arbitrary because
// validateTokenAudience intentionally does not verify it.
func makeJWT(t *testing.T, claims map[string]interface{}) string {
t.Helper()
header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","typ":"JWT"}`))
payloadBytes, err := json.Marshal(claims)
if err != nil {
t.Fatalf("marshal claims: %v", err)
}
payload := base64.RawURLEncoding.EncodeToString(payloadBytes)
return header + "." + payload + ".unverified-signature"
}
func TestValidateTokenAudience(t *testing.T) {
tests := []struct {
name string
token string
audience string
wantErr bool
}{
{
name: "empty token",
token: "",
audience: "netbird",
wantErr: true,
},
{
name: "not a JWT - no dots",
token: "notajwt",
audience: "netbird",
wantErr: true,
},
{
name: "not a JWT - two parts only",
token: "header.payload",
audience: "netbird",
wantErr: true,
},
{
name: "matching string audience",
token: makeJWT(t, map[string]interface{}{"aud": "netbird"}),
audience: "netbird",
wantErr: false,
},
{
name: "mismatching string audience",
token: makeJWT(t, map[string]interface{}{"aud": "other"}),
audience: "netbird",
wantErr: true,
},
{
name: "matching audience in array",
token: makeJWT(t, map[string]interface{}{"aud": []interface{}{"other", "netbird"}}),
audience: "netbird",
wantErr: false,
},
{
name: "mismatching audience array",
token: makeJWT(t, map[string]interface{}{"aud": []interface{}{"a", "b"}}),
audience: "netbird",
wantErr: true,
},
{
name: "missing audience claim",
token: makeJWT(t, map[string]interface{}{"sub": "user"}),
audience: "netbird",
wantErr: true,
},
{
name: "invalid base64 payload",
token: "header.!!!not-base64!!!.sig",
audience: "netbird",
wantErr: true,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
err := validateTokenAudience(tc.token, tc.audience)
if tc.wantErr && err == nil {
t.Fatalf("expected error, got nil")
}
if !tc.wantErr && err != nil {
t.Fatalf("expected no error, got %v", err)
}
})
}
}
// TestValidateTokenAudienceNoPanic guards the regression where a non-empty
// token without the JWT dot structure caused an index-out-of-range panic.
func TestValidateTokenAudienceNoPanic(t *testing.T) {
inputs := []string{"a", ".", "a.", "aaaa", "no-dots-here"}
for _, in := range inputs {
if err := validateTokenAudience(in, "netbird"); err == nil {
t.Fatalf("expected error for malformed token %q, got nil", in)
}
}
}
+50 -8
View File
@@ -34,6 +34,8 @@ const (
// - Handling connection establishment based on peer signaling
//
// The implementation is not thread-safe; it is protected by engine.syncMsgMux.
// The only exception is ActivatePeer, which is safe for concurrent use so the
// DNS warm-up path can call it without contending on the engine mutex.
type ConnMgr struct {
peerStore *peerstore.Store
statusRecorder *peer.Status
@@ -42,12 +44,26 @@ type ConnMgr struct {
rosenpassEnabled bool
lazyConnMgr *manager.Manager
// lazyConnMgrMu guards the lazyConnMgr pointer for readers outside the
// engine loop (ActivatePeer). Writers hold it in addition to
// engine.syncMsgMux; all other reads stay under engine.syncMsgMux only.
lazyConnMgrMu sync.RWMutex
// reconcileRoutedIPs re-applies a peer's routed allowed IPs after its lazy wake endpoint is
// (re)armed (Mode A at arm time). Injected by the engine; nil disables the reconcile.
reconcileRoutedIPs func(peerKey string) error
wg sync.WaitGroup
lazyCtx context.Context
lazyCtxCancel context.CancelFunc
}
// SetRoutedIPsReconciler injects the callback used to re-apply a peer's routed allowed IPs when
// its lazy wake endpoint is (re)armed. Must be called before the lazy manager starts.
func (e *ConnMgr) SetRoutedIPsReconciler(fn func(peerKey string) error) {
e.reconcileRoutedIPs = fn
}
func NewConnMgr(engineConfig *EngineConfig, statusRecorder *peer.Status, peerStore *peerstore.Store, iface lazyconn.WGIface) *ConnMgr {
e := &ConnMgr{
peerStore: peerStore,
@@ -238,12 +254,20 @@ func (e *ConnMgr) RemovePeerConn(peerKey string) {
conn.Log.Infof("removed peer from lazy conn manager")
}
// ActivatePeer wakes an idle lazy connection. Unlike the rest of ConnMgr it is
// safe for concurrent use: the lazy manager pointer is read under lazyConnMgrMu
// and the manager itself is internally synchronized, so callers outside the
// engine loop (DNS warm-up) do not need engine.syncMsgMux.
func (e *ConnMgr) ActivatePeer(ctx context.Context, conn *peer.Conn) {
if !e.isStartedWithLazyMgr() {
e.lazyConnMgrMu.RLock()
lazyConnMgr := e.lazyConnMgr
started := lazyConnMgr != nil && e.lazyCtxCancel != nil
e.lazyConnMgrMu.RUnlock()
if !started {
return
}
if found := e.lazyConnMgr.ActivatePeer(conn.GetKey()); found {
if found := lazyConnMgr.ActivatePeer(conn.GetKey()); found {
if err := conn.Open(ctx); err != nil {
conn.Log.Errorf("failed to open connection: %v", err)
}
@@ -268,16 +292,22 @@ func (e *ConnMgr) Close() {
e.lazyCtxCancel()
e.wg.Wait()
e.lazyConnMgrMu.Lock()
e.lazyConnMgr = nil
e.lazyConnMgrMu.Unlock()
}
func (e *ConnMgr) initLazyManager(engineCtx context.Context) {
cfg := manager.Config{
InactivityThreshold: inactivityThresholdEnv(),
ReconcileAllowedIPs: e.reconcileRoutedIPs,
}
e.lazyConnMgr = manager.NewManager(cfg, engineCtx, e.peerStore, e.iface)
e.lazyConnMgrMu.Lock()
e.lazyConnMgr = manager.NewManager(cfg, engineCtx, e.peerStore, e.iface)
e.lazyCtx, e.lazyCtxCancel = context.WithCancel(engineCtx)
e.lazyConnMgrMu.Unlock()
e.wg.Add(1)
go func() {
@@ -316,7 +346,10 @@ func (e *ConnMgr) closeManager(ctx context.Context) {
e.lazyCtxCancel()
e.wg.Wait()
e.lazyConnMgrMu.Lock()
e.lazyConnMgr = nil
e.lazyConnMgrMu.Unlock()
for _, peerID := range e.peerStore.PeersPubKey() {
e.peerStore.PeerConnOpen(ctx, peerID)
@@ -352,11 +385,20 @@ func inactivityThresholdEnv() *time.Duration {
return nil
}
parsedMinutes, err := strconv.Atoi(envValue)
if err != nil || parsedMinutes <= 0 {
return nil
// Documented format: a Go duration such as "30m" or "1h".
if d, err := time.ParseDuration(envValue); err == nil {
if d <= 0 {
return nil
}
return &d
}
d := time.Duration(parsedMinutes) * time.Minute
return &d
// Backwards compatibility: a bare integer used to be interpreted as minutes.
if parsedMinutes, err := strconv.Atoi(envValue); err == nil && parsedMinutes > 0 {
d := time.Duration(parsedMinutes) * time.Minute
return &d
}
log.Warnf("invalid %s value %q: expected a Go duration such as 30m or 1h", lazyconn.EnvInactivityThreshold, envValue)
return nil
}
+101
View File
@@ -1,10 +1,21 @@
package internal
import (
"context"
"net"
"net/netip"
"os"
"sync"
"testing"
"time"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/client/internal/lazyconn"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/peerstore"
"github.com/netbirdio/netbird/monotime"
)
func TestResolveLazyForce(t *testing.T) {
@@ -38,3 +49,93 @@ func TestResolveLazyForce(t *testing.T) {
})
}
}
type mockLazyWGIface struct{}
func (mockLazyWGIface) RemovePeer(string) error { return nil }
func (mockLazyWGIface) UpdatePeer(string, []netip.Prefix, time.Duration, *net.UDPAddr, *wgtypes.Key) error {
return nil
}
func (mockLazyWGIface) IsUserspaceBind() bool { return false }
func (mockLazyWGIface) Address() wgaddr.Address { return wgaddr.Address{} }
func (mockLazyWGIface) LastActivities() map[string]monotime.Time { return nil }
func (mockLazyWGIface) MTU() uint16 { return 1280 }
// TestConnMgr_ActivatePeerConcurrentWithLifecycle exercises ActivatePeer from
// non-engine goroutines (the DNS warm-up path) racing the manager lifecycle,
// which stays on the engine loop. Run with -race: it fails if ActivatePeer
// still requires engine.syncMsgMux for safety.
func TestConnMgr_ActivatePeerConcurrentWithLifecycle(t *testing.T) {
t.Setenv(lazyconn.EnvLazyConn, "on")
status := peer.NewRecorder("https://mgm")
store := peerstore.NewConnStore()
connMgr := NewConnMgr(&EngineConfig{}, status, store, mockLazyWGIface{})
conn := newTestPeerConn(t, "peerA")
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
connMgr.Start(ctx)
done := make(chan struct{})
var wg sync.WaitGroup
for range 4 {
wg.Add(1)
go func() {
defer wg.Done()
for {
select {
case <-done:
return
default:
connMgr.ActivatePeer(ctx, conn)
}
}
}()
}
// Let the activators spin against the started manager, then tear it down
// underneath them and let them spin against the stopped manager.
time.Sleep(100 * time.Millisecond)
connMgr.Close()
time.Sleep(50 * time.Millisecond)
close(done)
wg.Wait()
}
func TestInactivityThresholdEnv(t *testing.T) {
tests := []struct {
name string
val string
want *time.Duration
}{
{name: "unset", val: "", want: nil},
{name: "go duration minutes", val: "30m", want: durPtr(30 * time.Minute)},
{name: "go duration hours", val: "1h", want: durPtr(time.Hour)},
{name: "go duration seconds", val: "90s", want: durPtr(90 * time.Second)},
{name: "bare integer is minutes (backwards compat)", val: "5", want: durPtr(5 * time.Minute)},
{name: "zero duration", val: "0s", want: nil},
{name: "zero integer", val: "0", want: nil},
{name: "negative duration", val: "-5m", want: nil},
{name: "garbage", val: "abc", want: nil},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Setenv(lazyconn.EnvInactivityThreshold, tc.val)
got := inactivityThresholdEnv()
switch {
case tc.want == nil && got != nil:
t.Fatalf("want nil, got %v", *got)
case tc.want != nil && got == nil:
t.Fatalf("want %v, got nil", *tc.want)
case tc.want != nil && *got != *tc.want:
t.Fatalf("want %v, got %v", *tc.want, *got)
}
})
}
}
func durPtr(d time.Duration) *time.Duration { return &d }
+12 -3
View File
@@ -34,6 +34,7 @@ import (
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/internal/statemanager"
"github.com/netbirdio/netbird/client/internal/stdnet"
"github.com/netbirdio/netbird/client/internal/tunnelnotifier"
"github.com/netbirdio/netbird/client/internal/updater"
"github.com/netbirdio/netbird/client/internal/updater/installer"
nbnet "github.com/netbirdio/netbird/client/net"
@@ -136,10 +137,13 @@ func (c *ConnectClient) RunOniOS(
// Set GC percent to 5% to reduce memory usage as iOS only allows 50MB of memory for the extension.
debug.SetGCPercent(5)
notifier := tunnelnotifier.New(networkChangeListener, dnsManager)
defer notifier.Close()
mobileDependency := MobileDependency{
FileDescriptor: fileDescriptor,
NetworkChangeListener: networkChangeListener,
DnsManager: dnsManager,
NetworkChangeListener: notifier,
DnsManager: notifier,
StateFilePath: stateFilePath,
TempDir: cacheDir,
}
@@ -257,7 +261,10 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
log.Errorf("failed to clean up temporary installer file: %v", err)
}
defer c.statusRecorder.ClientStop()
defer func() {
c.statusRecorder.SetSessionExpiresAt(time.Time{})
c.statusRecorder.ClientStop()
}()
operation := func() error {
// if context cancelled we not start new backoff cycle
if c.ctx.Err() != nil {
@@ -618,6 +625,7 @@ func createEngineConfig(key wgtypes.Key, config *profilemanager.Config, peerConf
BlockLANAccess: config.BlockLANAccess,
BlockInbound: config.BlockInbound,
DisableIPv6: config.DisableIPv6,
SyncMessageVersion: config.SyncMessageVersion,
LazyConnection: lazyconn.ParseState(config.LazyConnection),
@@ -693,6 +701,7 @@ func loginToManagement(ctx context.Context, client mgm.Client, pubSSHKey []byte,
config.BlockLANAccess,
config.BlockInbound,
config.DisableIPv6,
config.SyncMessageVersion,
config.EnableSSHRoot,
config.EnableSSHSFTP,
config.EnableSSHLocalPortForwarding,
+15
View File
@@ -0,0 +1,15 @@
package daemonaddr
// DaemonRunsAsSelf reports whether the daemon listening at addr runs as this very
// user. That is what makes an unprivileged daemon authorize this process for the
// changes it otherwise restricts to root or an administrator, so a client can tell
// up front whether those controls are usable instead of letting a save fail.
//
// It is answered from the ownership of the socket or pipe the daemon created, so it
// costs no round trip and needs no cooperation from the daemon. Ownership that
// cannot be read is reported as false, including for a TCP address, so a caller
// reading this as "the daemon would allow it" fails closed. The daemon remains the
// only thing that authorizes anything: this only decides what a client offers.
func DaemonRunsAsSelf(addr string) bool {
return daemonRunsAsSelf(addr)
}
+40
View File
@@ -0,0 +1,40 @@
//go:build !windows
package daemonaddr
import (
"os"
"strings"
"syscall"
log "github.com/sirupsen/logrus"
)
// daemonRunsAsSelf compares the owner of the daemon's Unix socket with this
// process's uid. Root is not treated specially here: a root caller is privileged
// on its own merits, and a root-owned socket says nothing about the caller.
func daemonRunsAsSelf(addr string) bool {
path, ok := strings.CutPrefix(addr, "unix://")
if !ok {
return false
}
info, err := os.Stat(path)
if err != nil {
log.Debugf("stat daemon socket %s: %v", path, err)
return false
}
// Only a socket says anything about a daemon. A directory or a leftover
// regular file at that path is not one, and reading it as "the daemon runs as
// us" would offer controls the daemon then refuses.
if info.Mode()&os.ModeSocket == 0 {
return false
}
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return false
}
return stat.Uid == uint32(os.Getuid())
}
@@ -0,0 +1,62 @@
//go:build !windows
package daemonaddr
import (
"net"
"os"
"path/filepath"
"testing"
)
// A socket this user created means the daemon runs as this user, which is the
// rootless case where the daemon delegates its authority to its own identity.
func TestDaemonRunsAsSelf_OwnSocket(t *testing.T) {
path := filepath.Join(t.TempDir(), "netbird.sock")
ln, err := net.Listen("unix", path)
if err != nil {
t.Fatalf("listen: %v", err)
}
t.Cleanup(func() {
if err := ln.Close(); err != nil {
t.Logf("close listener: %v", err)
}
})
if !DaemonRunsAsSelf("unix://" + path) {
t.Error("a socket owned by this user must count as the daemon running as us")
}
}
// Everything that is not a readable socket of ours has to answer false, because
// the caller reads a true as "the daemon would authorize me".
func TestDaemonRunsAsSelf_FailsClosed(t *testing.T) {
dir := t.TempDir()
// A socket owned by another user, which is what a root-run daemon looks like
// to an unprivileged client. Only assertable when we are not root ourselves.
rootOwned := "unix:///var/run/netbird.sock"
if _, err := os.Stat("/var/run/netbird.sock"); err == nil && os.Getuid() != 0 {
if DaemonRunsAsSelf(rootOwned) {
t.Error("a socket owned by another user must not count as ours")
}
}
for name, addr := range map[string]string{
"missing socket": "unix://" + filepath.Join(dir, "absent.sock"),
"tcp address": "tcp://127.0.0.1:41731",
"named pipe": "npipe://netbird",
"empty": "",
"no scheme": filepath.Join(dir, "absent.sock"),
"directory": "unix://" + dir,
"unknown scheme": "http://localhost:8080",
"scheme only": "unix://",
"relative socket": "unix://netbird.sock",
} {
t.Run(name, func(t *testing.T) {
if DaemonRunsAsSelf(addr) {
t.Errorf("%q must not count as a daemon running as us", addr)
}
})
}
}
@@ -0,0 +1,42 @@
//go:build windows
package daemonaddr
import (
"context"
"strings"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/ipcauth"
)
// daemonRunsAsSelf reads the owner of the daemon's pipe. A daemon running as the
// service account owns its pipe as LocalSystem, and an elevated one as
// BUILTIN\Administrators, so only a daemon the user started themselves matches.
func daemonRunsAsSelf(addr string) bool {
name, ok := strings.CutPrefix(addr, pipeScheme)
if !ok {
return false
}
for _, path := range PipePaths(name) {
// Bounded: this runs on the UI's path for deciding which controls to
// offer, so a pipe that does not answer promptly must not stall it. A
// timeout leaves the caller unprivileged, which only disables controls.
ctx, cancel := context.WithTimeout(context.Background(), probeTimeout)
conn, err := dialPipe(ctx, path)
cancel()
if err != nil {
continue
}
owned := ipcauth.PipeOwnedBySelf(conn)
if cerr := conn.Close(); cerr != nil {
log.Debugf("close daemon pipe %s after ownership check: %v", path, cerr)
}
return owned
}
return false
}
+103
View File
@@ -0,0 +1,103 @@
package daemonaddr
import (
"context"
"net"
"runtime"
"strings"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
)
const (
// WindowsPipeAddr is the default daemon address on Windows. A named pipe
// carries the connecting process's token, which loopback TCP does not, so
// it is the only Windows transport on which the daemon can tell who is
// calling it.
WindowsPipeAddr = "npipe://netbird"
// legacyWindowsAddr is the loopback-TCP address the Windows daemon used
// before named-pipe support.
legacyWindowsAddr = "tcp://127.0.0.1:41731"
pipeScheme = "npipe://"
// protectedPrefix is the NPFS namespace in which only LocalSystem and
// members of BUILTIN\Administrators may create a pipe. A daemon running as
// the service account creates its pipe there so that an unprivileged process
// cannot pre-create the name, which would keep the daemon from starting and
// leave callers talking to the squatter. Opening such a pipe needs no
// privilege, so unprivileged clients still reach the daemon.
protectedPrefix = `ProtectedPrefix\Administrators\`
)
// DialTarget returns the gRPC dial target and transport options for a daemon
// address. The npipe scheme needs a context dialer because gRPC has no
// named-pipe resolver; unix and tcp are handled by gRPC itself.
func DialTarget(addr string) (string, []grpc.DialOption) {
opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())}
if name, ok := strings.CutPrefix(addr, pipeScheme); ok {
paths := PipePaths(name)
opts = append(opts, grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) {
return dialPipePaths(ctx, paths)
}))
return "passthrough:///netbird-daemon-pipe", opts
}
return strings.TrimPrefix(addr, "tcp://"), opts
}
// PipePath maps an npipe address name ("netbird", from "npipe://netbird") to a
// Windows named-pipe path (\\.\pipe\netbird). A fully qualified path is left as
// is.
func PipePath(name string) string {
if strings.HasPrefix(name, `\\`) {
return name
}
return `\\.\pipe\` + name
}
// PipePaths returns the paths a daemon control pipe may live at for an npipe
// address name, in the order both sides must try them: the protected name first,
// then the plain one.
//
// The daemon serves the first it can create, which is the protected name when it
// runs as the service account and the plain one when it runs as an ordinary user,
// as it does in netstack mode. Clients therefore have to try both, and because a
// client cannot tell from the name alone who created the pipe, the plain name is
// only usable once the server's identity has been checked: see
// verifyPipeServer.
//
// A fully qualified path is what the operator asked for and is used as is.
func PipePaths(name string) []string {
if strings.HasPrefix(name, `\\`) {
return []string{name}
}
return []string{PipePath(protectedPrefix + name), PipePath(name)}
}
// IsProtectedPipePath reports whether a pipe path is in the namespace only an
// administrator or LocalSystem can create in, which is what lets a client trust
// such a pipe from its name alone.
func IsProtectedPipePath(path string) bool {
return strings.HasPrefix(path, `\\.\pipe\`+protectedPrefix)
}
// MigrateLegacy upgrades the pre-named-pipe Windows daemon address to the named
// pipe, reporting whether it rewrote the address. Existing installs persist the
// daemon address, so without this an upgraded daemon would keep listening on
// loopback TCP, where callers carry no identity and privileged operations would
// have to be refused for everyone. Only the exact legacy default is rewritten:
// a deliberately chosen custom address is left alone.
func MigrateLegacy(addr string) (string, bool) {
return migrateLegacyForOS(runtime.GOOS, addr)
}
func migrateLegacyForOS(goos, addr string) (string, bool) {
if goos == "windows" && addr == legacyWindowsAddr {
return WindowsPipeAddr, true
}
return addr, false
}
+15
View File
@@ -0,0 +1,15 @@
//go:build !windows
package daemonaddr
import (
"context"
"fmt"
"net"
)
// dialPipePaths is Windows-only: no other platform serves the daemon on a named
// pipe.
func dialPipePaths(context.Context, []string) (net.Conn, error) {
return nil, fmt.Errorf("named pipes are only supported on Windows")
}
+30
View File
@@ -0,0 +1,30 @@
package daemonaddr
import (
"slices"
"testing"
)
// The protected name must be tried before the plain one on both sides: it is the
// one an unprivileged process cannot create, so preferring it is what keeps a
// squatter from owning the name the service daemon would otherwise use.
func TestPipePaths_PrefersTheProtectedName(t *testing.T) {
got := PipePaths("netbird")
want := []string{
`\\.\pipe\ProtectedPrefix\Administrators\netbird`,
`\\.\pipe\netbird`,
}
if !slices.Equal(got, want) {
t.Errorf("PipePaths = %q, want %q", got, want)
}
}
// An operator who passes a full path chose exactly one pipe, so neither side may
// look anywhere else.
func TestPipePaths_QualifiedPathIsUsedAsIs(t *testing.T) {
path := `\\.\pipe\custom-netbird`
got := PipePaths(path)
if !slices.Equal(got, []string{path}) {
t.Errorf("PipePaths = %q, want just %q", got, path)
}
}
@@ -0,0 +1,59 @@
//go:build windows
package daemonaddr
import (
"context"
"errors"
"fmt"
"net"
"github.com/Microsoft/go-winio"
log "github.com/sirupsen/logrus"
"golang.org/x/sys/windows"
"github.com/netbirdio/netbird/client/internal/ipcauth"
)
// dialPipePaths connects to the first path that answers with a pipe server this
// client may trust, and returns the last error when none does.
func dialPipePaths(ctx context.Context, paths []string) (net.Conn, error) {
var lastErr error
for _, path := range paths {
conn, err := dialPipe(ctx, path)
if err != nil {
log.Debugf("dial daemon pipe %s: %v", path, err)
lastErr = err
continue
}
// A pipe in the protected namespace could only have been created by an
// administrator or LocalSystem, so its name is the guarantee. Any other
// name has to be checked, because any local user can create one.
if !IsProtectedPipePath(path) {
if err := ipcauth.PipeServerTrusted(conn); err != nil {
if closeErr := conn.Close(); closeErr != nil {
log.Debugf("close untrusted pipe %s: %v", path, closeErr)
}
lastErr = fmt.Errorf("%s: %w", path, err)
continue
}
}
return conn, nil
}
if lastErr == nil {
lastErr = errors.New("no daemon pipe to connect to")
}
return nil, lastErr
}
// dialPipe connects to the daemon control pipe at SECURITY_IDENTIFICATION.
// winio's plain DialPipe connects at SECURITY_ANONYMOUS, under which the daemon
// cannot read the caller's token at all. Identification lets the daemon read the
// caller's SID and groups without granting it the ability to act as the caller.
func dialPipe(ctx context.Context, path string) (net.Conn, error) {
access := uint32(windows.GENERIC_READ | windows.GENERIC_WRITE)
return winio.DialPipeAccessImpLevel(ctx, path, access, winio.PipeImpLevelIdentification)
}
@@ -0,0 +1,9 @@
//go:build !windows
package daemonaddr
// ResolveDaemonAddr is a no-op off Windows, where there is no named-pipe
// default to fall back from.
func ResolveDaemonAddr(addr string) string {
return addr
}
@@ -0,0 +1,82 @@
//go:build windows
package daemonaddr
import (
"net"
"strings"
"time"
"github.com/Microsoft/go-winio"
log "github.com/sirupsen/logrus"
)
// probeTimeout bounds each transport probe. Both are local, so a daemon that is
// listening answers immediately and one that is not fails immediately.
const probeTimeout = 300 * time.Millisecond
// ResolveDaemonAddr keeps a client on the named pipe and never silently moves it
// off. When the pipe does not answer it checks the legacy loopback TCP address, so
// a client meeting a daemon that has not restarted since the upgrade can say what
// is wrong, but it does not connect there.
//
// Using that address automatically would be a downgrade the user never asked for:
// any local process can bind 127.0.0.1 while the daemon is not listening, and the
// transport carries no caller identity, so a client that accepted whatever answered
// would hand a setup key, a pre-shared key or an SSO prompt to a local impostor. An
// operator who needs the legacy address during the upgrade window can still pass
// --daemon-addr explicitly, which is a deliberate choice and still refuses the
// privileged operations.
//
// Only the pipe address is resolved. A custom address is left alone, though passing
// --daemon-addr npipe://netbird explicitly is indistinguishable from the default
// here, so it is treated the same way.
func ResolveDaemonAddr(addr string) string {
if addr != WindowsPipeAddr {
return addr
}
for _, path := range PipePaths("netbird") {
if pipeAvailable(path) {
return addr
}
}
if tcpAvailable(legacyWindowsAddr) {
log.Warnf("the daemon is not serving %s, but something is listening on the legacy %s. "+
"Restart the NetBird service so it serves the pipe. That address is not used automatically: "+
"any local user can bind it and it carries no caller identity, so pass --daemon-addr %s "+
"explicitly if you accept that",
WindowsPipeAddr, legacyWindowsAddr, legacyWindowsAddr)
}
return addr
}
func pipeAvailable(path string) bool {
timeout := probeTimeout
conn, err := winio.DialPipe(path, &timeout)
if err != nil {
return false
}
if err := conn.Close(); err != nil {
log.Debugf("close daemon pipe probe: %v", err)
}
return true
}
func tcpAvailable(addr string) bool {
host := addr
if _, after, ok := strings.Cut(addr, "://"); ok {
host = after
}
conn, err := net.DialTimeout("tcp", host, probeTimeout)
if err != nil {
return false
}
if err := conn.Close(); err != nil {
log.Debugf("close daemon TCP probe: %v", err)
}
return true
}
+49 -22
View File
@@ -229,7 +229,6 @@ scutil_dns.txt (macOS only):
const (
clientLogFile = "client.log"
uiLogFile = "gui-client.log"
errorLogFile = "netbird.err"
stdoutLogFile = "netbird.out"
@@ -248,6 +247,20 @@ type MetricsExporter interface {
Export(w io.Writer) error
}
// LogOpener opens a log file for inclusion in the bundle. It exists so that log
// files whose path was supplied by an IPC caller can be opened under a check
// the daemon defines, instead of being opened with the daemon's privileges
// unconditionally.
type LogOpener func(path string) (*os.File, error)
func openLogFile(path string) (*os.File, error) {
f, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("open %s: %w", path, err)
}
return f, nil
}
type BundleGenerator struct {
anonymizer *anonymize.Anonymizer
@@ -257,6 +270,7 @@ type BundleGenerator struct {
syncResponse *mgmProto.SyncResponse
logPath string
uiLogPath string
uiLogOpener LogOpener
tempDir string
statePath string
cpuProfile []byte
@@ -285,14 +299,20 @@ type GeneratorDependencies struct {
SyncResponse *mgmProto.SyncResponse
LogPath string
UILogPath string // Absolute path to the desktop UI's gui-client.log, reported via RegisterUILog. Empty if no UI registered one.
TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used.
StatePath string // Path to the state file. If empty, the ServiceManager default path is used.
CPUProfile []byte
CapturePath string
RefreshStatus func()
ClientMetrics MetricsExporter
DaemonVersion string
CliVersion string
// UILogOpener opens the UI log and its rotated siblings. The path comes from
// a local IPC caller, so the daemon must not open it with plain os.Open: the
// opener is where the caller's right to that file is enforced. Defaults to
// os.Open, which is only correct where the path is not caller-supplied
// (mobile).
UILogOpener LogOpener
TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used.
StatePath string // Path to the state file. If empty, the ServiceManager default path is used.
CPUProfile []byte
CapturePath string
RefreshStatus func()
ClientMetrics MetricsExporter
DaemonVersion string
CliVersion string
}
func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGenerator {
@@ -302,6 +322,11 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
logFileCount = 1
}
uiLogOpener := deps.UILogOpener
if uiLogOpener == nil {
uiLogOpener = openLogFile
}
return &BundleGenerator{
anonymizer: anonymize.NewAnonymizer(anonymize.DefaultAddresses()),
@@ -310,6 +335,7 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
syncResponse: deps.SyncResponse,
logPath: deps.LogPath,
uiLogPath: deps.UILogPath,
uiLogOpener: uiLogOpener,
tempDir: deps.TempDir,
statePath: deps.StatePath,
cpuProfile: deps.CPUProfile,
@@ -678,6 +704,7 @@ func (g *BundleGenerator) addCommonConfigFields(configContent *strings.Builder)
configContent.WriteString(fmt.Sprintf("DisableIPv6: %v\n", g.internalConfig.DisableIPv6))
configContent.WriteString(fmt.Sprintf("LocalMetricsEnabled: %v\n", g.internalConfig.LocalMetricsEnabled))
configContent.WriteString(fmt.Sprintf("LocalMetricsAddress: %s\n", g.internalConfig.LocalMetricsAddress))
configContent.WriteString(fmt.Sprintf("SyncMessageVersion: %v\n", g.internalConfig.SyncMessageVersion))
if g.internalConfig.DisableNotifications != nil {
configContent.WriteString(fmt.Sprintf("DisableNotifications: %v\n", *g.internalConfig.DisableNotifications))
@@ -997,11 +1024,11 @@ func (g *BundleGenerator) addLogfile() error {
logDir := filepath.Dir(g.logPath)
if err := g.addSingleLogfile(g.logPath, clientLogFile); err != nil {
if err := g.addSingleLogfile(openLogFile, g.logPath, clientLogFile); err != nil {
return fmt.Errorf("add client log file to zip: %w", err)
}
g.addRotatedLogFiles(logDir, clientLogPrefix)
g.addRotatedLogFiles(openLogFile, logDir, clientLogPrefix)
stdErrLogPath := filepath.Join(logDir, errorLogFile)
stdoutLogPath := filepath.Join(logDir, stdoutLogFile)
@@ -1010,11 +1037,11 @@ func (g *BundleGenerator) addLogfile() error {
stdoutLogPath = darwinStdoutLogPath
}
if err := g.addSingleLogfile(stdErrLogPath, errorLogFile); err != nil {
if err := g.addSingleLogfile(openLogFile, stdErrLogPath, errorLogFile); err != nil {
log.Warnf("Failed to add %s to zip: %v", errorLogFile, err)
}
if err := g.addSingleLogfile(stdoutLogPath, stdoutLogFile); err != nil {
if err := g.addSingleLogfile(openLogFile, stdoutLogPath, stdoutLogFile); err != nil {
log.Warnf("Failed to add %s to zip: %v", stdoutLogFile, err)
}
@@ -1031,18 +1058,18 @@ func (g *BundleGenerator) addUILog() error {
return nil
}
if err := g.addSingleLogfile(g.uiLogPath, uiLogFile); err != nil {
if err := g.addSingleLogfile(g.uiLogOpener, g.uiLogPath, configs.UILogFile); err != nil {
return fmt.Errorf("add UI log file to zip: %w", err)
}
g.addRotatedLogFiles(filepath.Dir(g.uiLogPath), uiLogPrefix)
g.addRotatedLogFiles(g.uiLogOpener, filepath.Dir(g.uiLogPath), uiLogPrefix)
return nil
}
// addSingleLogfile adds a single log file to the archive
func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error {
logFile, err := os.Open(logPath)
func (g *BundleGenerator) addSingleLogfile(open LogOpener, logPath, targetName string) error {
logFile, err := open(logPath)
if err != nil {
return fmt.Errorf("open log file %s: %w", targetName, err)
}
@@ -1067,8 +1094,8 @@ func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error {
}
// addSingleLogFileGz adds a single gzipped log file to the archive
func (g *BundleGenerator) addSingleLogFileGz(logPath, targetName string) error {
f, err := os.Open(logPath)
func (g *BundleGenerator) addSingleLogFileGz(open LogOpener, logPath, targetName string) error {
f, err := open(logPath)
if err != nil {
return fmt.Errorf("open gz log file %s: %w", targetName, err)
}
@@ -1115,7 +1142,7 @@ func (g *BundleGenerator) addSingleLogFileGz(logPath, targetName string) error {
// addRotatedLogFiles adds rotated log files to the bundle based on logFileCount.
// prefix is the base log name without extension (e.g. "client", "gui-client");
// the glob matches both files rotated by us and by logrotate on linux.
func (g *BundleGenerator) addRotatedLogFiles(logDir, prefix string) {
func (g *BundleGenerator) addRotatedLogFiles(open LogOpener, logDir, prefix string) {
if g.logFileCount == 0 {
return
}
@@ -1155,9 +1182,9 @@ func (g *BundleGenerator) addRotatedLogFiles(logDir, prefix string) {
for i := 0; i < maxFiles; i++ {
name := filepath.Base(files[i])
if strings.HasSuffix(name, ".gz") {
err = g.addSingleLogFileGz(files[i], name)
err = g.addSingleLogFileGz(open, files[i], name)
} else {
err = g.addSingleLogfile(files[i], name)
err = g.addSingleLogfile(open, files[i], name)
}
if err != nil {
log.Warnf("failed to add rotated log %s: %v", name, err)
+1 -1
View File
@@ -27,7 +27,7 @@ func (g *BundleGenerator) addPlatformLog() error {
}
swiftLogPath := filepath.Join(filepath.Dir(g.logPath), swiftLogFile)
if err := g.addSingleLogfile(swiftLogPath, swiftLogFile); err != nil {
if err := g.addSingleLogfile(openLogFile, swiftLogPath, swiftLogFile); err != nil {
// The Swift log is best-effort: the app may not have written it yet.
log.Warnf("failed to add %s to debug bundle: %v", swiftLogFile, err)
}
+1 -1
View File
@@ -97,7 +97,7 @@ func runAddRotatedLogFilesPrefix(t *testing.T, dir, prefix string, logFileCount
archive: zip.NewWriter(&buf),
logFileCount: logFileCount,
}
g.addRotatedLogFiles(dir, prefix)
g.addRotatedLogFiles(openLogFile, dir, prefix)
require.NoError(t, g.archive.Close())
zr, err := zip.NewReader(bytes.NewReader(buf.Bytes()), int64(buf.Len()))
+2
View File
@@ -887,6 +887,8 @@ func TestAddConfig_AllFieldsCovered(t *testing.T) {
ClientCertKeyPath: "/tmp/key",
LazyConnection: "on",
MTU: 1280,
DisableIPv6: true,
SyncMessageVersion: func(v int) *int { return &v }(1),
}
for _, anonymize := range []bool{false, true} {
+64
View File
@@ -0,0 +1,64 @@
package debug
import (
"archive/zip"
"bytes"
"fmt"
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/configs"
)
// bundleEntries generates a bundle with the given generator and returns the
// set of entry names in the resulting archive.
func bundleEntries(t *testing.T, g *BundleGenerator) map[string]struct{} {
t.Helper()
path, err := g.Generate()
require.NoError(t, err)
t.Cleanup(func() { _ = os.Remove(path) })
data, err := os.ReadFile(path)
require.NoError(t, err)
zr, err := zip.NewReader(bytes.NewReader(data), int64(len(data)))
require.NoError(t, err)
names := make(map[string]struct{}, len(zr.File))
for _, f := range zr.File {
names[f.Name] = struct{}{}
}
return names
}
func TestBundleIncludesUILogWhenOpenerAllows(t *testing.T) {
path := filepath.Join(t.TempDir(), configs.UILogFile)
require.NoError(t, os.WriteFile(path, []byte("gui log"), 0600))
g := NewBundleGenerator(GeneratorDependencies{
UILogPath: path,
UILogOpener: openLogFile,
}, BundleConfig{})
require.Contains(t, bundleEntries(t, g), configs.UILogFile)
}
// A UILogOpener that refuses (as the ownership check does for a foreign file)
// keeps the UI log out of the bundle without failing bundle generation.
func TestBundleExcludesUILogWhenOpenerRefuses(t *testing.T) {
path := filepath.Join(t.TempDir(), configs.UILogFile)
require.NoError(t, os.WriteFile(path, []byte("secret"), 0600))
g := NewBundleGenerator(GeneratorDependencies{
UILogPath: path,
UILogOpener: func(string) (*os.File, error) {
return nil, fmt.Errorf("not owned by the caller")
},
}, BundleConfig{})
require.NotContains(t, bundleEntries(t, g), configs.UILogFile)
}
+78 -9
View File
@@ -3,10 +3,12 @@ package debug
import (
"context"
"crypto/sha256"
"crypto/tls"
"encoding/json"
"fmt"
"io"
"net/http"
neturl "net/url"
"os"
"github.com/netbirdio/netbird/upload-server/types"
@@ -14,20 +16,80 @@ import (
const maxBundleUploadSize = 50 * 1024 * 1024
func UploadDebugBundle(ctx context.Context, url, managementURL, filePath string) (key string, err error) {
response, err := getUploadURL(ctx, url, managementURL)
// requireHTTPS refuses any URL the daemon would fetch or upload to that is not
// https. The daemon runs as root and the bundle carries its logs and state, so a
// plaintext hop is a place to intercept the bundle or the presigned redirect.
// The server-side gate already enforces this for the desktop path; this also
// covers the mobile and job-runner callers that reach this package directly.
// Skipped when the caller opted into an insecure upload (self-hosted server).
func requireHTTPS(what, rawURL string) error {
parsed, err := neturl.Parse(rawURL)
if err != nil {
return fmt.Errorf("parse %s: %w", what, err)
}
if parsed.Scheme != "https" {
return fmt.Errorf("%s must use https, got scheme %q", what, parsed.Scheme)
}
return nil
}
// uploadClient returns the HTTP client for the upload requests. The default
// client verifies TLS and refuses a redirect that would downgrade to a non-https
// hop, so a bundle can never leave over http after an https start. The insecure
// variant accepts http and untrusted certificates, and is only reachable for a
// privileged caller that passed --upload-bundle-insecure (see
// requirePrivilegeForUploadURL).
func uploadClient(insecure bool) *http.Client {
if !insecure {
return &http.Client{CheckRedirect: rejectInsecureRedirect}
}
return &http.Client{
Transport: &http.Transport{
//nolint:gosec // opt-in, privileged, self-hosted upload servers
TLSClientConfig: &tls.Config{InsecureSkipVerify: true, MinVersion: tls.VersionTLS12},
},
}
}
// rejectInsecureRedirect refuses a redirect to a non-https target and keeps the
// standard library's 10-hop limit that a custom CheckRedirect would otherwise
// disable.
func rejectInsecureRedirect(req *http.Request, via []*http.Request) error {
if req.URL.Scheme != "https" {
return fmt.Errorf("refusing redirect to non-https URL %s", req.URL.Redacted())
}
if len(via) >= 10 {
return fmt.Errorf("stopped after 10 redirects")
}
return nil
}
func UploadDebugBundle(ctx context.Context, url, managementURL, filePath string, insecure bool) (key string, err error) {
if !insecure {
if err := requireHTTPS("upload service URL", url); err != nil {
return "", err
}
}
response, err := getUploadURL(ctx, url, managementURL, insecure)
if err != nil {
return "", err
}
err = upload(ctx, filePath, response)
if !insecure {
if err := requireHTTPS("upload URL from service", response.URL); err != nil {
return "", err
}
}
err = upload(ctx, filePath, response, insecure)
if err != nil {
return "", err
}
return response.Key, nil
}
func upload(ctx context.Context, filePath string, response *types.GetURLResponse) error {
func upload(ctx context.Context, filePath string, response *types.GetURLResponse, insecure bool) error {
fileData, err := os.Open(filePath)
if err != nil {
return fmt.Errorf("open file: %w", err)
@@ -52,7 +114,7 @@ func upload(ctx context.Context, filePath string, response *types.GetURLResponse
req.ContentLength = stat.Size()
req.Header.Set("Content-Type", "application/octet-stream")
putResp, err := http.DefaultClient.Do(req)
putResp, err := uploadClient(insecure).Do(req)
if err != nil {
return fmt.Errorf("upload failed: %v", err)
}
@@ -65,16 +127,23 @@ func upload(ctx context.Context, filePath string, response *types.GetURLResponse
return nil
}
func getUploadURL(ctx context.Context, url string, managementURL string) (*types.GetURLResponse, error) {
id := getURLHash(managementURL)
getReq, err := http.NewRequestWithContext(ctx, "GET", url+"?id="+id, nil)
func getUploadURL(ctx context.Context, serviceURL string, managementURL string, insecure bool) (*types.GetURLResponse, error) {
parsed, err := neturl.Parse(serviceURL)
if err != nil {
return nil, fmt.Errorf("parse upload service URL: %w", err)
}
q := parsed.Query()
q.Set("id", getURLHash(managementURL))
parsed.RawQuery = q.Encode()
getReq, err := http.NewRequestWithContext(ctx, "GET", parsed.String(), nil)
if err != nil {
return nil, fmt.Errorf("create GET request: %w", err)
}
getReq.Header.Set(types.ClientHeader, types.ClientHeaderValue)
resp, err := http.DefaultClient.Do(getReq)
resp, err := uploadClient(insecure).Do(getReq)
if err != nil {
return nil, fmt.Errorf("get presigned URL: %w", err)
}
+46 -1
View File
@@ -5,6 +5,7 @@ import (
"errors"
"net"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
@@ -43,7 +44,7 @@ func TestUpload(t *testing.T) {
fileContent := []byte("test file content")
err := os.WriteFile(file, fileContent, 0640)
require.NoError(t, err)
key, err := UploadDebugBundle(context.Background(), testURL+types.GetURLPath, testURL, file)
key, err := UploadDebugBundle(context.Background(), testURL+types.GetURLPath, testURL, file, true)
require.NoError(t, err)
id := getURLHash(testURL)
require.Contains(t, key, id+"/")
@@ -79,3 +80,47 @@ func waitForServer(t *testing.T, addr string) {
}
t.Fatalf("server did not start listening on %s in time", addr)
}
func TestRequireHTTPS(t *testing.T) {
require.NoError(t, requireHTTPS("upload URL", "https://upload.example/path"))
require.Error(t, requireHTTPS("upload URL", "http://upload.example/path"))
require.Error(t, requireHTTPS("upload URL", "ftp://upload.example/path"))
require.Error(t, requireHTTPS("upload URL", "://malformed"))
}
func TestRejectInsecureRedirect(t *testing.T) {
httpsReq, err := http.NewRequest(http.MethodGet, "https://a.example/", nil)
require.NoError(t, err)
require.NoError(t, rejectInsecureRedirect(httpsReq, nil), "https redirect target must be allowed")
httpReq, err := http.NewRequest(http.MethodGet, "http://a.example/", nil)
require.NoError(t, err)
require.Error(t, rejectInsecureRedirect(httpReq, nil), "http redirect target must be refused")
require.Error(t, rejectInsecureRedirect(httpsReq, make([]*http.Request, 10)), "the 10-redirect limit must be enforced")
}
// The secure client refuses to follow an https response that redirects to http,
// so a bundle can't be downgraded onto plaintext mid-flight.
func TestUploadClientRefusesHTTPSToHTTPRedirect(t *testing.T) {
plain := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
}))
t.Cleanup(plain.Close)
secure := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, plain.URL, http.StatusFound)
}))
t.Cleanup(secure.Close)
client := uploadClient(false)
// Trust the test server's cert without disabling verification globally.
client.Transport = secure.Client().Transport
resp, err := client.Get(secure.URL)
if resp != nil {
_ = resp.Body.Close()
}
require.Error(t, err, "redirect from https to http must be refused")
require.Contains(t, err.Error(), "non-https")
}
+110 -20
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"net"
"net/netip"
"os"
"slices"
"strings"
"sync"
@@ -36,7 +37,43 @@ type resolver interface {
// record is left alone (it points at something outside our mesh, e.g.
// a non-peer upstream).
type PeerConnectivity interface {
IsConnectedByIP(ip string) (known, connected bool)
IsConnectedByIP(ip netip.Addr) (known, connected bool)
}
// PeerActivator wakes lazy-connection peers on demand. The local resolver calls
// it with the tunnel IPs an answer points at, so a peer that is idle (lazily
// disconnected) starts connecting at DNS-resolution time rather than racing the
// client's first request packet. nil disables warm-up.
type PeerActivator interface {
// ActivatePeersByIP triggers wake-up for the peer(s) owning addrs and blocks
// until one is connected or ctx (a short per-query budget) expires. It is a
// fast no-op for unknown or already-connected addresses.
ActivatePeersByIP(ctx context.Context, addrs []netip.Addr)
}
const (
defaultLazyWarmupTimeout = 2 * time.Second
envLazyWarmupTimeout = "NB_DNS_LAZY_WARMUP_TIMEOUT"
)
// lazyWarmupTimeoutFromEnv returns the per-query budget for waking a
// lazy-connection peer a DNS answer points at. Tunable via
// NB_DNS_LAZY_WARMUP_TIMEOUT (a Go duration). Parsed once at construction time.
func lazyWarmupTimeoutFromEnv() time.Duration {
v := os.Getenv(envLazyWarmupTimeout)
if v == "" {
return defaultLazyWarmupTimeout
}
d, err := time.ParseDuration(v)
if err != nil {
log.Warnf("invalid %s value %q, using default %s: %v", envLazyWarmupTimeout, v, defaultLazyWarmupTimeout, err)
return defaultLazyWarmupTimeout
}
if d <= 0 {
log.Warnf("non-positive %s value %q, using default %s", envLazyWarmupTimeout, v, defaultLazyWarmupTimeout)
return defaultLazyWarmupTimeout
}
return d
}
type Resolver struct {
@@ -51,6 +88,12 @@ type Resolver struct {
// filter and preserves the legacy "return whatever is registered"
// behaviour for callers that never wire a status source.
peerConn PeerConnectivity
// peerActivator, when non-nil, is called at resolution time to warm the
// lazy connection to the peer(s) an answer points at. nil disables warm-up.
peerActivator PeerActivator
// warmupTimeout is the per-query budget for the lazy-connection warm-up
// wait, resolved from the environment once at construction time.
warmupTimeout time.Duration
ctx context.Context
cancel context.CancelFunc
@@ -59,11 +102,12 @@ type Resolver struct {
func NewResolver() *Resolver {
ctx, cancel := context.WithCancel(context.Background())
return &Resolver{
records: make(map[dns.Question][]dns.RR),
domains: make(map[domain.Domain]struct{}),
zones: make(map[domain.Domain]bool),
ctx: ctx,
cancel: cancel,
records: make(map[dns.Question][]dns.RR),
domains: make(map[domain.Domain]struct{}),
zones: make(map[domain.Domain]bool),
warmupTimeout: lazyWarmupTimeoutFromEnv(),
ctx: ctx,
cancel: cancel,
}
}
@@ -76,6 +120,14 @@ func (d *Resolver) SetPeerConnectivity(p PeerConnectivity) {
d.peerConn = p
}
// SetPeerActivator wires the DNS-time lazy-connection warm-up. Pass nil to
// disable. Safe to call multiple times; the latest value wins.
func (d *Resolver) SetPeerActivator(a PeerActivator) {
d.mu.Lock()
defer d.mu.Unlock()
d.peerActivator = a
}
func (d *Resolver) MatchSubdomains() bool {
return true
}
@@ -122,6 +174,9 @@ func (d *Resolver) ServeDNS(w dns.ResponseWriter, r *dns.Msg) {
replyMessage.RecursionAvailable = true
result := d.lookupRecords(logger, question)
// Warm before filtering: activation flips a lazily-idle target to connected,
// which then lets it survive the disconnected-peer filter below.
d.warmLazyPeers(question, result.records)
result.records = d.filterDisconnectedPeerAnswers(logger, question, result.records)
replyMessage.Authoritative = !result.hasExternalData
replyMessage.Answer = result.records
@@ -495,8 +550,8 @@ func (d *Resolver) filterDisconnectedPeerAnswers(logger *log.Entry, question dns
kept := make([]dns.RR, 0, len(records))
var dropped int
for _, rr := range records {
ip := extractRecordIP(rr)
if ip == "" {
ip, ok := extractRecordAddr(rr)
if !ok {
kept = append(kept, rr)
continue
}
@@ -518,22 +573,57 @@ func (d *Resolver) filterDisconnectedPeerAnswers(logger *log.Entry, question dns
return kept
}
// extractRecordIP returns the dotted-decimal / colon-hex IP carried by
// an A or AAAA record, or "" for any other record type.
func extractRecordIP(rr dns.RR) string {
// warmLazyPeers triggers lazy-connection wake-up for the peers a resolved
// answer points at and waits briefly for one to connect, so the caller's first
// request doesn't race the connection establishment. Warm-up is scoped to
// match-only (non-authoritative) zones — the synthesized private-service zones
// and user-created zones whose records point at specific peers. The account's
// peer zone is authoritative, so plain peer-name lookups never trigger warm-up;
// otherwise resolving any peer's name would wake its idle connection, defeating
// laziness mesh-wide. No-op when no activator is wired (lazy connections
// disabled) or the answer carries no peer IPs.
func (d *Resolver) warmLazyPeers(question dns.Question, records []dns.RR) {
if len(records) < 2 {
return
}
d.mu.RLock()
activator := d.peerActivator
var nonAuth, found bool
if activator != nil {
nonAuth, found = d.findZone(question.Name)
}
d.mu.RUnlock()
if activator == nil || !found || !nonAuth {
return
}
var addrs []netip.Addr
for _, rr := range records {
if addr, ok := extractRecordAddr(rr); ok {
addrs = append(addrs, addr)
}
}
if len(addrs) == 0 {
return
}
ctx, cancel := context.WithTimeout(d.ctx, d.warmupTimeout)
defer cancel()
activator.ActivatePeersByIP(ctx, addrs)
}
// extractRecordAddr returns the IP address carried by an A or AAAA record.
// ok is false for any other record type or a record with no address.
func extractRecordAddr(rr dns.RR) (netip.Addr, bool) {
switch r := rr.(type) {
case *dns.A:
if r.A == nil {
return ""
}
return r.A.String()
addr, ok := netip.AddrFromSlice(r.A)
return addr.Unmap(), ok
case *dns.AAAA:
if r.AAAA == nil {
return ""
}
return r.AAAA.String()
addr, ok := netip.AddrFromSlice(r.AAAA)
return addr.Unmap(), ok
}
return ""
return netip.Addr{}, false
}
// Update replaces all zones and their records
+2 -2
View File
@@ -37,8 +37,8 @@ type mockPeerConnectivity struct {
byIP map[string]struct{ known, connected bool }
}
func (m mockPeerConnectivity) IsConnectedByIP(ip string) (known, connected bool) {
v, ok := m.byIP[ip]
func (m mockPeerConnectivity) IsConnectedByIP(ip netip.Addr) (known, connected bool) {
v, ok := m.byIP[ip.String()]
if !ok {
return false, false
}
+204
View File
@@ -0,0 +1,204 @@
package local
import (
"context"
"net"
"net/netip"
"sync"
"testing"
"time"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/internal/dns/test"
nbdns "github.com/netbirdio/netbird/dns"
)
// recordingActivator records the addresses it was asked to warm and returns
// immediately, so ServeDNS is not blocked by the test.
type recordingActivator struct {
mu sync.Mutex
called bool
addrs []netip.Addr
}
func (r *recordingActivator) ActivatePeersByIP(_ context.Context, addrs []netip.Addr) {
r.mu.Lock()
defer r.mu.Unlock()
r.called = true
r.addrs = append(r.addrs, addrs...)
}
func serveA(t *testing.T, resolver *Resolver, name string) *dns.Msg {
t.Helper()
var resp *dns.Msg
w := &test.MockResponseWriter{WriteMsgFunc: func(m *dns.Msg) error { resp = m; return nil }}
resolver.ServeDNS(w, new(dns.Msg).SetQuestion(name, dns.TypeA))
return resp
}
// serviceZone registers rec in a match-only (non-authoritative) zone, the shape
// the synthesized private-service zones arrive in.
func serviceZone(t *testing.T, resolver *Resolver, zone string, records ...nbdns.SimpleRecord) {
t.Helper()
resolver.Update([]nbdns.CustomZone{{
Domain: zone,
Records: records,
NonAuthoritative: true,
}})
}
func TestLocalResolver_WarmsLazyPeerOnResolve(t *testing.T) {
// Warm-up fires only for multi-record answers (the HA / round-robin shape of
// the synthesized private-service zones), so register two peer targets.
const name = "svc.proxy.netbird.cloud."
recs := []nbdns.SimpleRecord{
{Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"},
{Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.8"},
}
resolver := NewResolver()
serviceZone(t, resolver, "proxy.netbird.cloud", recs...)
act := &recordingActivator{}
resolver.SetPeerActivator(act)
resp := serveA(t, resolver, name)
require.NotNil(t, resp, "resolver must answer")
require.NotEmpty(t, resp.Answer, "answer must carry the A records")
act.mu.Lock()
defer act.mu.Unlock()
assert.True(t, act.called, "activator must be invoked for a multi-record service-zone answer")
assert.Contains(t, act.addrs, netip.MustParseAddr("100.64.0.7"), "activator must receive the first peer IP")
assert.Contains(t, act.addrs, netip.MustParseAddr("100.64.0.8"), "activator must receive the second peer IP")
}
func TestLocalResolver_NoWarmupForSingleRecord(t *testing.T) {
// A single-record answer does not trigger warm-up; the resolver only warms
// multi-record answers.
rec := nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"}
resolver := NewResolver()
serviceZone(t, resolver, "proxy.netbird.cloud", rec)
act := &recordingActivator{}
resolver.SetPeerActivator(act)
resp := serveA(t, resolver, rec.Name)
require.NotNil(t, resp, "resolver must answer")
require.NotEmpty(t, resp.Answer, "answer must carry the A record")
act.mu.Lock()
defer act.mu.Unlock()
assert.False(t, act.called, "activator must not be invoked for a single-record answer")
}
func TestLocalResolver_NoActivatorNoWarmup(t *testing.T) {
// With no activator wired the resolver behaves exactly as before.
rec := nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"}
resolver := NewResolver()
serviceZone(t, resolver, "proxy.netbird.cloud", rec)
resp := serveA(t, resolver, rec.Name)
require.NotNil(t, resp, "resolver must still answer without an activator")
require.NotEmpty(t, resp.Answer, "answer must carry the A record")
}
func TestLocalResolver_NoWarmupForMissingRecord(t *testing.T) {
// A query that resolves to nothing must not invoke the activator (no IPs).
resolver := NewResolver()
serviceZone(t, resolver, "proxy.netbird.cloud",
nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"})
act := &recordingActivator{}
resolver.SetPeerActivator(act)
serveA(t, resolver, "absent.proxy.netbird.cloud.")
act.mu.Lock()
defer act.mu.Unlock()
assert.False(t, act.called, "activator must not be invoked when there is no answer")
}
func TestLocalResolver_NoWarmupInAuthoritativeZone(t *testing.T) {
// The account's peer zone is authoritative; resolving a peer's name there
// must not wake its lazy connection — warm-up is scoped to match-only
// (non-authoritative) zones such as the synthesized private-service zones.
// Use a multi-record answer so the authoritative-zone scoping is the only
// reason warm-up is skipped, not the single-record guard.
const name = "peer.netbird.cloud."
recs := []nbdns.SimpleRecord{
{Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.9"},
{Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.10"},
}
resolver := NewResolver()
resolver.Update([]nbdns.CustomZone{{
Domain: "netbird.cloud",
Records: recs,
}})
act := &recordingActivator{}
resolver.SetPeerActivator(act)
resp := serveA(t, resolver, name)
require.NotNil(t, resp, "resolver must answer")
require.NotEmpty(t, resp.Answer, "answer must carry the A records")
act.mu.Lock()
defer act.mu.Unlock()
assert.False(t, act.called, "activator must not be invoked for authoritative-zone answers")
}
func TestLazyWarmupTimeoutFromEnv(t *testing.T) {
tests := []struct {
name string
value string
envSet bool
want time.Duration
}{
{name: "unset uses default", want: defaultLazyWarmupTimeout},
{name: "valid overrides", value: "5s", envSet: true, want: 5 * time.Second},
{name: "invalid falls back", value: "not-a-duration", envSet: true, want: defaultLazyWarmupTimeout},
{name: "negative falls back", value: "-1s", envSet: true, want: defaultLazyWarmupTimeout},
{name: "zero falls back", value: "0s", envSet: true, want: defaultLazyWarmupTimeout},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if tt.envSet {
t.Setenv(envLazyWarmupTimeout, tt.value)
}
assert.Equal(t, tt.want, lazyWarmupTimeoutFromEnv())
assert.Equal(t, tt.want, NewResolver().warmupTimeout, "constructor must resolve the timeout once")
})
}
}
func TestExtractRecordAddr(t *testing.T) {
t.Run("A record yields unmapped v4", func(t *testing.T) {
// net.ParseIP returns the 16-byte v4-in-v6 form, the same shape
// miekg/dns stores after parsing an A record; the extracted address
// must compare equal to a plain v4 netip.Addr.
addr, ok := extractRecordAddr(&dns.A{A: net.ParseIP("100.64.0.7")})
require.True(t, ok)
assert.True(t, addr.Is4())
assert.Equal(t, netip.MustParseAddr("100.64.0.7"), addr)
})
t.Run("AAAA record yields v6", func(t *testing.T) {
addr, ok := extractRecordAddr(&dns.AAAA{AAAA: net.ParseIP("fd00::1")})
require.True(t, ok)
assert.Equal(t, netip.MustParseAddr("fd00::1"), addr)
})
t.Run("A record without address", func(t *testing.T) {
_, ok := extractRecordAddr(&dns.A{})
assert.False(t, ok)
})
t.Run("non-address record", func(t *testing.T) {
_, ok := extractRecordAddr(&dns.CNAME{Target: "target.netbird.cloud."})
assert.False(t, ok)
})
}
+6
View File
@@ -8,6 +8,7 @@ import (
"github.com/miekg/dns"
dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config"
"github.com/netbirdio/netbird/client/internal/dns/local"
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/domain"
@@ -92,6 +93,11 @@ func (m *MockServer) SetFirewall(Firewall) {
// Mock implementation - no-op
}
// SetPeerActivator mock implementation of SetPeerActivator from Server interface
func (m *MockServer) SetPeerActivator(local.PeerActivator) {
// Mock implementation - no-op
}
// BeginBatch mock implementation of BeginBatch from Server interface
func (m *MockServer) BeginBatch() {
// Mock implementation - no-op
+10 -2
View File
@@ -82,6 +82,7 @@ type Server interface {
PopulateManagementDomain(mgmtURL *url.URL) error
SetRouteSources(selected, active func() route.HAMap)
SetFirewall(Firewall)
SetPeerActivator(local.PeerActivator)
}
type nsGroupsByDomain struct {
@@ -491,6 +492,13 @@ func (s *DefaultServer) SetFirewall(fw Firewall) {
}
}
// SetPeerActivator wires the DNS-time lazy-connection warm-up on the local
// resolver. Injected after the connection manager exists (it does not at
// DNS-server construction time). Pass nil to disable.
func (s *DefaultServer) SetPeerActivator(a local.PeerActivator) {
s.localResolver.SetPeerActivator(a)
}
// Stop stops the server
func (s *DefaultServer) Stop() {
s.ctxCancel()
@@ -1435,11 +1443,11 @@ type localPeerConnectivity struct {
// IsConnectedByIP looks the IP up in the peerstore and surfaces both
// the known and connected bits. Used by Resolver.filterDisconnectedPeerAnswers.
func (l localPeerConnectivity) IsConnectedByIP(ip string) (known, connected bool) {
func (l localPeerConnectivity) IsConnectedByIP(ip netip.Addr) (known, connected bool) {
if l.status == nil {
return false, false
}
state, ok := l.status.PeerStateByIP(ip)
state, ok := l.status.PeerStateByIP(ip.String())
if !ok {
return false, false
}
+4 -6
View File
@@ -292,18 +292,16 @@ func (s *serviceViaListener) generateFreePort() (uint16, error) {
return customPort, nil
}
udpAddr := net.UDPAddrFromAddrPort(netip.MustParseAddrPort("0.0.0.0:0"))
probeListener, err := net.ListenUDP("udp", udpAddr)
probeListener, err := net.ListenUDP("udp4", &net.UDPAddr{})
if err != nil {
log.Debugf("failed to bind random port for DNS: %s", err)
return 0, err
}
addrPort := netip.MustParseAddrPort(probeListener.LocalAddr().String()) // might panic if address is incorrect
err = probeListener.Close()
if err != nil {
port := uint16(probeListener.LocalAddr().(*net.UDPAddr).Port)
if err = probeListener.Close(); err != nil {
log.Debugf("failed to free up DNS port: %s", err)
return 0, err
}
return addrPort.Port(), nil
return port, nil
}
+76
View File
@@ -0,0 +1,76 @@
package internal
import (
"context"
"net/netip"
"time"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/peerstore"
)
const dnsActivationPollInterval = 50 * time.Millisecond
// dnsPeerActivator wakes lazy-connection peers from the DNS resolution path. It
// implements dns/local.PeerActivator. DNS queries run on their own goroutines,
// so it only touches state that is safe for concurrent use — ConnMgr.ActivatePeer,
// peerstore.Store and peer.Status — and never takes the engine's syncMsgMux,
// keeping DNS resolution from contending with network-map processing.
type dnsPeerActivator struct {
connMgr *ConnMgr
peerStore *peerstore.Store
status *peer.Status
// ctx is the engine's long-lived context. The connection dial is tied to it
// (not the per-query DNS wait budget) so a handshake that outlasts the wait
// still completes in the background rather than being cancelled at the deadline.
ctx context.Context
}
// ActivatePeersByIP triggers wake-up for the peer(s) owning addrs and waits
// until one is connected or ctx (the per-query DNS wait budget) expires.
// Activation itself is tied to the engine's long-lived context so the dial
// survives a wait that times out. Unknown or already-connected addresses are
// skipped, so the steady-state (warm) path adds no latency.
func (a *dnsPeerActivator) ActivatePeersByIP(ctx context.Context, addrs []netip.Addr) {
if a == nil || a.connMgr == nil {
return
}
var pending []string
for _, addr := range addrs {
ip := addr.String()
st, ok := a.status.PeerStateByIP(ip)
if !ok || st.ConnStatus == peer.StatusConnected {
continue
}
conn, ok := a.peerStore.PeerConn(st.PubKey)
if !ok {
continue
}
a.connMgr.ActivatePeer(a.ctx, conn)
pending = append(pending, ip)
}
if len(pending) == 0 {
return
}
a.waitConnected(ctx, pending)
}
// waitConnected blocks until any of ips reports a connected peer or ctx expires.
func (a *dnsPeerActivator) waitConnected(ctx context.Context, ips []string) {
ticker := time.NewTicker(dnsActivationPollInterval)
defer ticker.Stop()
for {
for _, ip := range ips {
if st, ok := a.status.PeerStateByIP(ip); ok && st.ConnStatus == peer.StatusConnected {
return
}
}
select {
case <-ctx.Done():
return
case <-ticker.C:
}
}
}
+129
View File
@@ -0,0 +1,129 @@
package internal
import (
"context"
"net/netip"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/peerstore"
)
func newTestPeerConn(t *testing.T, key string) *peer.Conn {
t.Helper()
conn, err := peer.NewConn(peer.ConnConfig{
Key: key,
LocalKey: "local",
WgConfig: peer.WgConfig{
AllowedIps: []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")},
},
}, peer.ServiceDependencies{})
require.NoError(t, err)
return conn
}
func newTestDNSPeerActivator(t *testing.T) (*dnsPeerActivator, *peer.Status, *peerstore.Store) {
t.Helper()
status := peer.NewRecorder("https://mgm")
store := peerstore.NewConnStore()
// ConnMgr without Start: the lazy manager is nil, so ActivatePeer is a
// no-op — these tests exercise the activator's skip/wait logic.
connMgr := NewConnMgr(&EngineConfig{}, status, store, nil)
return &dnsPeerActivator{
connMgr: connMgr,
peerStore: store,
status: status,
ctx: context.Background(),
}, status, store
}
func TestDNSPeerActivator_NilSafe(t *testing.T) {
var a *dnsPeerActivator
a.ActivatePeersByIP(context.Background(), []netip.Addr{netip.MustParseAddr("100.64.0.1")})
}
// TestDNSPeerActivator_SkipsUnknownAndConnectedPeers verifies the steady-state
// (warm) path adds no latency: already-connected and unknown addresses never
// enter the wait loop.
func TestDNSPeerActivator_SkipsUnknownAndConnectedPeers(t *testing.T) {
a, status, store := newTestDNSPeerActivator(t)
require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", "fd00::1"))
require.NoError(t, status.UpdatePeerState(peer.State{PubKey: "peerA", ConnStatus: peer.StatusConnected}))
store.AddPeerConn("peerA", newTestPeerConn(t, "peerA"))
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
start := time.Now()
a.ActivatePeersByIP(ctx, []netip.Addr{
netip.MustParseAddr("100.64.0.1"), // known, connected -> skipped
netip.MustParseAddr("fd00::1"), // known via IPv6, connected -> skipped
netip.MustParseAddr("100.64.0.99"), // unknown -> skipped
})
require.Less(t, time.Since(start), time.Second, "no pending peer must mean no wait")
}
// TestDNSPeerActivator_WaitsForPendingPeerToConnect verifies the wait loop
// returns as soon as a pending peer reports connected, well before the
// per-query budget expires.
func TestDNSPeerActivator_WaitsForPendingPeerToConnect(t *testing.T) {
a, status, store := newTestDNSPeerActivator(t)
require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", ""))
store.AddPeerConn("peerA", newTestPeerConn(t, "peerA"))
go func() {
time.Sleep(150 * time.Millisecond)
_ = status.UpdatePeerState(peer.State{PubKey: "peerA", ConnStatus: peer.StatusConnected})
}()
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
start := time.Now()
a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")})
elapsed := time.Since(start)
require.GreaterOrEqual(t, elapsed, 100*time.Millisecond, "must wait for the pending peer")
require.Less(t, elapsed, 5*time.Second, "must return on connect, not at the deadline")
}
// TestDNSPeerActivator_ReturnsAtBudgetWhenPeerStaysIdle verifies a peer that
// never connects releases the DNS response at the per-query budget instead of
// blocking it indefinitely.
func TestDNSPeerActivator_ReturnsAtBudgetWhenPeerStaysIdle(t *testing.T) {
a, status, store := newTestDNSPeerActivator(t)
require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", ""))
store.AddPeerConn("peerA", newTestPeerConn(t, "peerA"))
ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond)
defer cancel()
start := time.Now()
a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")})
elapsed := time.Since(start)
require.GreaterOrEqual(t, elapsed, 250*time.Millisecond, "must wait out the budget for a pending peer")
require.Less(t, elapsed, 5*time.Second, "must not block past the budget")
}
// TestDNSPeerActivator_NoWaitWithoutPeerConn verifies a known-but-idle peer
// with no connection object in the store is not waited on: there is nothing to
// activate, so waiting could only ever time out.
func TestDNSPeerActivator_NoWaitWithoutPeerConn(t *testing.T) {
a, status, _ := newTestDNSPeerActivator(t)
require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", ""))
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
start := time.Now()
a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")})
require.Less(t, time.Since(start), time.Second, "peer without a conn must not be waited on")
}
Binary file not shown.
Binary file not shown.
+3
View File
@@ -52,11 +52,14 @@ int xdp_dns_fwd(struct iphdr *ip, struct udphdr *udp) {
if (udp->dest == GENERAL_DNS_PORT && ip->daddr == dns_ip) {
udp->dest = dns_port;
// Clear the now-stale checksum; zero means "not computed" for IPv4.
udp->check = 0;
return XDP_PASS;
}
if (udp->source == dns_port && ip->saddr == dns_ip) {
udp->source = GENERAL_DNS_PORT;
udp->check = 0;
return XDP_PASS;
}
+6
View File
@@ -50,5 +50,11 @@ int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) {
__be16 new_dst_port = htons(proxy_port);
udp->dest = new_dst_port;
udp->source = new_src_port;
// The ports are covered by the UDP checksum. This is an IPv4 loopback hop
// and the payload is already integrity-protected, so clear the checksum (a
// zero UDP checksum means "not computed" for IPv4) rather than leave a
// stale value the kernel would drop as UDP_CSUM.
udp->check = 0;
return XDP_PASS;
}
+92 -3
View File
@@ -64,7 +64,10 @@ import (
"github.com/netbirdio/netbird/route"
mgm "github.com/netbirdio/netbird/shared/management/client"
"github.com/netbirdio/netbird/shared/management/domain"
sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc"
nbnetworkmap "github.com/netbirdio/netbird/shared/management/networkmap"
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
types "github.com/netbirdio/netbird/shared/management/types"
"github.com/netbirdio/netbird/shared/netiputil"
auth "github.com/netbirdio/netbird/shared/relay/auth/hmac"
relayClient "github.com/netbirdio/netbird/shared/relay/client"
@@ -147,6 +150,7 @@ type EngineConfig struct {
BlockLANAccess bool
BlockInbound bool
DisableIPv6 bool
SyncMessageVersion *int
// LazyConnection is the MDM-sourced lazy-connection override; StateUnset defers to
// the env var and management feature flag.
@@ -220,6 +224,13 @@ type Engine struct {
// networkSerial is the latest CurrentSerial (state ID) of the network sent by the Management service
networkSerial uint64
// latestComponents is the most-recent NetworkMapComponents decoded from
// a NetworkMapEnvelope (capability=3 peers only). Held alongside the
// NetworkMap that Calculate() produced from it so future incremental
// updates have a base to apply changes against. nil for legacy-format
// peers. Guarded by syncMsgMux.
latestComponents *types.NetworkMapComponents
networkMonitor *networkmonitor.NetworkMonitor
sshServer sshServer
@@ -652,8 +663,24 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
iceCfg := e.createICEConfig()
e.connMgr = NewConnMgr(e.config, e.statusRecorder, e.peerStore, wgIface)
e.connMgr.SetRoutedIPsReconciler(func(peerKey string) error {
if e.routeManager == nil {
return nil
}
return e.routeManager.ReconcilePeerAllowedIPs(peerKey)
})
e.connMgr.Start(e.ctx)
// Wire DNS-time lazy-connection warm-up now that the connection manager
// exists (it does not at DNS-server construction time). A DNS answer that
// points at an idle peer then wakes it before the client's first request.
e.dnsServer.SetPeerActivator(&dnsPeerActivator{
connMgr: e.connMgr,
peerStore: e.peerStore,
status: e.statusRecorder,
ctx: e.ctx,
})
e.srWatcher = guard.NewSRWatcher(e.signal, e.relayManager, e.mobileDep.IFaceDiscover, iceCfg)
e.srWatcher.Start(peer.IsForceRelayed())
@@ -963,8 +990,12 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
e.ApplySessionDeadline(update.GetSessionExpiresAt())
if update.NetworkMap != nil && update.NetworkMap.PeerConfig != nil {
e.handleAutoUpdateVersion(update.NetworkMap.PeerConfig.AutoUpdate)
// Envelope sync responses carry PeerConfig at the top level; legacy
// NetworkMap syncs carry it under NetworkMap.PeerConfig.
if pc := update.GetPeerConfig(); pc != nil {
e.handleAutoUpdateVersion(pc.GetAutoUpdate())
} else if nm := update.GetNetworkMap(); nm != nil && nm.GetPeerConfig() != nil {
e.handleAutoUpdateVersion(nm.GetPeerConfig().GetAutoUpdate())
}
done := e.phase("netbird_config")
@@ -974,12 +1005,47 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
return err
}
// Decode the network map from either the components envelope or the
// legacy proto.NetworkMap before the posture-check gating below, so the
// "is there a network map" decision covers both wire shapes.
var (
nm *mgmProto.NetworkMap
components *types.NetworkMapComponents
)
if version := update.GetVersion(); version == int32(sharedgrpc.ComponentNetworkMap) {
// Components-format peer: decode the envelope back to typed
// components, run Calculate() locally, and convert to the wire
// NetworkMap shape the rest of the engine consumes. Components are
// retained so future incremental updates can apply deltas instead
// of doing a full reconstruction.
envelope := update.GetNetworkMapEnvelope()
if envelope == nil {
return fmt.Errorf("received a SyncReponse indicating use of components network map, but components are missing")
}
localKey := e.config.WgPrivateKey.PublicKey().String()
dnsName := ""
if pc := update.GetPeerConfig(); pc != nil {
// PeerConfig.Fqdn = "<dns_label>.<dns_domain>" — extract the
// shared domain by stripping the peer's own label prefix. Falls
// back to empty if the FQDN doesn't have the expected shape.
dnsName = extractDNSDomainFromFQDN(pc.GetFqdn())
}
result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName)
if err != nil {
return fmt.Errorf("decode network map envelope: %w", err)
}
nm = result.NetworkMap
components = result.Components
} else {
nm = update.GetNetworkMap()
}
// Posture checks are bound to the network map presence:
// NetworkMap != nil, checks present -> apply the received checks
// NetworkMap != nil, checks nil -> posture checks were removed, clear them
// NetworkMap == nil -> config-only update (e.g. relay token rotation),
// leave the previously applied checks untouched
nm := update.GetNetworkMap()
if nm == nil {
return nil
}
@@ -992,6 +1058,14 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
}
done = e.phase("persist")
// Only retain the components view when the server sent the envelope
// path. A legacy proto.NetworkMap means components == nil; writing it
// here would clobber a previously-cached snapshot, breaking the
// incremental-delta base on a future envelope sync.
if components != nil {
e.latestComponents = components
}
e.persistSyncResponse(update)
done()
@@ -1005,6 +1079,19 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
return nil
}
// extractDNSDomainFromFQDN returns the trailing dotted domain part of the
// receiving peer's FQDN — the same value the management server fills as
// dnsName when it builds the legacy NetworkMap. "peer42.netbird.cloud" →
// "netbird.cloud". An empty string is returned for unrecognized formats.
func extractDNSDomainFromFQDN(fqdn string) string {
for i := 0; i < len(fqdn); i++ {
if fqdn[i] == '.' && i+1 < len(fqdn) {
return fqdn[i+1:]
}
}
return ""
}
// updateNetbirdConfig applies the management-provided NetBird configuration:
// STUN/TURN and relay servers, flow logging and DNS settings. A nil config is a no-op,
// which is the case for sync updates carrying only a network map.
@@ -1164,6 +1251,7 @@ func (e *Engine) applyInfoFlags(info *system.Info) {
e.config.BlockLANAccess,
e.config.BlockInbound,
e.config.DisableIPv6,
e.config.SyncMessageVersion,
e.config.EnableSSHRoot,
e.config.EnableSSHSFTP,
e.config.EnableSSHLocalPortForwarding,
@@ -2032,6 +2120,7 @@ func (e *Engine) readInitialSettings() ([]*route.Route, *nbdns.Config, bool, err
e.config.BlockLANAccess,
e.config.BlockInbound,
e.config.DisableIPv6,
e.config.SyncMessageVersion,
e.config.EnableSSHRoot,
e.config.EnableSSHSFTP,
e.config.EnableSSHLocalPortForwarding,
@@ -75,4 +75,14 @@ func TestApplySessionDeadline_ThreeState(t *testing.T) {
require.True(t, e.statusRecorder.GetSessionExpiresAt().IsZero(),
"invalid timestamp must clear the deadline")
})
t.Run("recently expired timestamp stays visible as expired", func(t *testing.T) {
e := newEngine()
expired := time.Now().Add(-5 * time.Minute).UTC().Truncate(time.Second)
e.ApplySessionDeadline(timestamppb.New(expired))
require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(expired),
"recently-expired deadline must stay on the recorder so consumers render it as expired")
})
}
+31
View File
@@ -0,0 +1,31 @@
//go:build !linux && !darwin && !freebsd && !windows
package ipcauth
import (
"errors"
"net"
"google.golang.org/grpc/credentials"
)
// errUnsupported is returned on platforms with no local peer-identity
// primitive, so consumers fail closed instead of guessing an identity.
var errUnsupported = errors.New("peer identity is not available on this platform")
// NewTransportCredentials returns nil: without a peer-identity primitive the
// daemon cannot authenticate local callers, and the caller must treat that as
// "authorization cannot be enforced".
func NewTransportCredentials() credentials.TransportCredentials {
return nil
}
// PeerIdentity always fails on this platform.
func PeerIdentity(net.Conn) (Identity, error) {
return Identity{}, errUnsupported
}
// ConnIdentity always fails on this platform.
func ConnIdentity(net.Conn) (Identity, error) {
return Identity{}, errUnsupported
}
+56
View File
@@ -0,0 +1,56 @@
//go:build linux || darwin || freebsd
package ipcauth
import (
"context"
"net"
"google.golang.org/grpc/credentials"
)
// NewTransportCredentials returns gRPC transport credentials that expose the
// caller's kernel-authenticated identity via IdentityFromContext. It returns
// nil on platforms that have no peer-identity primitive, which the caller must
// treat as "authorization cannot be enforced".
//
// The handshake exchanges no bytes on the wire, so a client dialing with
// insecure credentials interoperates with a server using these. That keeps
// older CLI and UI binaries working against an upgraded daemon.
func NewTransportCredentials() credentials.TransportCredentials {
return unixCreds{}
}
// ConnIdentity extracts the caller's identity from an accepted local IPC
// connection. It is shared by the gRPC transport credentials and by the JSON
// gateway, which reads the identity of its own HTTP clients.
func ConnIdentity(conn net.Conn) (Identity, error) {
return PeerIdentity(conn)
}
type unixCreds struct{}
func (unixCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
return conn, AuthInfo{}, nil
}
// ServerHandshake extracts the peer identity and fails closed when it cannot
// be read, so a connection whose caller is unknown never reaches a handler.
func (unixCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
id, err := ConnIdentity(conn)
if err != nil {
return nil, nil, err
}
return conn, AuthInfo{
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
Identity: id,
}, nil
}
func (unixCreds) Info() credentials.ProtocolInfo {
return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()}
}
func (unixCreds) Clone() credentials.TransportCredentials { return unixCreds{} }
func (unixCreds) OverrideServerName(string) error { return nil }
+194
View File
@@ -0,0 +1,194 @@
//go:build windows
package ipcauth
import (
"context"
"fmt"
"net"
"runtime"
log "github.com/sirupsen/logrus"
"golang.org/x/sys/windows"
"google.golang.org/grpc/credentials"
)
var (
modadvapi32 = windows.NewLazySystemDLL("advapi32.dll")
procImpersonateNamedPipeClient = modadvapi32.NewProc("ImpersonateNamedPipeClient")
)
// DefaultPipeSDDL is the security descriptor for the daemon control pipe.
//
// D:P protected DACL, no inheritance
// (A;;GA;;;SY) allow GENERIC_ALL to LocalSystem (the daemon's service account)
// (A;;GA;;;WD) allow GENERIC_ALL to Everyone
//
// Any local caller may connect, as with a Unix socket at 0666; what a caller may
// actually do is decided from its token, not from the DACL. Remote callers are not
// a concern here: winio.ListenPipe creates the pipe with
// FILE_PIPE_REJECT_REMOTE_CLIENTS, so NPFS rejects connections from other machines
// before the descriptor is consulted.
//
// A deny ACE on the NETWORK SID would not add anything and would break callers:
// that SID is present in any network-logon token, which includes OpenSSH and WinRM
// sessions, so it denies administrators driving the CLI over SSH and denies the
// daemon itself when started from such a session.
func DefaultPipeSDDL() string {
return "D:P(A;;GA;;;SY)(A;;GA;;;WD)"
}
// NewTransportCredentials returns gRPC transport credentials that derive the
// caller's identity from the named-pipe client token.
//
// The client must connect at SECURITY_IDENTIFICATION for the daemon to be able
// to read its token, which is what DialNamedPipe does.
func NewTransportCredentials() credentials.TransportCredentials {
return winpipeCreds{}
}
// ConnIdentity extracts the caller's identity from an accepted named-pipe
// connection by impersonating the pipe client and reading its token. It is
// shared by the gRPC transport credentials and by the JSON gateway, which
// reads the identity of its own HTTP clients.
func ConnIdentity(conn net.Conn) (Identity, error) {
// go-winio's pipe connection embeds *win32File, which exposes Fd().
fdConn, ok := conn.(interface{ Fd() uintptr })
if !ok {
return Identity{}, fmt.Errorf("connection %T does not expose a pipe handle", conn)
}
return pipeClientIdentity(windows.Handle(fdConn.Fd()))
}
type winpipeCreds struct{}
func (winpipeCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
return conn, AuthInfo{}, nil
}
// ServerHandshake extracts the connecting client's identity and fails closed
// when the handle or token cannot be read, so a connection whose caller is
// unknown never reaches a handler.
func (winpipeCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
id, err := ConnIdentity(conn)
if err != nil {
return nil, nil, err
}
return conn, AuthInfo{
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
Identity: id,
}, nil
}
func (winpipeCreds) Info() credentials.ProtocolInfo {
return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()}
}
func (winpipeCreds) Clone() credentials.TransportCredentials { return winpipeCreds{} }
func (winpipeCreds) OverrideServerName(string) error { return nil }
// pipeClientIdentity reads the connecting client's user SID, usable group
// SIDs, and elevation state by impersonating the pipe client on this thread
// and reading the resulting impersonation token.
func pipeClientIdentity(handle windows.Handle) (id Identity, err error) {
// Impersonation is per-thread, so the goroutine must stay on this thread
// until RevertToSelf, otherwise an unrelated goroutine could inherit the
// impersonated context.
runtime.LockOSThread()
// The thread only goes back to the runtime's pool once it is provably no
// longer impersonating the client. If the revert fails, leaving it locked
// makes Go terminate it when this goroutine exits, which costs one thread
// and keeps a thread running as the client from ever being reused.
clean := false
defer func() {
if clean {
runtime.UnlockOSThread()
}
}()
if err = impersonateNamedPipeClient(handle); err != nil {
clean = true
return Identity{}, fmt.Errorf("impersonate named pipe client: %w", err)
}
defer func() {
// Surface the revert failure only when nothing else failed: leaving
// the thread impersonated is worse than the original error.
revErr := windows.RevertToSelf()
if revErr != nil {
if err == nil {
err = fmt.Errorf("revert impersonation: %w", revErr)
}
return
}
clean = true
}()
// openAsSelf=true opens the token with the daemon's own process context
// rather than the impersonated client's, so the open cannot fail because
// the client lacks access to its own token.
var token windows.Token
if err = windows.OpenThreadToken(windows.CurrentThread(), windows.TOKEN_QUERY, true, &token); err != nil {
return Identity{}, fmt.Errorf("open thread token: %w", err)
}
defer func() {
if cerr := token.Close(); cerr != nil {
log.Debugf("close client token: %v", cerr)
}
}()
return identityFromToken(token)
}
// identityFromToken reads the user SID, usable group SIDs and elevation state
// out of a Windows token.
func identityFromToken(token windows.Token) (Identity, error) {
user, err := token.GetTokenUser()
if err != nil {
return Identity{}, fmt.Errorf("read token user: %w", err)
}
groups, err := tokenGroupSIDs(token)
if err != nil {
return Identity{}, err
}
return Identity{
SID: user.User.Sid.String(),
Groups: groups,
Elevated: token.IsElevated(),
}, nil
}
// tokenGroupSIDs returns the SIDs of the groups the token can actually
// exercise. Groups that are disabled or marked deny-only are skipped: a
// UAC-filtered administrator carries BUILTIN\Administrators as deny-only, and
// treating that as membership would hand every admin account privilege it
// cannot currently use.
func tokenGroupSIDs(token windows.Token) ([]string, error) {
tg, err := token.GetTokenGroups()
if err != nil {
return nil, fmt.Errorf("read token groups: %w", err)
}
var sids []string
for _, g := range tg.AllGroups() {
if g.Attributes&windows.SE_GROUP_ENABLED == 0 {
continue
}
if g.Attributes&windows.SE_GROUP_USE_FOR_DENY_ONLY != 0 {
continue
}
sids = append(sids, g.Sid.String())
}
return sids, nil
}
func impersonateNamedPipeClient(h windows.Handle) error {
r, _, e := procImpersonateNamedPipeClient.Call(uintptr(h))
if r == 0 {
return e
}
return nil
}
+272
View File
@@ -0,0 +1,272 @@
package ipcauth
import (
"context"
"crypto/rand"
"crypto/subtle"
"encoding/hex"
"fmt"
"slices"
"strconv"
"strings"
"google.golang.org/grpc/metadata"
)
// Metadata keys the local JSON gateway uses to forward the identity of its own
// HTTP client to the daemon. The gateway runs inside the daemon process and
// re-dials the daemon over the control socket, so without forwarding every
// JSON request would appear to come from the daemon itself.
const (
// mdFwd marks a request as forwarded by the JSON gateway. It is always
// set, even when the gateway could not read its client's identity, so the
// daemon can tell "no identity available" apart from "not forwarded".
mdFwd = "x-netbird-fwd"
mdFwdUID = "x-netbird-fwd-uid" // Unix user ID
mdFwdGID = "x-netbird-fwd-gid" // Unix primary group ID
mdFwdSID = "x-netbird-fwd-sid" // Windows user SID
mdFwdGroup = "x-netbird-fwd-group" // Windows group SID, repeated
mdFwdElevated = "x-netbird-fwd-elevated" // Windows, "1" when elevated
// mdFwdProof proves the forwarded identity was stamped by this process. The
// gateway runs inside the daemon, so a secret held in memory is available to
// the only legitimate producer and to nothing else.
mdFwdProof = "x-netbird-fwd-proof"
)
// forwardKeys is every metadata key the gateway sets. An HTTP client must never
// be able to supply one itself: see IsReservedForwardKey.
var forwardKeys = []string{mdFwd, mdFwdUID, mdFwdGID, mdFwdSID, mdFwdGroup, mdFwdElevated, mdFwdProof}
// forwardProof authenticates the gateway's forwarding metadata. It is generated
// once per daemon process and never leaves it: it is not written to disk, not
// logged, and not sent anywhere except over the daemon's own control socket to
// itself.
//
// Without it, trusting a forwarded identity rests on every layer in front of it
// stripping incoming forwarding keys, and on each key's value shape being
// distinguishable from an injected one. A single injected group SID or an
// injected "elevated" flag has the same shape as a legitimate one, so no
// cardinality rule can catch it. Requiring the proof means metadata that did not
// come from this process is refused whatever it contains.
var forwardProof = mustForwardProof()
func mustForwardProof() string {
var buf [32]byte
if _, err := rand.Read(buf[:]); err != nil {
// Continuing would leave the forwarded path authenticated by a
// predictable value, which is worse than not starting.
panic(fmt.Sprintf("generate identity forwarding proof: %v", err))
}
return hex.EncodeToString(buf[:])
}
// IsReservedForwardKey reports whether a gRPC metadata key belongs to the
// gateway's identity forwarding, and therefore must be dropped when it arrives
// from outside.
//
// grpc-gateway maps "Grpc-Metadata-<key>" request headers into gRPC metadata and
// joins them ahead of the values its own annotators add. Without dropping these,
// an HTTP client could hand the daemon "x-netbird-fwd-uid: 0" and be believed,
// because the daemon trusts forwarded metadata when the transport peer is the
// (privileged) gateway.
func IsReservedForwardKey(key string) bool {
key = strings.ToLower(key)
return slices.Contains(forwardKeys, key)
}
// ForwardIdentityMetadata encodes an HTTP client's identity for the JSON
// gateway to forward to the daemon. When known is false only the marker is
// set, which makes the daemon treat the caller as unidentified rather than as
// the daemon itself.
func ForwardIdentityMetadata(id Identity, known bool) metadata.MD {
md := metadata.MD{}
md.Set(mdFwd, "1")
md.Set(mdFwdProof, forwardProof)
if !known {
return md
}
if id.IsWindows() {
md.Set(mdFwdSID, id.SID)
if len(id.Groups) > 0 {
md.Set(mdFwdGroup, id.Groups...)
}
if id.Elevated {
md.Set(mdFwdElevated, "1")
}
return md
}
md.Set(mdFwdUID, strconv.FormatUint(uint64(id.UID), 10))
md.Set(mdFwdGID, strconv.FormatUint(uint64(id.GID), 10))
return md
}
// CallerIdentity returns the identity to authorize a request against. For a
// direct connection that is the transport peer's kernel identity. For a
// request relayed by the local JSON gateway it is the identity the gateway
// forwarded, since the transport peer is then the daemon itself.
//
// A forwarded identity is only honoured when the transport peer is the daemon's
// own identity and the metadata carries this process's forwarding proof, so
// forged forwarding metadata gains a caller nothing. A forwarded request that
// carries no identity is reported as unidentified, never as the daemon.
//
// The second return value is false when no identity could be established, and
// callers MUST fail closed in that case.
func CallerIdentity(ctx context.Context) (Identity, bool) {
id, ok := IdentityFromContext(ctx)
if !ok {
return Identity{}, false
}
// A forwarding key that arrives more than once did not come from the gateway
// alone, so nothing about the request can be trusted to describe its caller.
// Refusing outright matters because the alternative reading, "not forwarded",
// would authorize the request as the transport peer, which on the gateway's
// connection is the daemon itself.
if duplicatedForwardKey(ctx) {
return Identity{}, false
}
forwarded := isForwarded(ctx)
// Our own process on the other end of the socket is the JSON gateway, the only
// thing that dials the daemon from inside it. Such a call must carry a
// forwarded identity; without one there is no caller to authorize, and
// treating it as the daemon would authorize whatever reached the JSON socket.
// Only Linux reports the peer PID, so this is a belt on top of the gateway's
// interceptor rather than the sole guarantee.
if id.PID != 0 && int(id.PID) == selfPID && !forwarded {
return Identity{}, false
}
// Only the gateway's own connection may speak for someone else. Being
// privileged is not enough and not the point: the gateway runs inside the
// daemon, so it dials as the daemon's identity whatever user that is, which
// also covers a rootless container.
if !forwarded || !IsDaemonSelf(id) {
return id, true
}
// Speaking for someone else additionally requires the proof only this process
// holds. Refusing is the only safe reading: the transport peer here is the
// daemon itself, so falling back to it would authorize the request as the
// daemon. This is also what makes the forwarded values trustworthy once
// accepted, so they need no shape checks of their own.
if !authenticForward(ctx) {
return Identity{}, false
}
return forwardedIdentity(ctx)
}
// duplicatedForwardKey reports whether any forwarding key carries more than one
// value. The gateway's interceptor sets each key exactly once and replaces what
// was already there, so a repeat means a second source supplied it.
func duplicatedForwardKey(ctx context.Context) bool {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return false
}
for _, key := range forwardKeys {
// Group SIDs are legitimately repeated; the rest identify the caller.
if key == mdFwdGroup {
continue
}
if len(md.Get(key)) > 1 {
return true
}
}
return false
}
// authenticForward reports whether the request carries this process's forwarding
// proof, which only the in-process JSON gateway can supply.
func authenticForward(ctx context.Context) bool {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return false
}
got := mdSingle(md, mdFwdProof)
return subtle.ConstantTimeCompare([]byte(got), []byte(forwardProof)) == 1
}
// isForwarded reports whether the request carries the JSON gateway marker.
func isForwarded(ctx context.Context) bool {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return false
}
return mdSingle(md, mdFwd) != ""
}
// forwardedIdentity decodes the identity the JSON gateway attached.
func forwardedIdentity(ctx context.Context) (Identity, bool) {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return Identity{}, false
}
if sid := mdSingle(md, mdFwdSID); sid != "" {
return Identity{
SID: sid,
// Repeated by design, one value per group, and only reachable once
// the forwarding proof has been verified.
Groups: md.Get(mdFwdGroup),
Elevated: mdSingle(md, mdFwdElevated) == "1",
}, true
}
uid, err := strconv.ParseUint(mdSingle(md, mdFwdUID), 10, 32)
if err != nil {
return Identity{}, false
}
id := Identity{UID: uint32(uid)}
if gid, err := strconv.ParseUint(mdSingle(md, mdFwdGID), 10, 32); err == nil {
id.GID = uint32(gid)
}
return id, true
}
// mdSingle returns the value of a forwarded key only when exactly one was
// supplied. The gateway's interceptor sets each key exactly once, so more than one
// value means something else also supplied it, and the whole identity is treated as
// unknown rather than picking a winner. Defence in depth behind the gateway's
// header filter.
func mdSingle(md metadata.MD, key string) string {
if v := md.Get(key); len(v) == 1 {
return v[0]
}
return ""
}
// WithForwardedIdentity stamps id onto a context's outgoing metadata for the JSON
// gateway's call to the daemon, replacing any forwarding keys already present so
// values supplied from outside cannot survive alongside it.
//
// This is deliberately not done with runtime.WithMetadata: grpc-gateway skips its
// annotators entirely when no request header maps to metadata ("if len(pairs) == 0
// { return ctx, nil, nil }", runtime/context.go), which an HTTP/1.0 request with no
// Host header over a unix socket achieves. The daemon would then see an unmarked
// call whose transport peer is the daemon's own identity, and authorize it as the
// daemon. A client interceptor runs for every RPC regardless of headers.
func WithForwardedIdentity(ctx context.Context, id Identity, known bool) context.Context {
md, ok := metadata.FromOutgoingContext(ctx)
if !ok {
md = metadata.MD{}
} else {
md = md.Copy()
}
for _, key := range forwardKeys {
delete(md, key)
}
for key, values := range ForwardIdentityMetadata(id, known) {
md[key] = values
}
return metadata.NewOutgoingContext(ctx, md)
}
+214
View File
@@ -0,0 +1,214 @@
package ipcauth
import (
"context"
"testing"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/peer"
)
// transportCtx builds a request context as the daemon's transport credentials
// would: the identity of whoever opened the socket, plus whatever metadata the
// request carried.
func transportCtx(id Identity, md metadata.MD) context.Context {
ctx := peer.NewContext(context.Background(), &peer.Peer{
AuthInfo: AuthInfo{
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
Identity: id,
},
})
if md != nil {
ctx = metadata.NewIncomingContext(ctx, md)
}
return ctx
}
var (
root = Identity{UID: 0}
unprivUser = Identity{UID: 1000, GID: 1000}
)
// asDaemon pins which identity counts as this process for the duration of a test.
// Without it the test binary's own uid decides, which silently changes what
// "the gateway" means.
func asDaemon(t *testing.T, id Identity) {
t.Helper()
prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate
t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate })
selfIdentity, selfKnown = id, true
selfMayDelegate = !id.IsPrivileged()
}
func TestCallerIdentity_DirectConnections(t *testing.T) {
t.Run("no transport credentials is not an identity", func(t *testing.T) {
if _, ok := CallerIdentity(context.Background()); ok {
t.Fatal("a caller with no credentials must not be identified")
}
})
t.Run("a direct caller is its transport identity", func(t *testing.T) {
id, ok := CallerIdentity(transportCtx(unprivUser, nil))
if !ok || id.UID != 1000 {
t.Fatalf("got %v ok=%t, want uid 1000", id, ok)
}
})
// The whole point of honouring forwarded metadata only from a privileged
// transport peer: an unprivileged caller can set any metadata it likes on its
// own connection to the daemon socket.
t.Run("an unprivileged caller cannot forge an identity", func(t *testing.T) {
asDaemon(t, root)
forged := metadata.Pairs(mdFwd, "1", mdFwdUID, "0", mdFwdGID, "0")
id, ok := CallerIdentity(transportCtx(unprivUser, forged))
if !ok {
t.Fatal("caller should still be identified, as itself")
}
if id.IsPrivileged() || id.UID != 1000 {
t.Fatalf("forged metadata was believed: got %v", id)
}
})
}
func TestCallerIdentity_GatewayForwarding(t *testing.T) {
t.Run("the gateway's client identity is used, not the gateway's own", func(t *testing.T) {
asDaemon(t, root)
md := ForwardIdentityMetadata(unprivUser, true)
id, ok := CallerIdentity(transportCtx(root, md))
if !ok {
t.Fatal("forwarded identity should be usable")
}
if id.IsPrivileged() || id.UID != 1000 {
t.Fatalf("got %v, want the forwarded uid 1000 and not privileged", id)
}
})
t.Run("a privileged gateway client stays privileged", func(t *testing.T) {
asDaemon(t, root)
md := ForwardIdentityMetadata(root, true)
id, ok := CallerIdentity(transportCtx(root, md))
if !ok || !id.IsPrivileged() {
t.Fatalf("got %v ok=%t, want a privileged identity", id, ok)
}
})
// A JSON socket the gateway cannot read peer credentials from (a TCP socket,
// say) must not make every request look like the daemon itself.
t.Run("an unreadable client identity is unknown, not the daemon", func(t *testing.T) {
asDaemon(t, root)
md := ForwardIdentityMetadata(Identity{}, false)
if _, ok := CallerIdentity(transportCtx(root, md)); ok {
t.Fatal("a forwarded request with no identity must not be identified")
}
})
// grpc-gateway turns Grpc-Metadata-<key> headers into gRPC metadata and joins
// them ahead of its annotators' values. If an HTTP client's header survived
// that, this is the shape the daemon would see: the attacker's uid 0 first,
// the real uid second. The gateway filters those headers out, and reading a
// duplicated key as unknown makes the daemon safe even if it did not.
t.Run("a duplicated key from an injected header is not believed", func(t *testing.T) {
asDaemon(t, root)
md := metadata.MD{}
md.Append(mdFwd, "1")
md.Append(mdFwdUID, "0") // injected by the HTTP client
md.Append(mdFwdUID, "1000") // appended by the gateway's annotator
if id, ok := CallerIdentity(transportCtx(root, md)); ok {
t.Fatalf("injected uid was accepted: got %v", id)
}
})
t.Run("a duplicated marker is not believed either", func(t *testing.T) {
asDaemon(t, root)
md := metadata.MD{}
md.Append(mdFwd, "1")
md.Append(mdFwd, "1")
md.Append(mdFwdUID, "1000")
// A repeated marker must not be read as "not forwarded": that would
// authorize the request as the transport peer, which on the gateway's
// connection is the daemon itself.
if id, ok := CallerIdentity(transportCtx(root, md)); ok {
t.Fatalf("a duplicated marker was believed: got %v", id)
}
})
// The layers in front of this (the gateway's header matcher, and its
// interceptor replacing every forwarding key) are what keep outside metadata
// from arriving at all. The proof is what the daemon can check for itself, and
// it is the only defence that works for a value whose legitimate shape is
// indistinguishable from an injected one: a lone group SID, or "elevated".
t.Run("forwarding metadata without this process's proof is refused", func(t *testing.T) {
asDaemon(t, root)
for name, md := range map[string]metadata.MD{
"no proof": metadata.Pairs(mdFwd, "1", mdFwdUID, "0"),
"wrong proof": metadata.Pairs(mdFwd, "1", mdFwdUID, "0", mdFwdProof, "deadbeef"),
"windows identity without a proof": metadata.Pairs(mdFwd, "1",
mdFwdSID, "S-1-5-21-1-2-3-1001", mdFwdGroup, sidAdministrators, mdFwdElevated, "1"),
} {
t.Run(name, func(t *testing.T) {
if id, ok := CallerIdentity(transportCtx(root, md)); ok {
t.Fatalf("unstamped forwarding metadata was believed: got %v", id)
}
})
}
})
// A caller that reaches the gateway cannot see the proof, so it cannot append
// a group of its own to a genuine forwarded identity: doing so would have to
// go through the interceptor, which replaces the whole set.
t.Run("a group appended to a stamped identity does not survive the interceptor", func(t *testing.T) {
asDaemon(t, root)
injected := metadata.MD{}
injected.Append(mdFwdGroup, sidAdministrators)
ctx := WithForwardedIdentity(metadata.NewOutgoingContext(context.Background(), injected),
Identity{SID: "S-1-5-21-1-2-3-1001"}, true)
out, ok := metadata.FromOutgoingContext(ctx)
if !ok {
t.Fatal("no outgoing metadata")
}
if groups := out.Get(mdFwdGroup); len(groups) != 0 {
t.Fatalf("injected group survived: %v", groups)
}
})
}
func TestIsReservedForwardKey(t *testing.T) {
for _, key := range forwardKeys {
if !IsReservedForwardKey(key) {
t.Errorf("%q must be reserved", key)
}
}
// grpc-gateway canonicalises header names, so the check has to be
// case-insensitive.
if !IsReservedForwardKey("X-Netbird-Fwd-Uid") {
t.Error("the check must be case-insensitive")
}
for _, key := range []string{"authorization", "x-netbird", "x-netbird-fwd-uid-extra", ""} {
if IsReservedForwardKey(key) {
t.Errorf("%q must not be reserved", key)
}
}
}
func TestForwardIdentityMetadata_AlwaysMarksForwarded(t *testing.T) {
for _, tc := range []struct {
name string
id Identity
known bool
}{
{"known unix identity", unprivUser, true},
{"unknown identity", Identity{}, false},
{"windows identity", Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}, true},
} {
t.Run(tc.name, func(t *testing.T) {
md := ForwardIdentityMetadata(tc.id, tc.known)
if got := md.Get(mdFwd); len(got) != 1 || got[0] != "1" {
t.Fatalf("marker = %v, want exactly one \"1\"", got)
}
})
}
}
+127
View File
@@ -0,0 +1,127 @@
// Package ipcauth provides the kernel-authenticated identity of a local IPC
// (gRPC) caller and the transport credentials that surface it into the gRPC
// context, so the daemon can authorize individual RPCs by caller identity.
//
// On Unix the identity is read from the kernel via SO_PEERCRED (Linux) or
// LOCAL_PEERCRED (Darwin/FreeBSD). On Windows it is derived from the
// named-pipe client token. Platforms without a peer-identity primitive get no
// credentials, and every consumer must fail closed when no identity is
// available.
package ipcauth
import (
"context"
"fmt"
"slices"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/peer"
)
// Well-known Windows SIDs that identify a fully privileged principal.
const (
sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM
sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE
sidNetworkService = "S-1-5-20" // NT AUTHORITY\NETWORK SERVICE
sidAdministrators = "S-1-5-32-544" // BUILTIN\Administrators
)
// Identity is the kernel-authenticated identity of a local IPC caller. The
// zero value is not a valid identity: consumers must only use one obtained
// with a true ok/nil error return.
type Identity struct {
// UID and GID are the caller's Unix user ID and primary group ID. Both are
// zero on Windows, where SID is authoritative instead.
UID uint32
GID uint32
// SID is the caller's Windows security identifier, empty on Unix.
SID string
// Groups holds the caller's Windows group SIDs, captured from the client
// token at handshake time. Only groups that are enabled and not
// deny-only are captured, so a group listed here is one the caller can
// actually exercise. Empty on Unix.
Groups []string
// Elevated reports whether the Windows client token is elevated (running
// as administrator, or an administrator with UAC turned off). Always false
// on Unix, where privilege is uid 0.
Elevated bool
// PID is the caller's process ID where the platform reports it (Linux's
// SO_PEERCRED), and 0 where it does not. It identifies the daemon's own
// process dialling itself, which is what the JSON gateway does, and is never
// used to grant anything.
PID int32
}
// IsWindows reports whether this identity is a Windows principal (SID-based)
// rather than a Unix uid/gid principal.
func (i Identity) IsWindows() bool {
return i.SID != ""
}
// IsPrivileged reports whether the caller is the platform's administrative
// principal, which is what the daemon requires for changes that cross the
// user-to-root boundary.
//
// On Windows the decision comes from the caller's token rather than from
// account names or group RIDs: an elevated token, one of the service accounts
// the daemon itself may run as, or a token with BUILTIN\Administrators
// enabled. A UAC-filtered administrator has that group marked deny-only, and
// deny-only groups are dropped when the identity is captured, so such a
// caller is correctly reported as unprivileged. Domain group memberships
// (Domain Admins and friends) are deliberately not consulted: they say
// nothing about what this token may do on this machine.
func (i Identity) IsPrivileged() bool {
if !i.IsWindows() {
return i.UID == 0
}
if i.Elevated {
return true
}
switch i.SID {
case sidLocalSystem, sidLocalService, sidNetworkService:
return true
}
return slices.Contains(i.Groups, sidAdministrators)
}
// String renders the identity for audit logs and denial messages.
func (i Identity) String() string {
if i.IsWindows() {
return fmt.Sprintf("sid=%s elevated=%t", i.SID, i.Elevated)
}
return fmt.Sprintf("uid=%d gid=%d", i.UID, i.GID)
}
// AuthInfo carries the peer Identity as a gRPC credentials.AuthInfo so
// handlers can retrieve it from the request context via IdentityFromContext.
type AuthInfo struct {
credentials.CommonAuthInfo
Identity Identity
}
// AuthType identifies the authentication scheme.
func (AuthInfo) AuthType() string { return "netbird-ipc-peercred" }
// IdentityFromContext extracts the caller's kernel-authenticated identity from
// the gRPC peer context. The second return value is false when no IPC
// transport credentials were negotiated, which happens on a TCP daemon socket
// and on platforms without a peer-identity primitive. Callers MUST fail closed
// in that case.
func IdentityFromContext(ctx context.Context) (Identity, bool) {
p, ok := peer.FromContext(ctx)
if !ok {
return Identity{}, false
}
info, ok := p.AuthInfo.(AuthInfo)
if !ok {
return Identity{}, false
}
return info.Identity, true
}
+63
View File
@@ -0,0 +1,63 @@
package ipcauth
import (
"fmt"
"os"
)
// OpenOwnedFile opens path for reading on behalf of the IPC caller identified by
// id, and fails unless the opened file is a regular file that id owns.
//
// It exists for the paths a local caller hands to the daemon over the IPC. The
// daemon runs as root, so opening such a path unchecked lets any local user read
// any file through it. Ownership is the invariant that keeps the daemon from
// reading, with its own privileges, a file the caller could not read itself: a
// symlink or hard link planted at the path resolves to a file someone else owns
// and is refused.
//
// The check is made against the open descriptor rather than the path, so
// swapping the path between the check and the read cannot change the answer.
//
// A privileged caller is exempt: it can read the file directly, so refusing it
// here would protect nothing. The regular-file requirement still applies to
// everyone, since a fifo or device planted at the path is never a log file.
func OpenOwnedFile(id Identity, path string) (*os.File, error) {
f, err := openForRead(path)
if err != nil {
return nil, err
}
if err := checkOwnership(id, f); err != nil {
if cerr := f.Close(); cerr != nil {
return nil, fmt.Errorf("%w (close: %v)", err, cerr)
}
return nil, err
}
return f, nil
}
func checkOwnership(id Identity, f *os.File) error {
info, err := f.Stat()
if err != nil {
return fmt.Errorf("stat %s: %w", f.Name(), err)
}
if !info.Mode().IsRegular() {
return fmt.Errorf("%s is not a regular file", f.Name())
}
if IsPrivilegedCaller(id) {
return nil
}
owned, err := fileOwnedBy(id, f)
if err != nil {
return fmt.Errorf("read owner of %s: %w", f.Name(), err)
}
if !owned {
return fmt.Errorf("%s is not owned by the caller (%s)", f.Name(), id)
}
return nil
}
+64
View File
@@ -0,0 +1,64 @@
package ipcauth
import (
"io"
"os"
"path/filepath"
"runtime"
"testing"
"github.com/stretchr/testify/require"
)
// otherIdentity is an unprivileged caller that owns nothing the test creates.
func otherIdentity(t *testing.T) Identity {
t.Helper()
if runtime.GOOS == "windows" {
return Identity{SID: "S-1-5-21-1-2-3-1001"}
}
return Identity{UID: uint32(os.Geteuid() + 1), GID: uint32(os.Getegid() + 1)}
}
func TestOpenOwnedFileReadsFileOwnedByCaller(t *testing.T) {
path := filepath.Join(t.TempDir(), "gui-client.log")
require.NoError(t, os.WriteFile(path, []byte("hello"), 0600))
id, err := CurrentProcessIdentity()
require.NoError(t, err)
f, err := OpenOwnedFile(id, path)
require.NoError(t, err)
t.Cleanup(func() { _ = f.Close() })
content, err := io.ReadAll(f)
require.NoError(t, err)
require.Equal(t, "hello", string(content))
}
func TestOpenOwnedFileRefusesFileOwnedByAnother(t *testing.T) {
path := filepath.Join(t.TempDir(), "gui-client.log")
require.NoError(t, os.WriteFile(path, []byte("secret"), 0600))
_, err := OpenOwnedFile(otherIdentity(t), path)
require.ErrorContains(t, err, "not owned by the caller")
}
func TestOpenOwnedFileRefusesNonRegularFile(t *testing.T) {
dir := t.TempDir()
// The caller owns the directory, so this is the regular-file requirement
// talking, not the ownership check.
id, err := CurrentProcessIdentity()
require.NoError(t, err)
_, err = OpenOwnedFile(id, dir)
require.ErrorContains(t, err, "not a regular file")
}
func TestOpenOwnedFileRefusesMissingFile(t *testing.T) {
id, err := CurrentProcessIdentity()
require.NoError(t, err)
_, err = OpenOwnedFile(id, filepath.Join(t.TempDir(), "absent.log"))
require.Error(t, err)
}
+35
View File
@@ -0,0 +1,35 @@
//go:build !windows
package ipcauth
import (
"fmt"
"os"
"syscall"
)
// openForRead opens a caller-supplied path without following a symlink at its
// final component and without blocking: a fifo planted at the path would
// otherwise stall the open until a writer appears, and the daemon holds a lock
// while it collects the file.
func openForRead(path string) (*os.File, error) {
f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW|syscall.O_NONBLOCK, 0)
if err != nil {
return nil, fmt.Errorf("open %s: %w", path, err)
}
return f, nil
}
func fileOwnedBy(id Identity, f *os.File) (bool, error) {
info, err := f.Stat()
if err != nil {
return false, err
}
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return false, fmt.Errorf("no owner information in %T", info.Sys())
}
return stat.Uid == id.UID, nil
}
@@ -0,0 +1,57 @@
//go:build !windows
package ipcauth
import (
"errors"
"os"
"path/filepath"
"syscall"
"testing"
"time"
"github.com/stretchr/testify/require"
)
// A symlink is the shape the arbitrary-read attempt takes: the caller owns the
// link, the file it points at belongs to someone else.
func TestOpenOwnedFileRefusesSymlink(t *testing.T) {
dir := t.TempDir()
target := filepath.Join(dir, "target.log")
require.NoError(t, os.WriteFile(target, []byte("secret"), 0600))
link := filepath.Join(dir, "gui-client.log")
require.NoError(t, os.Symlink(target, link))
id, err := CurrentProcessIdentity()
require.NoError(t, err)
_, err = OpenOwnedFile(id, link)
// O_NOFOLLOW on a symlink reports ELOOP on Linux/Darwin and EMLINK on FreeBSD.
if !errors.Is(err, syscall.ELOOP) && !errors.Is(err, syscall.EMLINK) {
t.Fatalf("symlink open: got %v, want ELOOP or EMLINK", err)
}
}
// A fifo would block the open until a writer showed up, stalling the daemon
// while it holds its lock.
func TestOpenOwnedFileRefusesFifoWithoutBlocking(t *testing.T) {
path := filepath.Join(t.TempDir(), "gui-client.log")
require.NoError(t, syscall.Mkfifo(path, 0600))
id, err := CurrentProcessIdentity()
require.NoError(t, err)
done := make(chan error, 1)
go func() {
_, err := OpenOwnedFile(id, path)
done <- err
}()
select {
case err := <-done:
require.ErrorContains(t, err, "not a regular file")
case <-time.After(5 * time.Second):
t.Fatal("opening a fifo blocked")
}
}
@@ -0,0 +1,59 @@
//go:build windows
package ipcauth
import (
"fmt"
"os"
"golang.org/x/sys/windows"
)
// openForRead opens a caller-supplied path without following a reparse point at
// it. FILE_FLAG_OPEN_REPARSE_POINT is the Windows analogue of O_NOFOLLOW: it
// opens a symlink/junction itself rather than its target, so the regular-file
// check in checkOwnership refuses a link the caller planted to redirect the
// read. FILE_FLAG_BACKUP_SEMANTICS lets a directory open too (as os.Open does),
// so a directory planted at the path is refused as non-regular rather than
// erroring here. The share mode matches os.Open so a log being written stays
// openable.
func openForRead(path string) (*os.File, error) {
p, err := windows.UTF16PtrFromString(path)
if err != nil {
return nil, fmt.Errorf("convert path %s: %w", path, err)
}
handle, err := windows.CreateFile(
p,
windows.GENERIC_READ,
windows.FILE_SHARE_READ|windows.FILE_SHARE_WRITE|windows.FILE_SHARE_DELETE,
nil,
windows.OPEN_EXISTING,
windows.FILE_FLAG_OPEN_REPARSE_POINT|windows.FILE_FLAG_BACKUP_SEMANTICS,
0,
)
if err != nil {
return nil, fmt.Errorf("open %s: %w", path, err)
}
return os.NewFile(uintptr(handle), path), nil
}
// fileOwnedBy compares the file's owner SID with the caller's. Files an elevated
// process creates are owned by BUILTIN\Administrators rather than by the user,
// but such a caller is privileged and never reaches this check.
func fileOwnedBy(id Identity, f *os.File) (bool, error) {
// x/sys/windows GetSecurityInfo frees the OS buffer itself and returns a
// Go-heap copy, so there is nothing to LocalFree here.
sd, err := windows.GetSecurityInfo(windows.Handle(f.Fd()), windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION)
if err != nil {
return false, fmt.Errorf("read security info: %w", err)
}
owner, _, err := sd.Owner()
if err != nil {
return false, fmt.Errorf("read owner: %w", err)
}
return id.SID != "" && owner.String() == id.SID, nil
}
@@ -0,0 +1,78 @@
//go:build windows
package ipcauth
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/require"
"golang.org/x/sys/windows"
)
// fileOwnerSID reads the owner SID of path the same way OpenOwnedFile does, so
// the test can construct an Identity that matches (or deliberately does not).
func fileOwnerSID(t *testing.T, path string) string {
t.Helper()
f, err := os.Open(path)
require.NoError(t, err)
t.Cleanup(func() { _ = f.Close() })
sd, err := windows.GetSecurityInfo(windows.Handle(f.Fd()), windows.SE_FILE_OBJECT, windows.OWNER_SECURITY_INFORMATION)
require.NoError(t, err)
owner, _, err := sd.Owner()
require.NoError(t, err)
return owner.String()
}
// The allow branch of fileOwnedBy is the SID-equality path the legitimate GUI
// flow depends on. Running elevated, a created file is owned by
// BUILTIN\Administrators; an Identity carrying that SID with Elevated=false and
// no groups is unprivileged by IsPrivileged (which reads the token, not the
// SID's RID), so this exercises the real GetSecurityInfo equality rather than
// the privileged-caller shortcut.
func TestOpenOwnedFileWindowsOwnerMatchAllows(t *testing.T) {
path := filepath.Join(t.TempDir(), "gui-client.log")
require.NoError(t, os.WriteFile(path, []byte("hello"), 0600))
ownerSID := fileOwnerSID(t, path)
id := Identity{SID: ownerSID}
require.False(t, id.IsPrivileged(), "identity built from the owner SID must be unprivileged for this to test the match path")
f, err := OpenOwnedFile(id, path)
require.NoError(t, err)
_ = f.Close()
}
func TestOpenOwnedFileWindowsOwnerMismatchRefuses(t *testing.T) {
path := filepath.Join(t.TempDir(), "gui-client.log")
require.NoError(t, os.WriteFile(path, []byte("secret"), 0600))
other := Identity{SID: "S-1-5-21-9-9-9-9999"}
require.False(t, other.IsPrivileged())
_, err := OpenOwnedFile(other, path)
require.ErrorContains(t, err, "not owned by the caller")
}
// FILE_FLAG_OPEN_REPARSE_POINT must make OpenOwnedFile refuse a symlink the same
// way O_NOFOLLOW does on Unix, so a planted link can't redirect the read to
// another file. Creating a symlink needs a privilege the runner may lack, so the
// test skips rather than fails when it can't.
func TestOpenOwnedFileWindowsRefusesSymlink(t *testing.T) {
dir := t.TempDir()
target := filepath.Join(dir, "target.log")
require.NoError(t, os.WriteFile(target, []byte("secret"), 0600))
link := filepath.Join(dir, "gui-client.log")
if err := os.Symlink(target, link); err != nil {
t.Skipf("cannot create symlink (privilege not held?): %v", err)
}
id, err := CurrentProcessIdentity()
require.NoError(t, err)
_, err = OpenOwnedFile(id, link)
require.Error(t, err, "a symlink must be refused")
}
+43
View File
@@ -0,0 +1,43 @@
//go:build darwin || freebsd
package ipcauth
import (
"fmt"
"net"
"golang.org/x/sys/unix"
)
// PeerIdentity reads the kernel-authenticated identity of the process on the
// other end of a Unix socket via LOCAL_PEERCRED. The xucred is recorded by the
// kernel at connect() time and carries the peer's uid and its group list, of
// which the first entry is the primary group.
func PeerIdentity(conn net.Conn) (Identity, error) {
uc, ok := conn.(*net.UnixConn)
if !ok {
return Identity{}, fmt.Errorf("connection is not a unix socket: %T", conn)
}
raw, err := uc.SyscallConn()
if err != nil {
return Identity{}, fmt.Errorf("raw conn: %w", err)
}
var cred *unix.Xucred
var credErr error
if err := raw.Control(func(fd uintptr) {
cred, credErr = unix.GetsockoptXucred(int(fd), unix.SOL_LOCAL, unix.LOCAL_PEERCRED)
}); err != nil {
return Identity{}, fmt.Errorf("control raw conn: %w", err)
}
if credErr != nil {
return Identity{}, fmt.Errorf("read LOCAL_PEERCRED: %w", credErr)
}
id := Identity{UID: cred.Uid}
if cred.Ngroups > 0 {
id.GID = cred.Groups[0]
}
return id, nil
}
+39
View File
@@ -0,0 +1,39 @@
//go:build linux
package ipcauth
import (
"fmt"
"net"
"golang.org/x/sys/unix"
)
// PeerIdentity reads the kernel-authenticated identity of the process on the
// other end of a Unix socket via SO_PEERCRED. The credentials are recorded by
// the kernel at connect() time and cannot be changed for the life of the
// connection, so they are not spoofable by the caller.
func PeerIdentity(conn net.Conn) (Identity, error) {
uc, ok := conn.(*net.UnixConn)
if !ok {
return Identity{}, fmt.Errorf("connection is not a unix socket: %T", conn)
}
raw, err := uc.SyscallConn()
if err != nil {
return Identity{}, fmt.Errorf("raw conn: %w", err)
}
var cred *unix.Ucred
var credErr error
if err := raw.Control(func(fd uintptr) {
cred, credErr = unix.GetsockoptUcred(int(fd), unix.SOL_SOCKET, unix.SO_PEERCRED)
}); err != nil {
return Identity{}, fmt.Errorf("control raw conn: %w", err)
}
if credErr != nil {
return Identity{}, fmt.Errorf("read SO_PEERCRED: %w", credErr)
}
return Identity{UID: cred.Uid, GID: cred.Gid, PID: cred.Pid}, nil
}
@@ -0,0 +1,87 @@
//go:build windows
package ipcauth
import (
"fmt"
"net"
log "github.com/sirupsen/logrus"
"golang.org/x/sys/windows"
)
// PipeServerTrusted reports an error unless the pipe behind conn was created by a
// principal this client may hand secrets to. Clients call it for a pipe whose name
// carries no guarantee of its own, which is any name outside the
// ProtectedPrefix\Administrators namespace: that namespace already restricts
// creation to administrators and LocalSystem, while a plain name can be created by
// any local user before the daemon gets there.
//
// The decision is made from the pipe object's owner, not from the serving process,
// because a client cannot open a process running as another user at all, and the
// legitimate case is precisely an unprivileged client talking to a privileged
// daemon. Trusted owners are the service accounts, BUILTIN\Administrators, and
// this client's own user, the last of which is the daemon a user runs themselves
// as in netstack mode. A pipe owned by anyone else gets no setup key, pre-shared
// key or SSO prompt out of this client.
func PipeServerTrusted(conn net.Conn) error {
// go-winio's pipe connection embeds *win32File, which exposes Fd().
fdConn, ok := conn.(interface{ Fd() uintptr })
if !ok {
return fmt.Errorf("connection %T does not expose a pipe handle", conn)
}
owner, err := pipeOwnerSID(windows.Handle(fdConn.Fd()))
if err != nil {
return err
}
if !trustedPipeOwner(owner) {
return fmt.Errorf("pipe owned by %s, which is neither an administrator nor this user", owner)
}
return nil
}
// PipeOwnedBySelf reports whether the pipe behind conn was created by this very
// user, which is how a client recognises a daemon running as itself. Ownership it
// cannot read is reported as false.
func PipeOwnedBySelf(conn net.Conn) bool {
fdConn, ok := conn.(interface{ Fd() uintptr })
if !ok {
return false
}
owner, err := pipeOwnerSID(windows.Handle(fdConn.Fd()))
if err != nil {
log.Debugf("read daemon pipe owner: %v", err)
return false
}
return selfKnown && selfIdentity.SID != "" && owner == selfIdentity.SID
}
// pipeOwnerSID reads the owner of the pipe object a client is connected to. The
// handle was opened with GENERIC_READ, which includes READ_CONTROL, so no extra
// access is needed.
func pipeOwnerSID(handle windows.Handle) (string, error) {
sd, err := windows.GetSecurityInfo(handle, windows.SE_KERNEL_OBJECT, windows.OWNER_SECURITY_INFORMATION)
if err != nil {
return "", fmt.Errorf("read pipe security info: %w", err)
}
owner, _, err := sd.Owner()
if err != nil {
return "", fmt.Errorf("read pipe owner: %w", err)
}
return owner.String(), nil
}
// trustedPipeOwner reports whether a pipe's owner is a principal a client may
// speak to. An elevated process's objects are owned by BUILTIN\Administrators by
// default, an unelevated one's by the user, which is why both forms appear here.
func trustedPipeOwner(owner string) bool {
switch owner {
case sidLocalSystem, sidLocalService, sidNetworkService, sidAdministrators:
return true
}
return selfKnown && selfIdentity.SID != "" && owner == selfIdentity.SID
}
+125
View File
@@ -0,0 +1,125 @@
package ipcauth
import (
"os"
"runtime"
)
// Fields of the ErrorInfo detail the daemon attaches to a PermissionDenied it
// raises for an operation that requires root/administrator. Clients match on
// Reason and Domain rather than on the message text, and render the summary and
// command themselves so the user gets guidance instead of a gRPC error dump.
const (
// ErrorReasonPrivilegeRequired identifies the detail.
ErrorReasonPrivilegeRequired = "PRIVILEGE_REQUIRED"
// ErrorDomain scopes the reason to the NetBird daemon.
ErrorDomain = "daemon.netbird.io"
// ErrorMetaSummary is the one-sentence explanation of what was refused.
ErrorMetaSummary = "summary"
// ErrorMetaCommand is the command that performs the same operation with the
// privileges it needs, ready to copy and run.
ErrorMetaCommand = "command"
)
// The identity of the process evaluating callers, captured once because it cannot
// change. selfKnown is false when it could not be read, in which case nothing is
// ever treated as this process. selfMayDelegate additionally requires this
// process to be unprivileged: see IsPrivilegedCaller.
var (
selfIdentity Identity
selfKnown bool
selfMayDelegate bool
// selfPID is this process's PID, used to recognise the daemon dialling itself.
selfPID = os.Getpid()
)
func init() {
id, err := CurrentProcessIdentity()
if err != nil {
return
}
selfIdentity, selfKnown = id, true
// Only an unprivileged daemon delegates its authority to its own identity.
// When it is root or LocalSystem, sharing its identity does not mean sharing
// its power: on Windows a filtered and a full token carry the same SID, so
// matching there would let a non-elevated shell of an administrator account
// act as an administrator, which is the boundary the token check exists to
// keep.
selfMayDelegate = !id.IsPrivileged()
}
// IsDaemonSelf reports whether an identity is this very process. The JSON gateway
// runs inside the daemon and re-dials it locally, so this is what distinguishes
// the gateway from any other caller, whatever user the daemon runs as.
func IsDaemonSelf(id Identity) bool {
if !selfKnown || id.IsWindows() != selfIdentity.IsWindows() {
return false
}
if id.IsWindows() {
return id.SID != "" && id.SID == selfIdentity.SID
}
return id.UID == selfIdentity.UID
}
// IsPrivilegedCaller reports whether an identity may make the changes the daemon
// restricts to the platform administrator. This is the daemon's own rule and
// cannot be evaluated by a client, which does not know what the daemon runs as.
//
// Beyond root/administrator it accepts a caller running as the daemon's own
// identity when the daemon is itself unprivileged. That keeps a rootless container
// working, where there is no uid 0 at all, and a Windows daemon in netstack mode,
// which needs no administrator rights. In those setups a caller sharing the
// daemon's identity can already rewrite the config files it reads and replace the
// binary it runs, so refusing it a config change would protect nothing; and an
// unprivileged daemon cannot hand out a root shell in the first place.
func IsPrivilegedCaller(id Identity) bool {
if id.IsPrivileged() {
return true
}
return selfMayDelegate && IsDaemonSelf(id)
}
// SelfDelegatesTo returns the identity this process delegates its authority to,
// and whether it delegates at all. Only an unprivileged daemon does: see
// IsPrivilegedCaller. It exists so a refusal can name who may actually perform the
// operation, because on such a host root is neither required nor necessarily
// available.
func SelfDelegatesTo() (Identity, bool) {
if !selfKnown || !selfMayDelegate {
return Identity{}, false
}
return selfIdentity, true
}
// PrivilegedActor names the principal a privileged operation requires, for use
// in messages shown to the user.
func PrivilegedActor() string {
if runtime.GOOS == "windows" {
return "administrator privileges"
}
return "root"
}
// ElevatedCommand renders a command so that running it grants the privileges the
// operation needs. Windows has no in-line equivalent of sudo, so the command is
// returned unchanged and the user is expected to run it from an elevated
// terminal.
func ElevatedCommand(command string) string {
if runtime.GOOS == "windows" {
return command
}
return "sudo " + command
}
// UpCommand renders an elevated `netbird up` with the given flags, preceded by a
// `down`. The down is what makes the command work on a connected client: `netbird
// up` prints "Already connected" and returns without applying any config flag, so
// on its own the command would appear to do nothing. It is a no-op, exit 0, when
// the client is not connected.
//
// ";" rather than "&&" so the line can be pasted into any of the shells a user
// might have: PowerShell 5.1, still the default on Windows Server, rejects "&&"
// as a syntax error.
func UpCommand(flags string) string {
return ElevatedCommand("netbird down") + "; " + ElevatedCommand("netbird up "+flags)
}
+134
View File
@@ -0,0 +1,134 @@
package ipcauth
import "testing"
// The self rule is the one place privilege is granted to something other than the
// platform administrator, so its two guards matter: it must apply only when the
// daemon is itself unprivileged, and only to a caller with the daemon's identity.
func TestIsPrivilegedCaller_SelfRule(t *testing.T) {
tests := []struct {
name string
// self stands in for the process the daemon runs as.
self Identity
selfKnown bool
caller Identity
want bool
}{
{
name: "root is privileged whatever the daemon runs as",
self: Identity{UID: 1000},
selfKnown: true,
caller: Identity{UID: 0},
want: true,
},
{
name: "an unprivileged daemon delegates to its own user (rootless container)",
self: Identity{UID: 1000},
selfKnown: true,
caller: Identity{UID: 1000},
want: true,
},
{
name: "an unprivileged daemon delegates to nobody else",
self: Identity{UID: 1000},
selfKnown: true,
caller: Identity{UID: 1001},
want: false,
},
{
// The daemon is root on a normal install, so sharing its identity is
// already covered by being root; nothing else may match.
name: "a root daemon delegates to nobody",
self: Identity{UID: 0},
selfKnown: true,
caller: Identity{UID: 1000},
want: false,
},
{
// Windows netstack mode: the daemon needs no administrator rights.
name: "an unprivileged windows daemon delegates to its own SID",
self: Identity{SID: "S-1-5-21-1-2-3-1001"},
selfKnown: true,
caller: Identity{SID: "S-1-5-21-1-2-3-1001"},
want: true,
},
{
name: "an unprivileged windows daemon delegates to no other SID",
self: Identity{SID: "S-1-5-21-1-2-3-1001"},
selfKnown: true,
caller: Identity{SID: "S-1-5-21-1-2-3-1002"},
want: false,
},
{
// The UAC boundary: a filtered and a full token of the same account
// carry the same SID but not the same power, so an elevated daemon must
// never delegate to its own SID.
name: "an elevated windows daemon does not delegate to its own SID",
self: Identity{SID: "S-1-5-21-1-2-3-500", Elevated: true},
selfKnown: true,
caller: Identity{SID: "S-1-5-21-1-2-3-500"},
want: false,
},
{
name: "LocalSystem is privileged on its own merits, not by delegation",
self: Identity{SID: sidLocalSystem},
selfKnown: true,
caller: Identity{SID: sidLocalSystem},
want: true, // LocalSystem is privileged on its own merits
},
{
name: "identities of different kinds never match",
self: Identity{UID: 1000},
selfKnown: true,
caller: Identity{SID: "S-1-5-21-1-2-3-1001"},
want: false,
},
{
name: "an unknown self identity delegates to nobody",
self: Identity{},
selfKnown: false,
caller: Identity{UID: 1000},
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate
t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate })
selfIdentity, selfKnown = tt.self, tt.selfKnown
selfMayDelegate = tt.selfKnown && !tt.self.IsPrivileged()
if got := IsPrivilegedCaller(tt.caller); got != tt.want {
t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t",
tt.caller, tt.self, got, tt.want)
}
})
}
}
// The real process must never accidentally delegate: a test binary running as a
// normal user is unprivileged, so it may match itself, but nothing else.
func TestIsPrivilegedCaller_ThisProcess(t *testing.T) {
id, err := CurrentProcessIdentity()
if err != nil {
t.Skipf("cannot read this process's identity: %v", err)
}
// This process is always allowed to act as itself: either it is privileged, or
// it is unprivileged and therefore delegates to its own identity.
if !IsPrivilegedCaller(id) {
t.Errorf("this process %v was refused its own identity", id)
}
// A caller that is neither root nor this process must be refused, whatever
// this process happens to be.
other := Identity{UID: id.UID + 1}
if id.IsWindows() {
other = Identity{SID: id.SID + "9"}
}
if IsPrivilegedCaller(other) {
t.Errorf("an unrelated identity %v was treated as privileged", other)
}
}
+17
View File
@@ -0,0 +1,17 @@
//go:build !windows
package ipcauth
import "os"
// CurrentProcessIdentity returns this process's identity as the daemon would
// see it if this process connected to the local IPC. It lets a client (the UI)
// decide up front whether a privileged operation can succeed, without a
// round-trip and without duplicating the rules: the answer comes from the same
// Identity.IsPrivileged the daemon applies.
func CurrentProcessIdentity() (Identity, error) {
return Identity{
UID: uint32(os.Geteuid()),
GID: uint32(os.Getegid()),
}, nil
}
+35
View File
@@ -0,0 +1,35 @@
//go:build windows
package ipcauth
import (
"fmt"
"golang.org/x/sys/windows"
)
// CurrentProcessIdentity returns this process's identity as the daemon would see
// it if this process connected to the local IPC. It lets a client (the UI)
// decide up front whether a privileged operation can succeed, without a
// round-trip and without duplicating the rules: the answer comes from the same
// Identity.IsPrivileged the daemon applies to the token it reads off the pipe.
func CurrentProcessIdentity() (Identity, error) {
// A pseudo-token, so it must not be closed.
token := windows.GetCurrentProcessToken()
user, err := token.GetTokenUser()
if err != nil {
return Identity{}, fmt.Errorf("read token user: %w", err)
}
groups, err := tokenGroupSIDs(token)
if err != nil {
return Identity{}, err
}
return Identity{
SID: user.User.Sid.String(),
Groups: groups,
Elevated: token.IsElevated(),
}, nil
}
+37 -3
View File
@@ -29,6 +29,11 @@ type managedPeer struct {
type Config struct {
InactivityThreshold *time.Duration
// ReconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is
// armed. The activity listener creates the wake peer with the overlay /32 only; without the
// routed prefixes WireGuard would not steer subnet-bound traffic to the wake endpoint, so an
// idle routing peer could never be woken by that traffic. Optional; nil disables the reconcile.
ReconcileAllowedIPs func(peerKey string) error
}
// Manager manages lazy connections
@@ -56,6 +61,9 @@ type Manager struct {
peerToHAGroups map[string][]route.HAUniqueID // peer ID -> HA groups they belong to
haGroupToPeers map[route.HAUniqueID][]string // HA group -> peer IDs in the group
routesMu sync.RWMutex
// reconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is armed.
reconcileAllowedIPs func(peerKey string) error
}
// NewManager creates a new lazy connection manager
@@ -73,6 +81,7 @@ func NewManager(config Config, engineCtx context.Context, peerStore *peerstore.S
activityManager: activity.NewManager(wgIface),
peerToHAGroups: make(map[string][]route.HAUniqueID),
haGroupToPeers: make(map[route.HAUniqueID][]string),
reconcileAllowedIPs: config.ReconcileAllowedIPs,
}
if wgIface.IsUserspaceBind() {
@@ -201,7 +210,7 @@ func (m *Manager) AddPeer(peerCfg lazyconn.PeerConfig) (bool, error) {
return false, nil
}
if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil {
if err := m.armActivityListener(peerCfg); err != nil {
return false, err
}
@@ -288,7 +297,7 @@ func (m *Manager) DeactivatePeer(peerID peerid.ConnID) {
m.inactivityManager.RemovePeer(mp.peerCfg.PublicKey)
if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil {
if err := m.armActivityListener(*mp.peerCfg); err != nil {
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
return
}
@@ -465,6 +474,31 @@ func (m *Manager) close() {
}
// shouldDeferIdleForHA checks if peer should stay connected due to HA group requirements
// armRoutedAllowedIPs re-applies the peer's routed allowed IPs onto its freshly armed wake
// endpoint. The activity listener creates the wake peer with the overlay /32 only, so without
// this the routed prefixes would be missing and traffic to a routed subnet could not wake the
// idle routing peer. It is a no-op when no reconciler is configured.
// armActivityListener (re)arms the peer's wake endpoint via the activity manager and then
// re-applies its routed allowed IPs, so traffic to a routed subnet can wake an idle routing
// peer. The routed prefixes must be re-applied after the wake endpoint exists because the
// listener creates it with the overlay /32 only.
func (m *Manager) armActivityListener(peerCfg lazyconn.PeerConfig) error {
if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil {
return err
}
m.armRoutedAllowedIPs(&peerCfg)
return nil
}
func (m *Manager) armRoutedAllowedIPs(peerCfg *lazyconn.PeerConfig) {
if m.reconcileAllowedIPs == nil {
return
}
if err := m.reconcileAllowedIPs(peerCfg.PublicKey); err != nil {
peerCfg.Log.Errorf("failed to reconcile routed allowed IPs on wake endpoint: %v", err)
}
}
func (m *Manager) shouldDeferIdleForHA(inactivePeers map[string]struct{}, peerID string) bool {
m.routesMu.RLock()
defer m.routesMu.RUnlock()
@@ -577,7 +611,7 @@ func (m *Manager) onPeerInactivityTimedOut(peerIDs map[string]struct{}) {
mp.peerCfg.Log.Infof("start activity monitor")
if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil {
if err := m.armActivityListener(*mp.peerCfg); err != nil {
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
continue
}
+7 -5
View File
@@ -11,12 +11,14 @@ import (
// MobileDependency collect all dependencies for mobile platform
type MobileDependency struct {
// Android only
TunAdapter device.TunAdapter
IFaceDiscover stdnet.ExternalIFaceDiscover
// Android and iOS
NetworkChangeListener listener.NetworkChangeListener
HostDNSAddresses []netip.AddrPort
DnsReadyListener dns.ReadyListener
// Android only
TunAdapter device.TunAdapter
IFaceDiscover stdnet.ExternalIFaceDiscover
HostDNSAddresses []netip.AddrPort
DnsReadyListener dns.ReadyListener
// iOS only
DnsManager dns.IosDnsManager
+5 -10
View File
@@ -813,19 +813,14 @@ func (d *Status) SetSessionExpiresAt(deadline time.Time) {
}
// GetSessionExpiresAt returns the most recently recorded SSO session deadline,
// or the zero value when no deadline is tracked. A deadline that has already
// slipped into the past reports as "none": once the session has expired it is
// no longer a meaningful countdown, and the sessionwatch.Watcher does not
// arm a timer at the deadline itself to clear it (only the two pre-expiry
// warnings). Without this guard the UI would keep painting a stale
// "expires in …" against a moment that has passed until the next login,
// extend, or teardown rewrote the value.
// or the zero value when no deadline is tracked. A deadline in the past is
// returned as-is: it means the session has expired, and consumers (tray row,
// CLI status) render it as "expired" rather than hiding it — masking it as
// "none" would blank the UI at the exact moment it should say the session
// ended.
func (d *Status) GetSessionExpiresAt() time.Time {
d.mux.Lock()
defer d.mux.Unlock()
if !d.sessionExpiresAt.IsZero() && d.sessionExpiresAt.Before(time.Now()) {
return time.Time{}
}
return d.sessionExpiresAt
}
+15
View File
@@ -96,6 +96,7 @@ type ConfigInput struct {
BlockLANAccess *bool
BlockInbound *bool
DisableIPv6 *bool
SyncMessageVersion *int
DisableNotifications *bool
@@ -140,6 +141,7 @@ type Config struct {
BlockLANAccess bool
BlockInbound bool
DisableIPv6 bool
SyncMessageVersion *int
DisableNotifications *bool
@@ -607,6 +609,12 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true
}
if input.SyncMessageVersion != nil && *input.SyncMessageVersion != *config.SyncMessageVersion {
log.Infof("setting SyncMessageVersion to %v", *input.SyncMessageVersion)
*config.SyncMessageVersion = *input.SyncMessageVersion
updated = true
}
if input.DisableNotifications != nil && (config.DisableNotifications == nil || *input.DisableNotifications != *config.DisableNotifications) {
if *input.DisableNotifications {
log.Infof("disabling notifications")
@@ -764,6 +772,13 @@ func (config *Config) applyMDMPolicy(policy *mdm.Policy) {
// appended for https or ":80" for http. The serviceName parameter is
// used to contextualise error messages. On success returns the parsed
// *url.URL; on failure returns a non-nil error.
// ParseServiceURL normalises a service URL exactly as the config layer does when
// it stores one, so callers comparing a requested URL against a stored one do not
// have to reimplement the scheme validation and default-port handling.
func ParseServiceURL(serviceName, serviceURL string) (*url.URL, error) {
return parseURL(serviceName, serviceURL)
}
func parseURL(serviceName, serviceURL string) (*url.URL, error) {
parsedMgmtURL, err := url.ParseRequestURI(serviceURL)
if err != nil {
@@ -95,7 +95,7 @@ func (d *DnsInterceptor) RemoveRoute() error {
// AllowedIPs should use real IPs
if d.currentPeerKey != "" {
if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil {
if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err))
}
}
@@ -172,7 +172,7 @@ func (d *DnsInterceptor) removeAllowedIP(realPrefix netip.Prefix) error {
}
// AllowedIPs use real IPs
if _, err := d.allowedIPsRefcounter.Decrement(realPrefix); err != nil {
if _, err := d.allowedIPsRefcounter.Decrement(realPrefix, d.currentPeerKey); err != nil {
return fmt.Errorf("remove allowed IP %s: %v", realPrefix, err)
}
@@ -205,7 +205,7 @@ func (d *DnsInterceptor) RemoveAllowedIPs() error {
for _, prefixes := range d.interceptedDomains {
for _, prefix := range prefixes {
// AllowedIPs use real IPs
if _, err := d.allowedIPsRefcounter.Decrement(prefix); err != nil {
if _, err := d.allowedIPsRefcounter.Decrement(prefix, d.currentPeerKey); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %v", prefix, err))
}
}
+22 -8
View File
@@ -135,7 +135,7 @@ func (r *Route) RemoveAllowedIPs() error {
var merr *multierror.Error
for _, domainPrefixes := range r.dynamicDomains {
for _, prefix := range domainPrefixes {
if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil {
if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err))
}
}
@@ -185,7 +185,7 @@ func (r *Route) startResolver(ctx context.Context) {
}
func (r *Route) update(ctx context.Context) error {
resolved, err := r.resolveDomains()
resolved, err := r.resolveDomains(ctx)
if err != nil {
if len(resolved) == 0 {
return fmt.Errorf("resolve domains: %w", err)
@@ -199,9 +199,9 @@ func (r *Route) update(ctx context.Context) error {
return nil
}
func (r *Route) resolveDomains() (domainMap, error) {
func (r *Route) resolveDomains(ctx context.Context) (domainMap, error) {
results := make(chan resolveResult)
go r.resolve(results)
go r.resolve(ctx, results)
resolved := domainMap{}
var merr *multierror.Error
@@ -217,7 +217,7 @@ func (r *Route) resolveDomains() (domainMap, error) {
return resolved, nberrors.FormatErrorOrNil(merr)
}
func (r *Route) resolve(results chan resolveResult) {
func (r *Route) resolve(ctx context.Context, results chan resolveResult) {
var wg sync.WaitGroup
for _, d := range r.route.Domains {
@@ -225,10 +225,10 @@ func (r *Route) resolve(results chan resolveResult) {
go func(domain domain.Domain) {
defer wg.Done()
ips, err := r.getIPsFromResolver(domain)
ips, err := r.getIPsFromResolver(ctx, domain)
if err != nil {
log.Tracef("Failed to resolve domain %s with private resolver: %v", domain.SafeString(), err)
ips, err = net.LookupIP(domain.PunycodeString())
ips, err = lookupHostIPs(ctx, domain)
if err != nil {
results <- resolveResult{domain: domain, err: fmt.Errorf("resolve d %s: %w", domain.SafeString(), err)}
return
@@ -320,7 +320,7 @@ func (r *Route) removeRoutes(prefixes []netip.Prefix) ([]netip.Prefix, error) {
merr = multierror.Append(merr, fmt.Errorf("remove dynamic route for IP %s: %w", prefix, err))
}
if r.currentPeerKey != "" {
if _, err := r.allowedIPsRefcounter.Decrement(prefix); err != nil {
if _, err := r.allowedIPsRefcounter.Decrement(prefix, r.currentPeerKey); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %s: %w", prefix, err))
}
}
@@ -364,6 +364,20 @@ func determinePrefixChanges(oldPrefixes, newPrefixes []netip.Prefix) (toAdd, toR
return
}
// lookupHostIPs resolves d via the system resolver, honoring ctx cancellation.
func lookupHostIPs(ctx context.Context, d domain.Domain) ([]net.IP, error) {
addrs, err := net.DefaultResolver.LookupIPAddr(ctx, d.PunycodeString())
if err != nil {
return nil, err
}
ips := make([]net.IP, 0, len(addrs))
for _, addr := range addrs {
ips = append(ips, addr.IP)
}
return ips, nil
}
func combinePrefixes(oldPrefixes, removedPrefixes, addedPrefixes []netip.Prefix) []netip.Prefix {
prefixSet := make(map[netip.Prefix]struct{})
for _, prefix := range oldPrefixes {
@@ -3,11 +3,12 @@
package dynamic
import (
"context"
"net"
"github.com/netbirdio/netbird/shared/management/domain"
)
func (r *Route) getIPsFromResolver(domain domain.Domain) ([]net.IP, error) {
return net.LookupIP(domain.PunycodeString())
func (r *Route) getIPsFromResolver(ctx context.Context, domain domain.Domain) ([]net.IP, error) {
return lookupHostIPs(ctx, domain)
}
@@ -3,6 +3,7 @@
package dynamic
import (
"context"
"fmt"
"net"
"time"
@@ -16,7 +17,7 @@ import (
const dialTimeout = 10 * time.Second
func (r *Route) getIPsFromResolver(domain domain.Domain) ([]net.IP, error) {
func (r *Route) getIPsFromResolver(ctx context.Context, domain domain.Domain) ([]net.IP, error) {
privateClient, err := nbdns.GetClientPrivate(r.wgInterface, r.resolverAddr.Addr(), dialTimeout)
if err != nil {
return nil, fmt.Errorf("error while creating private client: %s", err)
@@ -32,7 +33,7 @@ func (r *Route) getIPsFromResolver(domain domain.Domain) ([]net.IP, error) {
msg := new(dns.Msg)
msg.SetQuestion(fqdn, qtype)
response, _, err := nbdns.ExchangeWithFallback(nil, privateClient, msg, r.resolverAddr.String())
response, _, err := nbdns.ExchangeWithFallback(ctx, privateClient, msg, r.resolverAddr.String())
if err != nil {
if queryErr == nil {
queryErr = fmt.Errorf("DNS query for %s (type %d) after %s: %w", domain.SafeString(), qtype, time.Since(startTime), err)
+31 -9
View File
@@ -52,6 +52,10 @@ type Manager interface {
UpdateRoutes(updateSerial uint64, serverRoutes map[route.ID]*route.Route, clientRoutes route.HAMap, useNewDNSRoute bool) error
ClassifyRoutes(newRoutes []*route.Route) (map[route.ID]*route.Route, route.HAMap)
TriggerSelection(route.HAMap)
SelectRoutes(ids []route.NetID, appendRoute bool) error
DeselectRoutes(ids []route.NetID) error
SelectAllRoutes()
DeselectAllRoutes()
GetRouteSelector() *routeselector.RouteSelector
GetClientRoutes() route.HAMap
GetSelectedClientRoutes() route.HAMap
@@ -61,6 +65,7 @@ type Manager interface {
InitialRouteRange() []string
SetFirewall(firewall.Manager) error
SetDNSForwarderPort(port uint16)
ReconcilePeerAllowedIPs(peerKey string) error
Stop(stateManager *statemanager.Manager)
}
@@ -215,7 +220,7 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) {
)
}
m.allowedIPsRefCounter = refcounter.New(
m.allowedIPsRefCounter = refcounter.NewAllowedIPs(
func(prefix netip.Prefix, peerKey string) (string, error) {
// save peerKey to use it in the remove function
return peerKey, m.wgInterface.AddAllowedIP(peerKey, prefix)
@@ -232,6 +237,30 @@ func (m *DefaultManager) setupRefCounters(useNoop bool) {
)
}
// ReconcilePeerAllowedIPs re-applies every routed allowed IP currently tracked for the peer
// onto the WireGuard device. The allowed-IP refcounter only calls its AddFunc (which pushes to
// the device) on a prefix's 0->1 transition, so a peer whose device entry was rebuilt without a
// matching refcounter change — e.g. a lazy connection cycling through idle->wake, which recreates
// the WireGuard peer with the overlay /32 only — ends up missing routed prefixes the refcounter
// still considers installed, and nothing retries. Calling this when the peer's WireGuard entry is
// (re)created restores convergence. It is add-only and idempotent: AddAllowedIP is update-only, so
// prefixes are re-added to an existing peer and an absent peer is left untouched.
func (m *DefaultManager) ReconcilePeerAllowedIPs(peerKey string) error {
if m.allowedIPsRefCounter == nil {
return nil
}
return m.allowedIPsRefCounter.ReapplyMatching(
func(out string) bool { return out == peerKey },
func(prefix netip.Prefix) error {
if err := m.wgInterface.AddAllowedIP(peerKey, prefix); err != nil {
return fmt.Errorf("add allowed IP %s for peer %s: %w", prefix, peerKey, err)
}
return nil
},
)
}
// Init sets up the routing
func (m *DefaultManager) Init() error {
m.routeSelector = m.initSelector()
@@ -775,7 +804,7 @@ func (m *DefaultManager) collectExitNodeInfo(clientRoutes route.HAMap) exitNodeI
var info exitNodeInfo
for haID, routes := range clientRoutes {
if !m.isExitNodeRoute(routes) {
if !isExitNodeRoutes(routes) {
continue
}
@@ -795,13 +824,6 @@ func (m *DefaultManager) collectExitNodeInfo(clientRoutes route.HAMap) exitNodeI
return info
}
func (m *DefaultManager) isExitNodeRoute(routes []*route.Route) bool {
if len(routes) == 0 {
return false
}
return route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network)
}
func (m *DefaultManager) categorizeUserSelection(netID route.NetID, info *exitNodeInfo) {
if m.routeSelector.IsSelected(netID) {
info.userSelected = append(info.userSelected, netID)
+31
View File
@@ -16,6 +16,8 @@ type MockManager struct {
ClassifyRoutesFunc func(routes []*route.Route) (map[route.ID]*route.Route, route.HAMap)
UpdateRoutesFunc func(updateSerial uint64, serverRoutes map[route.ID]*route.Route, clientRoutes route.HAMap, useNewDNSRoute bool) error
TriggerSelectionFunc func(haMap route.HAMap)
SelectRoutesFunc func(ids []route.NetID, appendRoute bool) error
DeselectRoutesFunc func(ids []route.NetID) error
GetRouteSelectorFunc func() *routeselector.RouteSelector
GetClientRoutesFunc func() route.HAMap
GetSelectedClientRoutesFunc func() route.HAMap
@@ -55,6 +57,30 @@ func (m *MockManager) TriggerSelection(networks route.HAMap) {
}
}
// SelectRoutes mock implementation of SelectRoutes from Manager interface
func (m *MockManager) SelectRoutes(ids []route.NetID, appendRoute bool) error {
if m.SelectRoutesFunc != nil {
return m.SelectRoutesFunc(ids, appendRoute)
}
return nil
}
// DeselectRoutes mock implementation of DeselectRoutes from Manager interface
func (m *MockManager) DeselectRoutes(ids []route.NetID) error {
if m.DeselectRoutesFunc != nil {
return m.DeselectRoutesFunc(ids)
}
return nil
}
// SelectAllRoutes mock implementation of SelectAllRoutes from Manager interface
func (m *MockManager) SelectAllRoutes() {
}
// DeselectAllRoutes mock implementation of DeselectAllRoutes from Manager interface
func (m *MockManager) DeselectAllRoutes() {
}
// GetRouteSelector mock implementation of GetRouteSelector from Manager interface
func (m *MockManager) GetRouteSelector() *routeselector.RouteSelector {
if m.GetRouteSelectorFunc != nil {
@@ -112,6 +138,11 @@ func (m *MockManager) SetFirewall(firewall.Manager) error {
func (m *MockManager) SetDNSForwarderPort(port uint16) {
}
// ReconcilePeerAllowedIPs mock implementation of ReconcilePeerAllowedIPs from Manager interface
func (m *MockManager) ReconcilePeerAllowedIPs(peerKey string) error {
return nil
}
// Stop mock implementation of Stop from Manager interface
func (m *MockManager) Stop(stateManager *statemanager.Manager) {
if m.StopFunc != nil {
@@ -3,7 +3,6 @@
package notifier
import (
"container/list"
"net/netip"
"slices"
"sort"
@@ -16,20 +15,12 @@ import (
type Notifier struct {
mu sync.Mutex
cond *sync.Cond
currentPrefixes []string
listener listener.NetworkChangeListener
queue *list.List
closed bool
}
func NewNotifier() *Notifier {
n := &Notifier{
queue: list.New(),
}
n.cond = sync.NewCond(&n.mu)
go n.deliverLoop()
return n
return &Notifier{}
}
func (n *Notifier) SetListener(listener listener.NetworkChangeListener) {
@@ -59,44 +50,19 @@ func (n *Notifier) OnNewPrefixes(prefixes []netip.Prefix) {
sort.Strings(newNets)
n.mu.Lock()
defer n.mu.Unlock()
if slices.Equal(n.currentPrefixes, newNets) {
n.mu.Unlock()
return
}
n.currentPrefixes = newNets
routes := strings.Join(n.currentPrefixes, ",")
n.queue.PushBack(routes)
n.cond.Signal()
n.mu.Unlock()
if n.listener != nil {
n.listener.OnNetworkChanged(strings.Join(n.currentPrefixes, ","))
}
}
func (n *Notifier) Close() {
n.mu.Lock()
n.closed = true
n.cond.Signal()
n.mu.Unlock()
}
func (n *Notifier) GetInitialRouteRanges() []string {
return nil
}
func (n *Notifier) deliverLoop() {
for {
n.mu.Lock()
for n.queue.Len() == 0 && !n.closed {
n.cond.Wait()
}
if n.closed && n.queue.Len() == 0 {
n.mu.Unlock()
return
}
routes := n.queue.Remove(n.queue.Front()).(string)
l := n.listener
n.mu.Unlock()
if l != nil {
l.OnNetworkChanged(routes)
}
}
}
@@ -0,0 +1,90 @@
//go:build !windows
package routemanager
import (
"net"
"net/netip"
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/tun/netstack"
"github.com/netbirdio/netbird/client/iface/device"
"github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/client/internal/routemanager/refcounter"
)
// reconcileWGMock is a minimal iface.WGIface that only records AddAllowedIP calls; every other
// method is an inert stub because ReconcilePeerAllowedIPs exercises none of them.
type reconcileWGMock struct {
mu sync.Mutex
adds map[string][]netip.Prefix
}
func (m *reconcileWGMock) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
m.mu.Lock()
defer m.mu.Unlock()
if m.adds == nil {
m.adds = map[string][]netip.Prefix{}
}
m.adds[peerKey] = append(m.adds[peerKey], allowedIP)
return nil
}
func (m *reconcileWGMock) added(peerKey string) []netip.Prefix {
m.mu.Lock()
defer m.mu.Unlock()
return m.adds[peerKey]
}
func (m *reconcileWGMock) RemoveAllowedIP(string, netip.Prefix) error { return nil }
func (m *reconcileWGMock) Name() string { return "utun-test" }
func (m *reconcileWGMock) Address() wgaddr.Address { return wgaddr.Address{} }
func (m *reconcileWGMock) ToInterface() *net.Interface { return nil }
func (m *reconcileWGMock) IsUserspaceBind() bool { return false }
func (m *reconcileWGMock) GetFilter() device.PacketFilter { return nil }
func (m *reconcileWGMock) GetDevice() *device.FilteredDevice { return nil }
func (m *reconcileWGMock) GetNet() *netstack.Net { return nil }
// TestReconcilePeerAllowedIPs verifies the declarative reconcile re-applies every routed prefix
// tracked for the peer (self-heal, independent of refcount level) and stays scoped to that peer.
func TestReconcilePeerAllowedIPs(t *testing.T) {
wg := &reconcileWGMock{}
m := &DefaultManager{wgInterface: wg}
m.allowedIPsRefCounter = refcounter.NewAllowedIPs(
func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil },
func(netip.Prefix, string) error { return nil },
)
peerA1 := netip.MustParsePrefix("10.0.0.0/24")
peerA2 := netip.MustParsePrefix("10.1.0.0/24")
peerB1 := netip.MustParsePrefix("10.2.0.0/24")
for prefix, peer := range map[netip.Prefix]string{peerA1: "peerA", peerA2: "peerA", peerB1: "peerB"} {
_, err := m.allowedIPsRefCounter.Increment(prefix, peer)
require.NoError(t, err)
}
// Extra reference: reconcile must still re-apply the prefix even though its refcount never
// hit 0 again (the exact case the plain incremental path skips).
_, err := m.allowedIPsRefCounter.Increment(peerA1, "peerA")
require.NoError(t, err)
require.NoError(t, m.ReconcilePeerAllowedIPs("peerA"))
assert.ElementsMatch(t, []netip.Prefix{peerA1, peerA2}, wg.added("peerA"),
"reconcile must re-apply all routed prefixes of the peer")
assert.Empty(t, wg.added("peerB"), "reconcile must not touch another peer's prefixes")
}
// TestReconcilePeerAllowedIPsNoCounter verifies reconcile is a safe no-op before the refcounter is
// set up.
func TestReconcilePeerAllowedIPsNoCounter(t *testing.T) {
wg := &reconcileWGMock{}
m := &DefaultManager{wgInterface: wg}
require.NoError(t, m.ReconcilePeerAllowedIPs("peerA"))
assert.Empty(t, wg.added("peerA"))
}
@@ -0,0 +1,206 @@
package refcounter
import (
"errors"
"fmt"
"net/netip"
"sort"
"sync"
"github.com/hashicorp/go-multierror"
nberrors "github.com/netbirdio/netbird/client/errors"
)
// allowedIPsEntry holds the per-peer reference counts for a single prefix and which peer is
// currently installed in WireGuard. WireGuard allows a prefix on exactly one peer, so at most
// one peer is active at a time even when several peers reference the prefix.
type allowedIPsEntry struct {
// peers maps a peerKey to the number of references holding the prefix for that peer.
peers map[string]int
// active is the peerKey currently installed in WireGuard for this prefix ("" if none).
active string
// total is the sum of all per-peer reference counts (kept in sync with peers).
total int
}
// AllowedIPsRefCounter is a peer-aware reference counter for WireGuard AllowedIPs.
//
// The generic Counter keys only by prefix and remembers a single Out value set by the first
// caller, which it never changes. That is wrong for AllowedIPs: two independent watchers (or
// multiple resolved domains) can reference the same prefix through different peers, and when the
// peer currently installed in WireGuard releases its last reference the prefix must be handed over
// to a surviving peer instead of being left pointing at the released one.
//
// It calls add/remove (which program WireGuard) only on the transitions that matter:
// - add on the first reference for a prefix, or when swapping the active peer;
// - remove on the last reference for a prefix, or on the old peer during a swap.
type AllowedIPsRefCounter struct {
mu sync.Mutex
entries map[netip.Prefix]*allowedIPsEntry
add AddFunc[netip.Prefix, string, string]
remove RemoveFunc[netip.Prefix, string]
}
// NewAllowedIPs creates a new peer-aware AllowedIPs reference counter.
// add programs a prefix on a peer in WireGuard and returns the peerKey to store as the active peer.
// remove unprograms the prefix from the given peer.
func NewAllowedIPs(add AddFunc[netip.Prefix, string, string], remove RemoveFunc[netip.Prefix, string]) *AllowedIPsRefCounter {
return &AllowedIPsRefCounter{
entries: map[netip.Prefix]*allowedIPsEntry{},
add: add,
remove: remove,
}
}
// Increment adds a reference to prefix for peerKey. WireGuard is programmed only for the first
// reference to a prefix; while a different peer is already installed the prefix is left with it
// (first peer wins, HA at the WireGuard layer is not possible) and only the reference count is kept.
func (rm *AllowedIPsRefCounter) Increment(prefix netip.Prefix, peerKey string) (Ref[string], error) {
rm.mu.Lock()
defer rm.mu.Unlock()
e, ok := rm.entries[prefix]
if !ok {
e = &allowedIPsEntry{peers: map[string]int{}}
rm.entries[prefix] = e
}
logCallerF("Increasing allowed IP ref count for prefix %v peer %s [peer %d -> %d, total %d -> %d, active %q]",
prefix, peerKey, e.peers[peerKey], e.peers[peerKey]+1, e.total, e.total+1, e.active)
// Program WireGuard only when nothing is installed yet for this prefix.
if e.active == "" {
out, err := rm.add(prefix, peerKey)
if errors.Is(err, ErrIgnore) {
if e.total == 0 {
delete(rm.entries, prefix)
}
return Ref[string]{Count: e.total, Out: e.active}, nil
}
if err != nil {
if e.total == 0 {
delete(rm.entries, prefix)
}
return Ref[string]{}, fmt.Errorf("failed to add allowed IP %v for peer %s: %w", prefix, peerKey, err)
}
e.active = out
}
e.peers[peerKey]++
e.total++
return Ref[string]{Count: e.total, Out: e.active}, nil
}
// Decrement removes a reference to prefix for peerKey. When the peer currently installed in
// WireGuard releases its last reference, the prefix is swapped to a surviving peer if one exists,
// otherwise it is removed from WireGuard.
func (rm *AllowedIPsRefCounter) Decrement(prefix netip.Prefix, peerKey string) (Ref[string], error) {
rm.mu.Lock()
defer rm.mu.Unlock()
e, ok := rm.entries[prefix]
if !ok {
logCallerF("No allowed IP reference found for prefix %v", prefix)
return Ref[string]{}, nil
}
if e.peers[peerKey] > 0 {
logCallerF("Decreasing allowed IP ref count for prefix %v peer %s [peer %d -> %d, total %d -> %d, active %q]",
prefix, peerKey, e.peers[peerKey], e.peers[peerKey]-1, e.total, e.total-1, e.active)
e.peers[peerKey]--
e.total--
if e.peers[peerKey] == 0 {
delete(e.peers, peerKey)
}
} else {
logCallerF("No allowed IP reference found for prefix %v peer %s", prefix, peerKey)
}
// If the peer currently installed in WireGuard still holds references, nothing to reprogram.
// Keying the check on the active peer (not the one just released) makes this self-healing:
// a prior swap whose remove/add failed leaves e.active pointing at a peer with no references,
// and this retries the hand-off on the next Decrement instead of getting stuck.
if e.active != "" && e.peers[e.active] > 0 {
return Ref[string]{Count: e.total, Out: e.active}, nil
}
// Detach the stale/gone active peer from WireGuard before reprogramming.
if e.active != "" {
if err := rm.remove(prefix, e.active); err != nil {
return Ref[string]{Count: e.total, Out: e.active}, fmt.Errorf("remove allowed IP %v for peer %s: %w", prefix, e.active, err)
}
e.active = ""
}
// Hand the prefix over to a surviving peer, or drop the entry when none remain.
if survivor, ok := pickSurvivor(e.peers); ok {
out, err := rm.add(prefix, survivor)
if err != nil {
return Ref[string]{Count: e.total, Out: ""}, fmt.Errorf("swap allowed IP %v to peer %s: %w", prefix, survivor, err)
}
e.active = out
return Ref[string]{Count: e.total, Out: e.active}, nil
}
delete(rm.entries, prefix)
return Ref[string]{Count: 0, Out: ""}, nil
}
// Flush removes all prefixes from WireGuard and clears the counter.
func (rm *AllowedIPsRefCounter) Flush() error {
rm.mu.Lock()
defer rm.mu.Unlock()
var merr *multierror.Error
for prefix, e := range rm.entries {
if e.active == "" {
continue
}
logCallerF("Flushing allowed IP for prefix %v peer %s", prefix, e.active)
if err := rm.remove(prefix, e.active); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove allowed IP %v for peer %s: %w", prefix, e.active, err))
}
}
clear(rm.entries)
return nberrors.FormatErrorOrNil(merr)
}
// ReapplyMatching calls apply for every prefix whose currently installed (active) peer satisfies
// pred, holding the lock for the whole pass. It is used to re-push allowed IPs onto a peer whose
// WireGuard entry was rebuilt (e.g. a lazy connection cycling idle->wake) without a matching
// refcounter change, which would otherwise leave the prefix installed in the counter but missing
// on the device. Only the active peer is considered — a prefix that lost its installed peer to a
// failed swap is skipped here and reconciled by the next Increment/Decrement.
func (rm *AllowedIPsRefCounter) ReapplyMatching(pred func(out string) bool, apply func(key netip.Prefix) error) error {
rm.mu.Lock()
defer rm.mu.Unlock()
var merr *multierror.Error
for prefix, e := range rm.entries {
if e.active != "" && pred(e.active) {
if err := apply(prefix); err != nil {
merr = multierror.Append(merr, err)
}
}
}
return nberrors.FormatErrorOrNil(merr)
}
// pickSurvivor deterministically selects a peer still referencing the prefix. WireGuard cannot do
// multipath for a single prefix, so any surviving peer is a valid winner; the choice is made stable
// (lowest peerKey) for predictable behavior and testability.
func pickSurvivor(peers map[string]int) (string, bool) {
if len(peers) == 0 {
return "", false
}
keys := make([]string, 0, len(peers))
for k := range peers {
keys = append(keys, k)
}
sort.Strings(keys)
return keys[0], true
}
@@ -0,0 +1,241 @@
package refcounter
import (
"errors"
"net/netip"
"testing"
)
// fakeWG models WireGuard's cryptokey routing: a prefix can be installed on exactly one peer.
// failAdd/failRemove make the next add/remove fail once, to exercise the self-healing error paths.
type fakeWG struct {
installed map[netip.Prefix]string
adds int
removes int
failAdd bool
failRemove bool
}
func newFakeWG() *fakeWG {
return &fakeWG{installed: map[netip.Prefix]string{}}
}
func (f *fakeWG) counter() *AllowedIPsRefCounter {
return NewAllowedIPs(
func(prefix netip.Prefix, peerKey string) (string, error) {
if f.failAdd {
f.failAdd = false
return "", errors.New("add failed")
}
f.adds++
f.installed[prefix] = peerKey
return peerKey, nil
},
func(prefix netip.Prefix, peerKey string) error {
if f.failRemove {
f.failRemove = false
return errors.New("remove failed")
}
f.removes++
// only clear if this peer is the one installed, mirroring wg semantics
if f.installed[prefix] == peerKey {
delete(f.installed, prefix)
}
return nil
},
)
}
func mustPrefix(t *testing.T, s string) netip.Prefix {
t.Helper()
p, err := netip.ParsePrefix(s)
if err != nil {
t.Fatalf("parse prefix %q: %v", s, err)
}
return p
}
func mustIncrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] {
t.Helper()
ref, err := c.Increment(p, peer)
if err != nil {
t.Fatalf("Increment(%v, %s): %v", p, peer, err)
}
return ref
}
func mustDecrement(t *testing.T, c *AllowedIPsRefCounter, p netip.Prefix, peer string) Ref[string] {
t.Helper()
ref, err := c.Decrement(p, peer)
if err != nil {
t.Fatalf("Decrement(%v, %s): %v", p, peer, err)
}
return ref
}
// TestAllowedIPs_SwapOnActivePeerRemoval reproduces the reported bug: two networks with the same
// prefix routed by different peers. Removing the network whose peer is installed must hand the
// prefix over to the surviving peer instead of leaving it on the removed one.
func TestAllowedIPs_SwapOnActivePeerRemoval(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerB")
// First peer wins while both are present.
if got := f.installed[p]; got != "peerA" {
t.Fatalf("expected peerA installed, got %q", got)
}
// Remove the active peer's network -> must swap to peerB.
mustDecrement(t, c, p, "peerA")
if got := f.installed[p]; got != "peerB" {
t.Fatalf("BUG: prefix stuck on removed peer, want peerB got %q", got)
}
// Remove the last one -> prefix gone.
mustDecrement(t, c, p, "peerB")
if _, ok := f.installed[p]; ok {
t.Fatalf("expected prefix removed, still installed on %q", f.installed[p])
}
}
// TestAllowedIPs_RemoveNonActivePeer removing a non-installed peer must not touch WireGuard.
func TestAllowedIPs_RemoveNonActivePeer(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerB")
removesBefore := f.removes
mustDecrement(t, c, p, "peerB")
if f.installed[p] != "peerA" {
t.Fatalf("active peer must stay peerA, got %q", f.installed[p])
}
if f.removes != removesBefore {
t.Fatalf("removing a non-active peer must not call wg remove")
}
}
// TestAllowedIPs_SamePeerMultipleRefs two references via the same peer must keep the prefix until
// the last reference is released (the reason the per-peer count must be an int, not a set).
func TestAllowedIPs_SamePeerMultipleRefs(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerA")
if f.adds != 1 {
t.Fatalf("expected a single wg add for the same peer, got %d", f.adds)
}
mustDecrement(t, c, p, "peerA")
if f.installed[p] != "peerA" {
t.Fatalf("prefix must stay while a reference remains, got %q", f.installed[p])
}
if f.removes != 0 {
t.Fatalf("no wg remove expected while a reference remains, got %d", f.removes)
}
mustDecrement(t, c, p, "peerA")
if _, ok := f.installed[p]; ok {
t.Fatalf("prefix must be removed after last reference")
}
}
// TestAllowedIPs_RefCountAndActive checks the Ref returned to callers (used for the HA-disabled log).
func TestAllowedIPs_RefCountAndActive(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
ref := mustIncrement(t, c, p, "peerA")
if ref.Count != 1 || ref.Out != "peerA" {
t.Fatalf("want {1, peerA}, got {%d, %q}", ref.Count, ref.Out)
}
ref = mustIncrement(t, c, p, "peerB")
if ref.Count != 2 || ref.Out != "peerA" {
t.Fatalf("want {2, peerA}, got {%d, %q}", ref.Count, ref.Out)
}
}
// TestAllowedIPs_Flush removes everything installed and clears the counter.
func TestAllowedIPs_Flush(t *testing.T) {
f := newFakeWG()
c := f.counter()
p1 := mustPrefix(t, "10.44.8.0/24")
p2 := mustPrefix(t, "10.44.9.0/24")
mustIncrement(t, c, p1, "peerA")
mustIncrement(t, c, p2, "peerB")
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(f.installed) != 0 {
t.Fatalf("expected all prefixes removed, got %v", f.installed)
}
// After flush, a fresh increment must add again.
mustIncrement(t, c, p1, "peerC")
if f.installed[p1] != "peerC" {
t.Fatalf("counter not reset after flush")
}
}
// TestAllowedIPs_SelfHealAfterSwapAddError ensures a failed add during a swap does not permanently
// strand the prefix: the next Decrement (or Increment) must retry and install a surviving peer.
func TestAllowedIPs_SelfHealAfterSwapAddError(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerB")
mustIncrement(t, c, p, "peerC")
// Removing the active peerA triggers a swap to a survivor; make the add fail once.
f.failAdd = true
if _, err := c.Decrement(p, "peerA"); err == nil {
t.Fatalf("expected error from failed swap add")
}
if _, ok := f.installed[p]; ok {
t.Fatalf("nothing should be installed after a failed swap add, got %q", f.installed[p])
}
// A later Decrement of a non-active survivor must retry the hand-off (self-heal), not stay stuck.
ref := mustDecrement(t, c, p, "peerC")
if got := f.installed[p]; got == "" {
t.Fatalf("self-heal failed: prefix left unrouted after add recovered")
}
if ref.Out == "" {
t.Fatalf("expected an active peer after self-heal, got empty")
}
}
// TestAllowedIPs_SelfHealAfterRemoveError ensures a failed remove during a swap is retried instead
// of leaving e.active stuck on a peer that no longer holds references.
func TestAllowedIPs_SelfHealAfterRemoveError(t *testing.T) {
f := newFakeWG()
c := f.counter()
p := mustPrefix(t, "10.44.8.0/24")
mustIncrement(t, c, p, "peerA")
mustIncrement(t, c, p, "peerB")
// Releasing active peerA must detach it (remove) then add peerB; fail the remove once.
f.failRemove = true
if _, err := c.Decrement(p, "peerA"); err == nil {
t.Fatalf("expected error from failed remove")
}
// Next Decrement of the non-active survivor retries: removes stale peerA, installs peerB.
mustDecrement(t, c, p, "peerB")
// peerB had only one ref, so after retry the prefix is fully released.
if _, ok := f.installed[p]; ok {
t.Fatalf("expected prefix released after self-heal, still on %q", f.installed[p])
}
}
@@ -94,6 +94,26 @@ func (rm *Counter[Key, I, O]) Get(key Key) (Ref[O], bool) {
return ref, ok
}
// ReapplyMatching calls apply for every key whose stored Out satisfies pred, holding the
// counter lock for the whole pass. Running apply under the lock keeps it atomic with respect
// to Increment/Decrement: a prefix dropped to zero is removed from the map (and had its
// RemoveFunc called) before this pass observes it, so a stale key can never be re-applied.
// pred and apply are invoked under the lock, so they must not call back into the counter.
func (rm *Counter[Key, I, O]) ReapplyMatching(pred func(out O) bool, apply func(key Key) error) error {
rm.mu.Lock()
defer rm.mu.Unlock()
var merr *multierror.Error
for key, ref := range rm.refCountMap {
if pred(ref.Out) {
if err := apply(key); err != nil {
merr = multierror.Append(merr, err)
}
}
}
return nberrors.FormatErrorOrNil(merr)
}
// Increment increments the reference count for the given key.
// If this is the first reference to the key, the AddFunc is called.
func (rm *Counter[Key, I, O]) Increment(key Key, in I) (Ref[O], error) {
@@ -0,0 +1,47 @@
package refcounter
import (
"net/netip"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// TestReapplyMatching verifies ReapplyMatching invokes apply for exactly the keys whose stored
// Out satisfies the predicate (no duplicates for multiply-referenced keys) — the primitive
// ReconcilePeerAllowedIPs relies on to re-apply a single peer's routed prefixes.
func TestReapplyMatching(t *testing.T) {
rc := New[netip.Prefix, string, string](
func(_ netip.Prefix, peerKey string) (string, error) { return peerKey, nil },
func(netip.Prefix, string) error { return nil },
)
peerA1 := netip.MustParsePrefix("10.0.0.0/24")
peerA2 := netip.MustParsePrefix("10.1.0.0/24")
peerB1 := netip.MustParsePrefix("10.2.0.0/24")
for prefix, peer := range map[netip.Prefix]string{peerA1: "peerA", peerA2: "peerA", peerB1: "peerB"} {
_, err := rc.Increment(prefix, peer)
require.NoError(t, err)
}
// a second reference must not make the key applied twice
_, err := rc.Increment(peerA1, "peerA")
require.NoError(t, err)
var applied []netip.Prefix
err = rc.ReapplyMatching(
func(out string) bool { return out == "peerA" },
func(key netip.Prefix) error { applied = append(applied, key); return nil },
)
require.NoError(t, err)
assert.ElementsMatch(t, []netip.Prefix{peerA1, peerA2}, applied)
var none []netip.Prefix
err = rc.ReapplyMatching(
func(out string) bool { return out == "missing" },
func(key netip.Prefix) error { none = append(none, key); return nil },
)
require.NoError(t, err)
assert.Empty(t, none)
}
@@ -5,5 +5,7 @@ import "net/netip"
// RouteRefCounter is a Counter for Route, it doesn't take any input on Increment and doesn't use any output on Decrement
type RouteRefCounter = Counter[netip.Prefix, struct{}, struct{}]
// AllowedIPsRefCounter is a Counter for AllowedIPs, it takes a peer key on Increment and passes it back to Decrement
type AllowedIPsRefCounter = Counter[netip.Prefix, string, string]
// AllowedIPsRefCounter tracks WireGuard AllowedIPs per prefix. Unlike the generic Counter it is peer-aware:
// a prefix can be claimed by several peers at once and WireGuard allows a given prefix on exactly one peer,
// so the counter records the per-peer reference count and swaps the installed peer when the active one is released.
// See allowedips.go.
+138
View File
@@ -0,0 +1,138 @@
package routemanager
import (
"fmt"
"slices"
"github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus"
"golang.org/x/exp/maps"
nberrors "github.com/netbirdio/netbird/client/errors"
"github.com/netbirdio/netbird/route"
)
// SelectRoutes selects the routes with the given network IDs and applies the
// new selection. V4/v6 exit-node pairs are expanded automatically. Exit nodes
// are mutually exclusive: if the selection activates an exit node, every other
// available exit node is deselected so two can't be active at once. With
// appendRoute=false the previous selection is replaced instead of extended.
func (m *DefaultManager) SelectRoutes(ids []route.NetID, appendRoute bool) error {
if err := m.selectRoutes(ids, appendRoute); err != nil {
return err
}
m.TriggerSelection(m.GetClientRoutes())
return nil
}
// DeselectRoutes removes the routes with the given network IDs from the
// selection and applies the change. V4/v6 exit-node pairs are expanded
// automatically.
func (m *DefaultManager) DeselectRoutes(ids []route.NetID) error {
if err := m.deselectRoutes(ids); err != nil {
return err
}
m.TriggerSelection(m.GetClientRoutes())
return nil
}
func (m *DefaultManager) deselectRoutes(ids []route.NetID) error {
routesMap := m.GetClientRoutesWithNetID()
routes := route.ExpandV6ExitPairs(slices.Clone(ids), routesMap)
log.Debugf("deselecting routes with ids: %v", routes)
if err := m.routeSelector.DeselectRoutes(routes, maps.Keys(routesMap)); err != nil {
return fmt.Errorf("deselect routes: %w", err)
}
return nil
}
// SelectAllRoutes selects every available route and applies the selection.
// Exit nodes stay mutually exclusive: at most one remains active.
func (m *DefaultManager) SelectAllRoutes() {
m.selectAllRoutes()
m.TriggerSelection(m.GetClientRoutes())
}
func (m *DefaultManager) selectAllRoutes() {
m.routeSelector.SelectAllRoutes()
// Select-all wipes every explicit selection, so exit nodes fall back to
// management's auto-apply flags — which may mark several at once.
// Reconcile immediately so at most one exit node stays active instead of
// waiting for the next network map to enforce it.
m.mux.Lock()
defer m.mux.Unlock()
m.updateRouteSelectorFromManagement(m.clientRoutes)
}
// DeselectAllRoutes deselects every route and applies the change.
func (m *DefaultManager) DeselectAllRoutes() {
m.routeSelector.DeselectAllRoutes()
m.TriggerSelection(m.GetClientRoutes())
}
func (m *DefaultManager) selectRoutes(ids []route.NetID, appendRoute bool) error {
routesMap := m.GetClientRoutesWithNetID()
routes := route.ExpandV6ExitPairs(slices.Clone(ids), routesMap)
allIDs := maps.Keys(routesMap)
log.Debugf("selecting routes with ids: %v", routes)
// A partial failure (e.g. an unknown ID in the request) still selects the
// valid routes, so exclusivity below must run regardless of the error.
var merr *multierror.Error
if err := m.routeSelector.SelectRoutes(routes, appendRoute, allIDs); err != nil {
merr = multierror.Append(merr, fmt.Errorf("select routes: %w", err))
}
// Exit nodes are mutually exclusive: if this selection activates an
// exit node, deselect every other available exit node so two can't be
// selected at once. Non-exit route selections are left untouched.
if requestActivatesExitNode(routes, routesMap) {
if others := otherExitNodeIDs(routesMap, routes); len(others) > 0 {
if err := m.routeSelector.DeselectRoutes(others, allIDs); err != nil {
merr = multierror.Append(merr, fmt.Errorf("deselect sibling exit nodes: %w", err))
}
}
}
return nberrors.FormatErrorOrNil(merr)
}
func isExitNodeRoutes(routes []*route.Route) bool {
return len(routes) > 0 && (route.IsV4DefaultRoute(routes[0].Network) || route.IsV6DefaultRoute(routes[0].Network))
}
// requestActivatesExitNode reports whether any requested NetID maps to an exit
// node (default route) in the current route table.
func requestActivatesExitNode(requested []route.NetID, routesMap map[route.NetID][]*route.Route) bool {
for _, id := range requested {
if isExitNodeRoutes(routesMap[id]) {
return true
}
}
return false
}
// otherExitNodeIDs returns every available exit-node NetID that is not in the
// requested set — the siblings to deselect so a single exit node stays active.
func otherExitNodeIDs(routesMap map[route.NetID][]*route.Route, requested []route.NetID) []route.NetID {
keep := make(map[route.NetID]struct{}, len(requested))
for _, id := range requested {
keep[id] = struct{}{}
}
var others []route.NetID
for id, routes := range routesMap {
if !isExitNodeRoutes(routes) {
continue
}
if _, ok := keep[id]; ok {
continue
}
others = append(others, id)
}
return others
}
@@ -0,0 +1,129 @@
package routemanager
import (
"net/netip"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/internal/routeselector"
"github.com/netbirdio/netbird/route"
)
func v6ExitRoute(netID, peer string) *route.Route {
return &route.Route{
NetID: route.NetID(netID),
Network: netip.MustParsePrefix("::/0"),
Peer: peer,
}
}
func newSelectionTestManager() *DefaultManager {
return &DefaultManager{
routeSelector: routeselector.NewRouteSelector(),
clientRoutes: route.HAMap{
"exitA|0.0.0.0/0": {exitRoute("exitA", "p1", true)},
"exitA-v6|::/0": {v6ExitRoute("exitA-v6", "p1")},
"exitB|0.0.0.0/0": {exitRoute("exitB", "p2", true)},
"lan|192.168.1.0/24": {{NetID: "lan", Network: netip.MustParsePrefix("192.168.1.0/24"), Peer: "p3"}},
},
}
}
func TestSelectRoutes_ExitNodeExclusivity(t *testing.T) {
m := newSelectionTestManager()
// Selecting an exit node selects its v6 pair and deselects the sibling.
require.NoError(t, m.selectRoutes([]route.NetID{"exitA"}, true))
assert.True(t, m.routeSelector.IsSelected("exitA"), "exitA should be selected")
assert.True(t, m.routeSelector.IsSelected("exitA-v6"), "the v6 pair follows its v4 base")
assert.False(t, m.routeSelector.IsSelected("exitB"), "the sibling exit node must be deselected")
// Switching to the sibling deselects the previous exit node and its v6 pair.
require.NoError(t, m.selectRoutes([]route.NetID{"exitB"}, true))
assert.True(t, m.routeSelector.IsSelected("exitB"), "exitB should now be selected")
assert.False(t, m.routeSelector.IsSelected("exitA"), "the previous exit node must be deselected")
assert.False(t, m.routeSelector.IsSelected("exitA-v6"), "the previous exit node's v6 pair must be deselected")
assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit route selection is untouched")
// Selecting a non-exit route leaves the active exit node alone.
require.NoError(t, m.selectRoutes([]route.NetID{"lan"}, true))
assert.True(t, m.routeSelector.IsSelected("exitB"), "selecting a non-exit route keeps the exit node")
// Deselecting the active exit node turns every exit node off.
require.NoError(t, m.deselectRoutes([]route.NetID{"exitB"}))
assert.False(t, m.routeSelector.IsSelected("exitB"), "exitB should be deselected")
assert.False(t, m.routeSelector.IsSelected("exitA"), "exitA stays deselected")
assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit route selection is untouched")
}
func TestSelectRoutes_PartialErrorStillEnforcesExclusivity(t *testing.T) {
// The unknown ID must be reported, but the valid exit node in the same
// request is still selected — so its sibling must still be deselected.
// Both orderings are covered: processing must continue past the invalid
// ID wherever it sits in the request.
requests := map[string][]route.NetID{
"invalid id first": {"missing", "exitB"},
"invalid id last": {"exitB", "missing"},
}
for name, ids := range requests {
t.Run(name, func(t *testing.T) {
m := newSelectionTestManager()
require.NoError(t, m.selectRoutes([]route.NetID{"exitA"}, true))
err := m.selectRoutes(ids, true)
assert.Error(t, err, "unknown id must be reported")
assert.True(t, m.routeSelector.IsSelected("exitB"), "valid exit node from the request is selected")
assert.False(t, m.routeSelector.IsSelected("exitA"), "sibling exit node must be deselected despite the error")
assert.False(t, m.routeSelector.IsSelected("exitA-v6"), "sibling's v6 pair must be deselected too")
})
}
}
func TestSelectAllRoutes_KeepsSingleExitNode(t *testing.T) {
// Both exit nodes are marked for auto-apply by management
// (SkipAutoApply=false), the state where select-all could turn on two at
// once without the immediate reconciliation.
m := &DefaultManager{
routeSelector: routeselector.NewRouteSelector(),
clientRoutes: route.HAMap{
"exitA|0.0.0.0/0": {exitRoute("exitA", "p1", false)},
"exitB|0.0.0.0/0": {exitRoute("exitB", "p2", false)},
"lan|192.168.1.0/24": {{NetID: "lan", Network: netip.MustParsePrefix("192.168.1.0/24"), Peer: "p3"}},
},
}
require.NoError(t, m.selectRoutes([]route.NetID{"exitB"}, true))
m.selectAllRoutes()
assert.True(t, m.routeSelector.IsSelected("lan"), "non-exit routes are all selected")
assert.True(t, m.routeSelector.IsSelected("exitA"), "the deterministic management pick stays active")
assert.False(t, m.routeSelector.IsSelected("exitB"), "select-all must not leave a second exit node active")
}
func TestSelectRoutes_UnknownRoute(t *testing.T) {
m := newSelectionTestManager()
assert.Error(t, m.selectRoutes([]route.NetID{"missing"}, true), "selecting an unavailable route must fail")
assert.Error(t, m.deselectRoutes([]route.NetID{"missing"}), "deselecting an unavailable route must fail")
}
func TestExitNodeSelectionHelpers(t *testing.T) {
routesMap := map[route.NetID][]*route.Route{
"exitA": {{Network: netip.MustParsePrefix("0.0.0.0/0")}},
"exitB": {{Network: netip.MustParsePrefix("::/0")}},
"lan": {{Network: netip.MustParsePrefix("192.168.0.0/16")}},
}
assert.True(t, requestActivatesExitNode([]route.NetID{"exitA"}, routesMap), "v4 default route is an exit node")
assert.True(t, requestActivatesExitNode([]route.NetID{"exitB"}, routesMap), "v6 default route is an exit node")
assert.False(t, requestActivatesExitNode([]route.NetID{"lan"}, routesMap), "lan route is not an exit node")
assert.False(t, requestActivatesExitNode([]route.NetID{"missing"}, routesMap), "unknown id is not an exit node")
others := otherExitNodeIDs(routesMap, []route.NetID{"exitB"})
assert.ElementsMatch(t, []route.NetID{"exitA"}, others, "only the other exit node is a sibling; the lan route is ignored")
}
+11 -3
View File
@@ -15,6 +15,11 @@ type Route struct {
route *route.Route
routeRefCounter *refcounter.RouteRefCounter
allowedIPsRefcounter *refcounter.AllowedIPsRefCounter
// currentPeerKey is the routing peer this watcher currently has the prefix installed on
// (the HA winner elected by the watcher). It can differ from route.Peer and change on
// failover, so it is recorded on AddAllowedIPs and used on RemoveAllowedIPs to decrement
// the exact peer that was incremented.
currentPeerKey string
}
func NewRoute(params common.HandlerParams) *Route {
@@ -52,12 +57,15 @@ func (r *Route) AddAllowedIPs(peerKey string) error {
ref.Out,
)
}
r.currentPeerKey = peerKey
return nil
}
func (r *Route) RemoveAllowedIPs() error {
if _, err := r.allowedIPsRefcounter.Decrement(r.route.Network); err != nil {
return err
var err error
if _, decErr := r.allowedIPsRefcounter.Decrement(r.route.Network, r.currentPeerKey); decErr != nil {
err = fmt.Errorf("remove allowed IP %s: %w", r.route.Network, decErr)
}
return nil
r.currentPeerKey = ""
return err
}
@@ -20,6 +20,8 @@ const (
rpFilterPath = "net.ipv4.conf.all.rp_filter"
rpFilterInterfacePath = "net.ipv4.conf.%s.rp_filter"
srcValidMarkPath = "net.ipv4.conf.all.src_valid_mark"
percentEscape = "%25"
dotEscape = "%2E"
)
type iface interface {
@@ -56,7 +58,11 @@ func Setup(wgIface iface) (map[string]int, error) {
continue
}
i := fmt.Sprintf(rpFilterInterfacePath, intf.Name)
// Escape '%' and '.' so they survive the dot-to-slash conversion in Set()
safeName := strings.ReplaceAll(intf.Name, "%", percentEscape)
safeName = strings.ReplaceAll(safeName, ".", dotEscape)
i := fmt.Sprintf(rpFilterInterfacePath, safeName)
oldVal, err := Set(i, 2, true)
if err != nil {
result = multierror.Append(result, err)
@@ -70,7 +76,11 @@ func Setup(wgIface iface) (map[string]int, error) {
// Set sets a sysctl configuration, if onlyIfOne is true it will only set the new value if it's set to 1
func Set(key string, desiredValue int, onlyIfOne bool) (int, error) {
path := fmt.Sprintf("/proc/sys/%s", strings.ReplaceAll(key, ".", "/"))
path := strings.ReplaceAll(key, ".", "/")
// Unescape interface dots and percent signs
path = strings.ReplaceAll(path, dotEscape, ".")
path = strings.ReplaceAll(path, percentEscape, "%")
path = fmt.Sprintf("/proc/sys/%s", path)
currentValue, err := os.ReadFile(path)
if err != nil {
return -1, fmt.Errorf("read sysctl %s: %w", key, err)
+6
View File
@@ -1,6 +1,7 @@
package statemanager
import (
"bytes"
"context"
"encoding/json"
"errors"
@@ -305,6 +306,11 @@ func (m *Manager) loadStateFile(deleteCorrupt bool) (map[string]json.RawMessage,
var rawStates map[string]json.RawMessage
if err := json.Unmarshal(data, &rawStates); err != nil {
if len(bytes.TrimSpace(data)) == 0 {
log.Warnf("state file %s is empty (%d bytes)", m.filePath, len(data))
} else {
log.Warnf("state file %s has malformed content (%d bytes)", m.filePath, len(data))
}
m.handleCorruptedState(deleteCorrupt)
return nil, fmt.Errorf("unmarshal states: %w", err)
}
+124
View File
@@ -0,0 +1,124 @@
package tunnelnotifier
import (
"container/list"
"sync"
"github.com/netbirdio/netbird/client/internal/dns"
"github.com/netbirdio/netbird/client/internal/listener"
)
type eventKind int
const (
eventRoutes eventKind = iota
eventIfaceIP
eventIfaceIPv6
eventDNS
)
var (
_ listener.NetworkChangeListener = (*Notifier)(nil)
_ dns.IosDnsManager = (*Notifier)(nil)
)
type event struct {
kind eventKind
payload string
}
type Notifier struct {
mu sync.Mutex
cond *sync.Cond
queue *list.List
closed bool
done chan struct{}
listener listener.NetworkChangeListener
dnsManager dns.IosDnsManager
}
func New(l listener.NetworkChangeListener, dm dns.IosDnsManager) *Notifier {
n := &Notifier{
queue: list.New(),
done: make(chan struct{}),
listener: l,
dnsManager: dm,
}
n.cond = sync.NewCond(&n.mu)
go n.deliverLoop()
return n
}
func (n *Notifier) OnNetworkChanged(routes string) {
n.enqueue(event{kind: eventRoutes, payload: routes})
}
func (n *Notifier) SetInterfaceIP(ip string) {
n.enqueue(event{kind: eventIfaceIP, payload: ip})
}
func (n *Notifier) SetInterfaceIPv6(ip string) {
n.enqueue(event{kind: eventIfaceIPv6, payload: ip})
}
func (n *Notifier) ApplyDns(config string) {
n.enqueue(event{kind: eventDNS, payload: config})
}
// Close stops accepting new events and blocks until the delivery loop has
// drained all queued events and exited.
func (n *Notifier) Close() {
n.mu.Lock()
n.closed = true
n.cond.Signal()
n.mu.Unlock()
<-n.done
}
func (n *Notifier) enqueue(ev event) {
n.mu.Lock()
defer n.mu.Unlock()
if n.closed {
return
}
n.queue.PushBack(ev)
n.cond.Signal()
}
func (n *Notifier) deliverLoop() {
defer close(n.done)
for {
n.mu.Lock()
for n.queue.Len() == 0 && !n.closed {
n.cond.Wait()
}
if n.closed && n.queue.Len() == 0 {
n.mu.Unlock()
return
}
ev := n.queue.Remove(n.queue.Front()).(event)
l := n.listener
dm := n.dnsManager
n.mu.Unlock()
switch ev.kind {
case eventRoutes:
if l != nil {
l.OnNetworkChanged(ev.payload)
}
case eventIfaceIP:
if l != nil {
l.SetInterfaceIP(ev.payload)
}
case eventIfaceIPv6:
if l != nil {
l.SetInterfaceIPv6(ev.payload)
}
case eventDNS:
if dm != nil {
dm.ApplyDns(ev.payload)
}
}
}
}
@@ -0,0 +1,192 @@
package tunnelnotifier
import (
"fmt"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type call struct {
kind string
payload string
}
type recorder struct {
mu sync.Mutex
calls []call
inFlight atomic.Int32
overlap atomic.Bool
delay time.Duration
}
func (r *recorder) record(kind, payload string) {
if r.inFlight.Add(1) != 1 {
r.overlap.Store(true)
}
if r.delay > 0 {
time.Sleep(r.delay)
}
r.mu.Lock()
r.calls = append(r.calls, call{kind: kind, payload: payload})
r.mu.Unlock()
r.inFlight.Add(-1)
}
func (r *recorder) count() int {
r.mu.Lock()
defer r.mu.Unlock()
return len(r.calls)
}
func (r *recorder) snapshot() []call {
r.mu.Lock()
defer r.mu.Unlock()
out := make([]call, len(r.calls))
copy(out, r.calls)
return out
}
type fakeListener struct {
rec *recorder
}
func (f *fakeListener) OnNetworkChanged(routes string) {
f.rec.record("routes", routes)
}
func (f *fakeListener) SetInterfaceIP(ip string) {
f.rec.record("ip", ip)
}
func (f *fakeListener) SetInterfaceIPv6(ip string) {
f.rec.record("ipv6", ip)
}
type fakeDNSManager struct {
rec *recorder
}
func (f *fakeDNSManager) ApplyDns(config string) {
f.rec.record("dns", config)
}
func TestFIFOOrder(t *testing.T) {
rec := &recorder{}
n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec})
defer n.Close()
n.SetInterfaceIP("10.0.0.1")
n.SetInterfaceIPv6("fd00::1")
n.ApplyDns(`{"domains":[]}`)
n.OnNetworkChanged("10.0.0.0/8,192.168.0.0/16")
n.ApplyDns(`{"domains":["example.com"]}`)
require.Eventually(t, func() bool { return rec.count() == 5 }, time.Second, time.Millisecond)
expected := []call{
{kind: "ip", payload: "10.0.0.1"},
{kind: "ipv6", payload: "fd00::1"},
{kind: "dns", payload: `{"domains":[]}`},
{kind: "routes", payload: "10.0.0.0/8,192.168.0.0/16"},
{kind: "dns", payload: `{"domains":["example.com"]}`},
}
assert.Equal(t, expected, rec.snapshot())
}
func TestNoOverlappingCalls(t *testing.T) {
rec := &recorder{delay: 100 * time.Microsecond}
n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec})
defer n.Close()
const producers = 8
const perProducer = 25
var wg sync.WaitGroup
for i := 0; i < producers; i++ {
wg.Add(1)
go func(id int) {
defer wg.Done()
for j := 0; j < perProducer; j++ {
payload := fmt.Sprintf("%d-%d", id, j)
switch j % 4 {
case 0:
n.OnNetworkChanged(payload)
case 1:
n.SetInterfaceIP(payload)
case 2:
n.SetInterfaceIPv6(payload)
case 3:
n.ApplyDns(payload)
}
}
}(i)
}
wg.Wait()
require.Eventually(t, func() bool { return rec.count() == producers*perProducer }, 5*time.Second, time.Millisecond)
assert.False(t, rec.overlap.Load())
}
func TestDNSAndRoutesInterleaved(t *testing.T) {
rec := &recorder{delay: 100 * time.Microsecond}
n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec})
defer n.Close()
const events = 50
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
for i := 0; i < events; i++ {
n.ApplyDns(fmt.Sprintf("dns-%d", i))
}
}()
go func() {
defer wg.Done()
for i := 0; i < events; i++ {
n.OnNetworkChanged(fmt.Sprintf("routes-%d", i))
}
}()
wg.Wait()
require.Eventually(t, func() bool { return rec.count() == 2*events }, 5*time.Second, time.Millisecond)
assert.False(t, rec.overlap.Load())
var dnsSeen, routesSeen int
for _, c := range rec.snapshot() {
switch c.kind {
case "dns":
assert.Equal(t, fmt.Sprintf("dns-%d", dnsSeen), c.payload)
dnsSeen++
case "routes":
assert.Equal(t, fmt.Sprintf("routes-%d", routesSeen), c.payload)
routesSeen++
}
}
assert.Equal(t, events, dnsSeen)
assert.Equal(t, events, routesSeen)
}
func TestCloseDrainsQueue(t *testing.T) {
rec := &recorder{delay: time.Millisecond}
n := New(&fakeListener{rec: rec}, &fakeDNSManager{rec: rec})
const events = 20
for i := 0; i < events; i++ {
n.OnNetworkChanged(fmt.Sprintf("routes-%d", i))
}
n.Close()
require.Equal(t, events, rec.count(), "Close must not return before all queued events are delivered")
n.OnNetworkChanged("after-close")
n.ApplyDns("after-close")
time.Sleep(50 * time.Millisecond)
assert.Equal(t, events, rec.count())
}