Merge main into poc/certificate-posture

This commit is contained in:
Viktor Liu
2026-10-05 19:09:43 +02:00
390 changed files with 28443 additions and 14545 deletions
+67 -3
View File
@@ -90,8 +90,9 @@ type StatusRecorder interface {
// fallback T-FinalWarningLead dialog (suppressed when the user dismissed
// the first one for the same deadline). Safe for concurrent use.
type Watcher struct {
lead time.Duration
finalLead time.Duration
lead time.Duration
finalLead time.Duration
deadlineOnly bool
mu sync.Mutex
current time.Time
@@ -102,6 +103,7 @@ type Watcher struct {
dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal
closed bool
recorder StatusRecorder
nowFn func() time.Time
}
// New returns a watcher with the package defaults WarningLead and
@@ -122,9 +124,17 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher {
lead: lead,
finalLead: final,
recorder: recorder,
nowFn: time.Now,
}
}
// NewDeadlineOnly returns a watcher that validates and records deadlines but arms no warning timers.
func NewDeadlineOnly(recorder StatusRecorder) *Watcher {
w := New(recorder)
w.deadlineOnly = true
return w
}
// Update sets the latest deadline. Pass the zero time to clear (e.g. when
// a Sync push from the server omits the field because login expiration
// was disabled).
@@ -181,7 +191,7 @@ func (w *Watcher) Update(deadline time.Time) error {
w.finalFiredAt = time.Time{}
w.dismissedAt = time.Time{}
if deadline.After(now) {
if deadline.After(now) && !w.deadlineOnly {
w.armTimerLocked(deadline)
}
recorder := w.recorder
@@ -303,6 +313,11 @@ func (w *Watcher) fire(armedFor time.Time) {
w.mu.Unlock()
return
}
now := w.nowFn()
if isLate(now, armedFor, max(w.finalLead, 0)) {
w.fireLateLocked(armedFor, now)
return
}
w.firedAt = armedFor
recorder := w.recorder
w.mu.Unlock()
@@ -331,6 +346,14 @@ func (w *Watcher) fireFinal(armedFor time.Time) {
log.Infof("auth session final-warning skipped (dismissed by user)")
return
}
now := w.nowFn()
if isLate(now, armedFor, 0) {
w.finalFiredAt = armedFor
w.mu.Unlock()
log.Infof("auth session final-warning skipped for deadline %s (passed %s ago)",
armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second))
return
}
w.finalFiredAt = armedFor
recorder := w.recorder
w.mu.Unlock()
@@ -341,6 +364,39 @@ func (w *Watcher) fireFinal(armedFor time.Time) {
publishWarning(recorder, armedFor, true)
}
// fireLateLocked handles a T-WarningLead callback that fired inside the
// final-warning window: it sends the final warning in its place while the
// deadline has not passed and the user has not dismissed it, so a resume
// with time left still warns. The caller must hold w.mu; this helper
// releases it.
func (w *Watcher) fireLateLocked(armedFor, now time.Time) {
w.firedAt = armedFor
switch {
case w.dismissedAt.Equal(armedFor):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped (dismissed by user)")
return
case w.finalFiredAt.Equal(armedFor):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped (final warning already fired)")
return
case isLate(now, armedFor, 0):
w.mu.Unlock()
log.Infof("auth session expiry soon warning skipped for deadline %s (passed %s ago)",
armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second))
return
}
w.finalFiredAt = armedFor
recorder := w.recorder
w.mu.Unlock()
if recorder == nil {
return
}
log.Infof("auth session expiry soon warning fired inside the final-warning window, sending final warning for deadline %s",
armedFor.Format(time.RFC3339))
publishWarning(recorder, armedFor, true)
}
// armOneShotLocked schedules cb at fireAt. When fireAt is already in the
// past it dispatches on the next scheduler tick so a state-change recorder
// notification (invoked after w.mu is released) lands first. Caller must
@@ -380,3 +436,11 @@ func publishWarning(recorder StatusRecorder, deadline time.Time, final bool) {
meta,
)
}
// isLate reports whether the wall clock now has already reached armedFor
// minus cutoffLead. The timers run on the monotonic clock, which can stall
// while the host sleeps, so a timer can fire long after the window it was
// armed for.
func isLate(now, armedFor time.Time, cutoffLead time.Duration) bool {
return !now.Round(0).Before(armedFor.Add(-cutoffLead).Round(0))
}
@@ -527,3 +527,201 @@ func TestDismissBeforeUpdateIsNoop(t *testing.T) {
}
t.Fatalf("final-warning did not publish after no-op pre-Update Dismiss, events=%+v", r.snapshot())
}
func TestIsLate(t *testing.T) {
armedFor := time.Date(2026, 10, 1, 12, 0, 0, 0, time.UTC)
lead := 2 * time.Minute
tests := []struct {
name string
now time.Time
cutoffLead time.Duration
want bool
}{
{"before cutoff", armedFor.Add(-3 * time.Minute), lead, false},
{"at cutoff", armedFor.Add(-lead), lead, true},
{"after cutoff", armedFor.Add(-time.Minute), lead, true},
{"zero lead before deadline", armedFor.Add(-time.Second), 0, false},
{"zero lead at deadline", armedFor, 0, true},
{"zero lead after deadline", armedFor.Add(time.Second), 0, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := isLate(tt.now, armedFor, tt.cutoffLead); got != tt.want {
t.Fatalf("isLate(%s, %s, %s) = %v, want %v", tt.now, armedFor, tt.cutoffLead, got, tt.want)
}
})
}
}
func TestIsLateIgnoresMonotonicReading(t *testing.T) {
now := time.Now()
wallOnly := now.Round(0)
if isLate(now, wallOnly.Add(time.Second), 0) {
t.Fatalf("now with monotonic reading must compare as wall clock before a later wall-only deadline")
}
if !isLate(now, wallOnly, 0) {
t.Fatalf("now with monotonic reading must compare as wall clock at an equal wall-only deadline")
}
}
func TestLateTimerFiring(t *testing.T) {
tests := []struct {
name string
final bool
beforeDl time.Duration
wantWarns int
wantFinals int
}{
{"warning on resume inside window", false, 3 * time.Minute, 1, 0},
{"warning promoted to final inside final window", false, time.Minute, 0, 1},
{"warning skipped past deadline", false, -time.Minute, 0, 0},
{"final on resume before deadline", true, time.Minute, 0, 1},
{"final skipped past deadline", true, -time.Minute, 0, 0},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
// The deadline is an hour out so the real timers never fire
// during the test; the late callback is invoked directly with an
// injected clock that simulates a resume near the deadline.
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-tt.beforeDl) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
if tt.final {
w.fireFinal(d)
} else {
w.fire(d)
}
events := r.snapshot()
if got := countWhere(events, event.isWarning); got != tt.wantWarns {
t.Fatalf("expected %d warning publishes, got %d: %+v", tt.wantWarns, got, events)
}
if got := countWhere(events, event.isFinalWarning); got != tt.wantFinals {
t.Fatalf("expected %d final-warning publishes, got %d: %+v", tt.wantFinals, got, events)
}
})
}
}
func TestPromotedFinalWarningIsNotRepeated(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
d := time.Now().Add(time.Hour).Round(0)
now := d.Add(-time.Minute)
w.nowFn = func() time.Time { return now }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.fire(d)
// The final timer was suspended too, so it fires even later than the
// warning timer, here still just before the deadline.
now = d.Add(-30 * time.Second)
w.fireFinal(d)
events := r.snapshot()
if got := countWhere(events, event.isFinalWarning); got != 1 {
t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events)
}
if got := countWhere(events, event.isWarning); got != 0 {
t.Fatalf("expected no regular warning publish, got %d: %+v", got, events)
}
}
func TestPromotionRespectsDismiss(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-time.Minute) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.Dismiss()
w.fire(d)
events := r.snapshot()
if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 {
t.Fatalf("expected no publish after dismiss, got %d: %+v", got, events)
}
}
func TestPromotionSkippedWhenFinalAlreadyFired(t *testing.T) {
r := &fakeRecorder{}
w := New(r)
defer w.Close()
// Both timers fall in the past after a long suspend and are dispatched
// with a zero delay, so the final callback can run before the warning one.
d := time.Now().Add(time.Hour).Round(0)
w.nowFn = func() time.Time { return d.Add(-time.Minute) }
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
w.fireFinal(d)
w.fire(d)
events := r.snapshot()
if got := countWhere(events, event.isFinalWarning); got != 1 {
t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events)
}
if got := countWhere(events, event.isWarning); got != 0 {
t.Fatalf("expected no regular warning publish, got %d: %+v", got, events)
}
}
func TestDeadlineOnlyRecordsDeadlineWithoutWarnings(t *testing.T) {
r := &fakeRecorder{}
w := NewDeadlineOnly(r)
defer w.Close()
// With the default leads this deadline would otherwise fire both
// timers on the next tick.
d := time.Now().Add(50 * time.Millisecond).Round(0)
if err := w.Update(d); err != nil {
t.Fatalf("Update: %v", err)
}
if got := r.deadline(); !got.Equal(d) {
t.Fatalf("expected recorder deadline %v, got %v", d, got)
}
time.Sleep(100 * time.Millisecond)
events := r.snapshot()
if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 {
t.Fatalf("expected no publish in deadline-only mode, got %d: %+v", got, events)
}
if w.timer != nil || w.finalTimer != nil {
t.Fatal("expected no timers armed in deadline-only mode")
}
}
func TestDeadlineOnlyStillRejectsOutOfRangeDeadlines(t *testing.T) {
r := &fakeRecorder{}
w := NewDeadlineOnly(r)
defer w.Close()
if err := w.Update(time.Now().Add(time.Hour)); err != nil {
t.Fatalf("Update: %v", err)
}
err := w.Update(time.Now().Add(-maxPastHorizon - time.Hour))
if !errors.Is(err, ErrDeadlineInPast) {
t.Fatalf("expected ErrDeadlineInPast, got %v", err)
}
if got := r.deadline(); !got.IsZero() {
t.Fatalf("expected recorder cleared after rejection, got %v", got)
}
}
+42
View File
@@ -0,0 +1,42 @@
package daemonaddr
import (
"os"
"strconv"
log "github.com/sirupsen/logrus"
)
const (
// EnvMaxRecvMsgSize overrides the default gRPC max receive message size for
// connections to the daemon. Value is in bytes.
EnvMaxRecvMsgSize = "NB_DAEMON_GRPC_MAX_MSG_SIZE"
// defaultMaxRecvMsgSize is the max gRPC receive message size used for daemon
// connections when EnvMaxRecvMsgSize is unset or invalid. It overrides the
// gRPC library default of 4 MB, which a detailed status already exceeds on a
// network of a few thousand peers.
defaultMaxRecvMsgSize = 1024 * 1024 * 16
)
// MaxRecvMsgSize returns the max gRPC receive message size for daemon connections
// from the environment, or defaultMaxRecvMsgSize (16 MB) if unset or invalid.
func MaxRecvMsgSize() int {
val := os.Getenv(EnvMaxRecvMsgSize)
if val == "" {
return defaultMaxRecvMsgSize
}
size, err := strconv.Atoi(val)
if err != nil {
log.Warnf("invalid %s value %q, using default: %v", EnvMaxRecvMsgSize, val, err)
return defaultMaxRecvMsgSize
}
if size <= 0 {
log.Warnf("invalid %s value %d, must be positive, using default", EnvMaxRecvMsgSize, size)
return defaultMaxRecvMsgSize
}
return size
}
+112
View File
@@ -0,0 +1,112 @@
package daemonaddr
import (
"context"
"net"
"os"
"strings"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/proto"
)
func TestMaxRecvMsgSize(t *testing.T) {
tests := []struct {
name string
envValue string
expected int
}{
{name: "unset returns default", envValue: "", expected: defaultMaxRecvMsgSize},
{name: "non-numeric returns default", envValue: "abc", expected: defaultMaxRecvMsgSize},
{name: "negative returns default", envValue: "-1", expected: defaultMaxRecvMsgSize},
{name: "zero returns default", envValue: "0", expected: defaultMaxRecvMsgSize},
{name: "valid value is used", envValue: "33554432", expected: 33554432},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
// Set first so the previous value is restored on cleanup, then unset to
// exercise the absent case.
t.Setenv(EnvMaxRecvMsgSize, tc.envValue)
if tc.envValue == "" {
require.NoError(t, os.Unsetenv(EnvMaxRecvMsgSize), "unset the override")
}
assert.Equal(t, tc.expected, MaxRecvMsgSize(), "max receive message size")
})
}
}
// bigStatusServer answers Status with a response larger than gRPC's 4 MB default
// receive limit, which is what a detailed status on a large network looks like.
type bigStatusServer struct {
proto.UnimplementedDaemonServiceServer
payload string
}
func (s *bigStatusServer) Status(context.Context, *proto.StatusRequest) (*proto.StatusResponse, error) {
return &proto.StatusResponse{Status: s.payload}, nil
}
func startBigStatusServer(t *testing.T, payload string) string {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err, "listen on loopback")
srv := grpc.NewServer()
proto.RegisterDaemonServiceServer(srv, &bigStatusServer{payload: payload})
go func() {
_ = srv.Serve(listener)
}()
t.Cleanup(srv.Stop)
return "tcp://" + listener.Addr().String()
}
func TestDialTargetAcceptsAStatusOverTheGrpcDefault(t *testing.T) {
payload := strings.Repeat("x", 5*1024*1024)
addr := startBigStatusServer(t, payload)
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
target, opts := DialTarget(addr)
conn, err := grpc.NewClient(target, opts...)
require.NoError(t, err, "dial the daemon")
t.Cleanup(func() { _ = conn.Close() })
resp, err := proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{})
require.NoError(t, err, "a detailed status must not be rejected for its size")
assert.Len(t, resp.GetStatus(), len(payload), "the whole response must arrive")
}
// TestDialTargetRaisesTheDefaultLimit is the negative control: the same response
// over a connection carrying gRPC's own defaults is refused, which is the failure
// reported by `netbird status -d` on a large deployment.
func TestDialTargetRaisesTheDefaultLimit(t *testing.T) {
payload := strings.Repeat("x", 5*1024*1024)
addr := startBigStatusServer(t, payload)
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
conn, err := grpc.NewClient(
strings.TrimPrefix(addr, "tcp://"),
grpc.WithTransportCredentials(insecure.NewCredentials()),
)
require.NoError(t, err, "dial with the library defaults")
t.Cleanup(func() { _ = conn.Close() })
_, err = proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{})
require.Error(t, err, "the library default must reject this response")
assert.Equal(t, codes.ResourceExhausted, status.Code(err), "gRPC rejects an oversized message")
}
+4 -1
View File
@@ -36,7 +36,10 @@ const (
// 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())}
opts := []grpc.DialOption{
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithDefaultCallOptions(grpc.MaxCallRecvMsgSize(MaxRecvMsgSize())),
}
if name, ok := strings.CutPrefix(addr, pipeScheme); ok {
paths := PipePaths(name)
+53 -1
View File
@@ -379,9 +379,38 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
}
}
// bundleFilePattern names the bundle zips Generate creates in tempDir; the
// asterisk is filled in by os.CreateTemp.
const bundleFilePattern = "netbird.debug.*.zip"
const exportedBundlePrefix = "netbird.debug-file."
const exportedBundleMaxAge = 24 * time.Hour
// RemoveStaleBundles deletes bundle zips that an interrupted generation or
// upload left behind in dir. Only files older than maxAge go, so a bundle that
// another caller is still writing or uploading in the same directory survives.
// Exported bundles are kept for exportedBundleMaxAge instead.
func RemoveStaleBundles(dir string, maxAge time.Duration) {
removeStaleFiles(dir, bundleFilePattern, maxAge)
removeStaleFiles(dir, exportedBundlePrefix+"*.zip", exportedBundleMaxAge)
}
// ExportBundle renames a generated bundle out of the RemoveStaleBundles pattern
// and returns the new path. The caller owns the file from then on; an export
// abandoned for longer than exportedBundleMaxAge is removed by RemoveStaleBundles.
func ExportBundle(path string) (string, error) {
base := strings.TrimPrefix(filepath.Base(path), strings.SplitN(bundleFilePattern, "*", 2)[0])
exported := filepath.Join(filepath.Dir(path), exportedBundlePrefix+base)
if err := os.Rename(path, exported); err != nil {
return "", fmt.Errorf("export debug bundle: %w", err)
}
return exported, nil
}
// Generate creates a debug bundle and returns the location.
func (g *BundleGenerator) Generate() (resp string, err error) {
bundlePath, err := os.CreateTemp(g.tempDir, "netbird.debug.*.zip")
bundlePath, err := os.CreateTemp(g.tempDir, bundleFilePattern)
if err != nil {
return "", fmt.Errorf("create zip file: %w", err)
}
@@ -1725,3 +1754,26 @@ func anonymizeSlice(v []any, anonymizer *anonymize.Anonymizer) []any {
}
return v
}
func removeStaleFiles(dir, pattern string, maxAge time.Duration) {
matches, err := filepath.Glob(filepath.Join(dir, pattern))
if err != nil {
log.Debugf("glob stale debug bundles in %s: %v", dir, err)
return
}
cutoff := time.Now().Add(-maxAge)
for _, path := range matches {
info, err := os.Stat(path)
if err != nil || info.ModTime().After(cutoff) {
continue
}
if err := os.Remove(path); err != nil {
if !errors.Is(err, fs.ErrNotExist) {
log.Warnf("remove stale debug bundle %s: %v", path, err)
}
continue
}
log.Infof("removed stale debug bundle %s", path)
}
}
+50
View File
@@ -4,6 +4,7 @@ import (
"archive/zip"
"bytes"
"encoding/json"
"fmt"
"net"
"net/netip"
"net/url"
@@ -969,3 +970,52 @@ func renderAddConfigSpecific(g *BundleGenerator) string {
func newAnonymizerForTest() *anonymize.Anonymizer {
return anonymize.NewAnonymizer(anonymize.DefaultAddresses())
}
func TestRemoveStaleBundles(t *testing.T) {
dir := t.TempDir()
stale := filepath.Join(dir, "netbird.debug.111.zip")
fresh := filepath.Join(dir, "netbird.debug.222.zip")
other := filepath.Join(dir, "netbird.debug.333.txt")
owned := filepath.Join(dir, "netbird.debug.444.zip")
abandoned := filepath.Join(dir, "netbird.debug.555.zip")
for _, p := range []string{stale, fresh, other, owned, abandoned} {
require.NoError(t, os.WriteFile(p, []byte("x"), 0o600))
}
exported, err := ExportBundle(owned)
require.NoError(t, err)
exportedAbandoned, err := ExportBundle(abandoned)
require.NoError(t, err)
old := time.Now().Add(-2 * time.Hour)
for _, p := range []string{stale, other, exported} {
require.NoError(t, os.Chtimes(p, old, old))
}
ancient := time.Now().Add(-exportedBundleMaxAge - time.Hour)
require.NoError(t, os.Chtimes(exportedAbandoned, ancient, ancient))
RemoveStaleBundles(dir, time.Hour)
assert.NoFileExists(t, stale, "bundle older than maxAge should be removed")
assert.FileExists(t, fresh, "bundle younger than maxAge must survive, it may still be uploading")
assert.FileExists(t, other, "files outside the bundle pattern must not be touched")
assert.NoFileExists(t, owned)
assert.FileExists(t, exported, "exported bundle is caller-owned and must survive maxAge")
assert.NoFileExists(t, exportedAbandoned, "exported bundle older than exportedBundleMaxAge is abandoned")
}
func TestBundleIncludesNetworkMap(t *testing.T) {
for _, anonymize := range []bool{false, true} {
t.Run(fmt.Sprintf("anonymize=%t", anonymize), func(t *testing.T) {
g := NewBundleGenerator(GeneratorDependencies{
SyncResponse: &mgmProto.SyncResponse{NetworkMap: &mgmProto.NetworkMap{Serial: 1}},
}, BundleConfig{Anonymize: anonymize})
require.Contains(t, bundleEntries(t, g), "network_map.json")
})
}
}
func TestBundleOmitsNetworkMapWithoutSyncResponse(t *testing.T) {
g := NewBundleGenerator(GeneratorDependencies{}, BundleConfig{})
require.NotContains(t, bundleEntries(t, g), "network_map.json")
}
+84 -14
View File
@@ -124,19 +124,9 @@ func newHostManager(wgInterface WGIface) (*registryConfigurator, error) {
return nil, err
}
var useGPO bool
k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE)
if err != nil {
log.Debugf("failed to open GPO DNS policy root: %v", err)
} else {
closer(k)
useGPO = true
log.Infof("detected GPO DNS policy configuration, using policy store")
}
configurator := &registryConfigurator{
guid: guid,
gpo: useGPO,
gpo: useGPOPolicyStore(),
}
origNameservers, err := configurator.captureOriginalNameservers()
@@ -576,14 +566,22 @@ func (r *registryConfigurator) setInterfaceRegistryKeyStringValue(key, value str
return nil
}
// deleteInterfaceRegistryKeyProperty removes a value from the interface key.
// A value that is already gone, or an interface key that is, is not an error:
// the caller asked for the value not to be there, and a cleanup that runs twice
// has to reach its later steps on the second run as well.
func (r *registryConfigurator) deleteInterfaceRegistryKeyProperty(propertyKey string) error {
regKey, err := r.getInterfaceRegistryKey()
if err != nil {
switch {
case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND):
log.Debugf("interface key of %s does not exist, nothing to delete %s from", r.guid, propertyKey)
return nil
case err != nil:
return fmt.Errorf("get interface registry key: %w", err)
}
defer closer(regKey)
if err := regKey.DeleteValue(propertyKey); err != nil {
if err := regKey.DeleteValue(propertyKey); err != nil && !errors.Is(err, registry.ErrNotExist) {
return fmt.Errorf("delete registry key %s: %w", propertyKey, err)
}
return nil
@@ -612,7 +610,12 @@ func (r *registryConfigurator) restoreHostDNS() error {
go r.flushDNSCache()
return nil
// Last, and only on the way out, once no rule of ours is left: during a
// session the store is where the rules of this run live, and emptying it
// mid-session would have the next rule recreate it anyway. Propagated so a
// failure keeps the shutdown state for the next run to retry, rather than
// leaving the store to hold up every rule change from here on.
return removeEmptyGPOPolicyStore()
}
// removeDNSMatchPolicies deletes every NRPT rule this client may have created,
@@ -651,6 +654,73 @@ func (r *registryConfigurator) restoreUncleanShutdownDNS() error {
return r.restoreHostDNS()
}
// useGPOPolicyStore reports whether NRPT rules have to go into the group policy
// store, and clears an empty one out of the way first.
//
// The order is the point. A store left empty by an earlier run would otherwise
// decide this run too, sending its rules somewhere the resolver only reads when
// the policy engine next applies DNS client policy. Removing it before the
// choice is made leaves the local store authoritative for the whole session,
// including the first one after an upgrade.
func useGPOPolicyStore() bool {
if err := removeEmptyGPOPolicyStore(); err != nil {
// Nothing to retry against here: the worst case is the run going
// through the group policy store, which is where it would have gone
// before this check existed.
log.Warnf("%v", err)
}
k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE)
if err != nil {
log.Debugf("failed to open GPO DNS policy root: %v", err)
return false
}
closer(k)
log.Infof("detected GPO DNS policy configuration, using policy store")
return true
}
// removeEmptyGPOPolicyStore deletes the group policy DnsPolicyConfig key once
// nothing is left in it. The key survives the deletion of the last rule it
// held, and the client treats its presence as "group policy configures the
// NRPT", so an empty one left behind keeps every later run writing rules there.
// Rules in that store reach the resolver only when the policy engine next
// applies DNS client policy, and a rule this client writes belongs to no GPO,
// so nothing schedules that application: both adding and removing a rule are
// held up by a minute or more, and for a removal that is a catch-all rule
// resolving every name over an interface that no longer exists. With the store
// absent the local one is authoritative and a change applies at once.
//
// A store that still holds rules, values or subkeys of somebody else's is left
// alone.
func removeEmptyGPOPolicyStore() error {
k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE)
switch {
case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND):
return nil
case err != nil:
return fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err)
}
info, err := k.Stat()
closer(k)
if err != nil {
return fmt.Errorf("stat HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err)
}
if info.SubKeyCount != 0 || info.ValueCount != 0 {
return nil
}
if err := registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot); err != nil {
return fmt.Errorf("delete empty HKEY_LOCAL_MACHINE\\%s: %w", GPODNSPolicyConfigRoot, err)
}
log.Infof("removed the empty GPO DNS policy store, leaving the local one authoritative")
return nil
}
// listNRPTRuleKeys returns the names of our NRPT rule keys under a policy store
// root. An absent root holds nothing to clean up, which is the normal state of
// the GPO store on a machine without DNS Client policy.
+129
View File
@@ -8,6 +8,8 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/sys/windows/registry"
"github.com/netbirdio/netbird/client/internal/winregistry"
)
// TestNRPTEntriesCleanupOnConfigChange tests that old NRPT entries are properly cleaned up
@@ -405,3 +407,130 @@ func TestNRPTDomainBatching(t *testing.T) {
})
}
}
// TestRemoveEmptyGPOPolicyStore verifies that cleanup takes the GPO policy
// store itself with it once our rules are gone, since the store existing keeps
// the local one from being applied, and that a store with somebody else's rule
// in it is left alone.
func TestRemoveEmptyGPOPolicyStore(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
t.Cleanup(func() { cleanupRegistryKeys(t) })
cleanupRegistryKeys(t)
testIP := netip.MustParseAddr("100.64.0.1")
cfg := &registryConfigurator{gpo: true}
// a store holding a rule of ours is kept, because the rule is still applied
require.NoError(t, cfg.addDNSMatchPolicy([]string{".example.com"}, testIP))
exists, err := registryKeyExists(gpoDnsPolicyConfigMatchPath + "-0")
require.NoError(t, err)
require.True(t, exists, "Should write the rule to the GPO policy store")
require.NoError(t, removeEmptyGPOPolicyStore())
exists, err = registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.True(t, exists, "Should keep a policy store that still holds a rule")
// once the rules are gone the store goes with them
require.NoError(t, cfg.removeDNSMatchPolicies())
require.NoError(t, removeEmptyGPOPolicyStore())
exists, err = registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.False(t, exists, "Should remove the GPO policy store once it is empty")
// A store is not ours to remove while somebody else has a rule in it. The
// rule is written volatile like our own: the rules above created the parent
// chain volatile, and Windows refuses a stable subkey under a volatile
// parent.
foreignRule := GPODNSPolicyConfigRoot + `\{2A3B4C5D-6E7F-4041-8283-84858687888A}`
foreignKey, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, foreignRule, registry.SET_VALUE)
require.NoError(t, err, "Should create a foreign GPO rule")
foreignKey.Close()
t.Cleanup(func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignRule)
_ = registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot)
})
require.NoError(t, cfg.removeDNSMatchPolicies())
require.NoError(t, removeEmptyGPOPolicyStore())
exists, err = registryKeyExists(foreignRule)
require.NoError(t, err)
assert.True(t, exists, "Should not remove a foreign rule")
exists, err = registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.True(t, exists, "Should keep a policy store that still holds a foreign rule")
}
// TestDeleteInterfaceRegistryKeyPropertyTwice verifies that removing a value
// that is already gone, or one on an interface key that is, reports success.
// Teardown runs again after a failed cleanup, and the steps that follow this
// one have to be reached on that second run.
func TestDeleteInterfaceRegistryKeyPropertyTwice(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
testGUID := "{12345678-1234-1234-1234-123456789ABC}"
interfacePath := InterfaceConfigPath + `\` + testGUID
testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE)
require.NoError(t, err, "Should create test interface registry key")
testKey.Close()
t.Cleanup(func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath)
})
cfg := &registryConfigurator{guid: testGUID}
require.NoError(t, cfg.setInterfaceRegistryKeyStringValue(interfaceConfigSearchListKey, "example.com"))
require.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey))
assert.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey),
"Should report success for a value that is already gone")
// and with the interface key itself gone, as it is once the adapter is
require.NoError(t, registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath))
assert.NoError(t, cfg.deleteInterfaceRegistryKeyProperty(interfaceConfigSearchListKey),
"Should report success when the interface key does not exist")
}
// TestUseGPOPolicyStoreClearsEmptyStore verifies that the store is cleared
// before it is consulted, so an empty one left by an earlier run does not send
// this run's rules to the group policy store. A store somebody else has a rule
// in still decides where the rules go.
func TestUseGPOPolicyStoreClearsEmptyStore(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
t.Cleanup(func() { cleanupRegistryKeys(t) })
cleanupRegistryKeys(t)
// the leftover an earlier run used to keep, which the client read as
// "group policy configures the NRPT" for every run after it
emptyStore, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.SET_VALUE)
require.NoError(t, err, "Should create the GPO policy store")
emptyStore.Close()
assert.False(t, useGPOPolicyStore(), "An empty store should not decide where the rules go")
exists, err := registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.False(t, exists, "Should clear the empty store before consulting it")
foreignRule := GPODNSPolicyConfigRoot + `\{2A3B4C5D-6E7F-4041-8283-84858687888A}`
foreignKey, _, err := winregistry.CreateVolatileKey(registry.LOCAL_MACHINE, foreignRule, registry.SET_VALUE)
require.NoError(t, err, "Should create a foreign GPO rule")
foreignKey.Close()
t.Cleanup(func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignRule)
_ = registry.DeleteKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot)
})
assert.True(t, useGPOPolicyStore(), "A store holding a rule should decide where the rules go")
exists, err = registryKeyExists(GPODNSPolicyConfigRoot)
require.NoError(t, err)
assert.True(t, exists, "Should keep a store that holds a rule")
}
+7 -10
View File
@@ -9,9 +9,9 @@ import (
"os"
"testing"
"go.uber.org/mock/gomock"
"github.com/miekg/dns"
"github.com/stretchr/testify/assert"
"go.uber.org/mock/gomock"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/netbirdio/netbird/client/iface"
@@ -24,6 +24,10 @@ import (
nbdns "github.com/netbirdio/netbird/dns"
)
// testIFaceBlackList mirrors the overlay prefixes profilemanager.DefaultInterfaceBlacklist
// carries. Declared here rather than imported because profilemanager imports this package.
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
func TestUpdateDNSServer(t *testing.T) {
nameServers := []nbdns.NameServer{
@@ -243,10 +247,7 @@ func TestUpdateDNSServer(t *testing.T) {
for n, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
privKey, _ := wgtypes.GenerateKey()
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
opts := iface.WGIFaceOpts{
IFaceName: fmt.Sprintf("utun230%d", n),
@@ -348,11 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) {
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"})
if err != nil {
t.Errorf("create stdnet: %v", err)
return
}
newNet := stdnet.NewNet(context.Background(), []string{"utun2301"})
privKey, _ := wgtypes.GeneratePrivateKey()
opts := iface.WGIFaceOpts{
+1 -5
View File
@@ -394,11 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) {
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"})
if err != nil {
t.Fatalf("create stdnet: %v", err)
return nil, err
}
newNet := stdnet.NewNet(context.Background(), []string{"utun2301"})
privKey, _ := wgtypes.GeneratePrivateKey()
-148
View File
@@ -1,148 +0,0 @@
// Code generated by bpf2go; DO NOT EDIT.
//go:build mips || mips64 || ppc64 || s390x
package ebpf
import (
"bytes"
_ "embed"
"fmt"
"io"
"github.com/cilium/ebpf"
)
// loadBpf returns the embedded CollectionSpec for bpf.
func loadBpf() (*ebpf.CollectionSpec, error) {
reader := bytes.NewReader(_BpfBytes)
spec, err := ebpf.LoadCollectionSpecFromReader(reader)
if err != nil {
return nil, fmt.Errorf("can't load bpf: %w", err)
}
return spec, err
}
// loadBpfObjects loads bpf and converts it into a struct.
//
// The following types are suitable as obj argument:
//
// *bpfObjects
// *bpfPrograms
// *bpfMaps
//
// See ebpf.CollectionSpec.LoadAndAssign documentation for details.
func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error {
spec, err := loadBpf()
if err != nil {
return err
}
return spec.LoadAndAssign(obj, opts)
}
// bpfSpecs contains maps and programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfSpecs struct {
bpfProgramSpecs
bpfMapSpecs
bpfVariableSpecs
}
// bpfProgramSpecs contains programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfProgramSpecs struct {
NbXdpProg *ebpf.ProgramSpec `ebpf:"nb_xdp_prog"`
}
// bpfMapSpecs contains maps before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfMapSpecs struct {
NbFeatures *ebpf.MapSpec `ebpf:"nb_features"`
NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"`
}
// bpfVariableSpecs contains global variables before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfVariableSpecs struct {
FlagFeatureWgProxy *ebpf.VariableSpec `ebpf:"flag_feature_wg_proxy"`
MapKeyFeatures *ebpf.VariableSpec `ebpf:"map_key_features"`
MapKeyProxyPort *ebpf.VariableSpec `ebpf:"map_key_proxy_port"`
MapKeyWgPort *ebpf.VariableSpec `ebpf:"map_key_wg_port"`
ProxyPort *ebpf.VariableSpec `ebpf:"proxy_port"`
WgPort *ebpf.VariableSpec `ebpf:"wg_port"`
}
// bpfObjects contains all objects after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfObjects struct {
bpfPrograms
bpfMaps
bpfVariables
}
func (o *bpfObjects) Close() error {
return _BpfClose(
&o.bpfPrograms,
&o.bpfMaps,
)
}
// bpfMaps contains all maps after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfMaps struct {
NbFeatures *ebpf.Map `ebpf:"nb_features"`
NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"`
}
func (m *bpfMaps) Close() error {
return _BpfClose(
m.NbFeatures,
m.NbWgProxySettingsMap,
)
}
// bpfVariables contains all global variables after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfVariables struct {
FlagFeatureWgProxy *ebpf.Variable `ebpf:"flag_feature_wg_proxy"`
MapKeyFeatures *ebpf.Variable `ebpf:"map_key_features"`
MapKeyProxyPort *ebpf.Variable `ebpf:"map_key_proxy_port"`
MapKeyWgPort *ebpf.Variable `ebpf:"map_key_wg_port"`
ProxyPort *ebpf.Variable `ebpf:"proxy_port"`
WgPort *ebpf.Variable `ebpf:"wg_port"`
}
// bpfPrograms contains all programs after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfPrograms struct {
NbXdpProg *ebpf.Program `ebpf:"nb_xdp_prog"`
}
func (p *bpfPrograms) Close() error {
return _BpfClose(
p.NbXdpProg,
)
}
func _BpfClose(closers ...io.Closer) error {
for _, closer := range closers {
if err := closer.Close(); err != nil {
return err
}
}
return nil
}
// Do not access this directly.
//
//go:embed bpf_bpfeb.o
var _BpfBytes []byte
Binary file not shown.
-148
View File
@@ -1,148 +0,0 @@
// Code generated by bpf2go; DO NOT EDIT.
//go:build 386 || amd64 || arm || arm64 || loong64 || mips64le || mipsle || ppc64le || riscv64 || wasm
package ebpf
import (
"bytes"
_ "embed"
"fmt"
"io"
"github.com/cilium/ebpf"
)
// loadBpf returns the embedded CollectionSpec for bpf.
func loadBpf() (*ebpf.CollectionSpec, error) {
reader := bytes.NewReader(_BpfBytes)
spec, err := ebpf.LoadCollectionSpecFromReader(reader)
if err != nil {
return nil, fmt.Errorf("can't load bpf: %w", err)
}
return spec, err
}
// loadBpfObjects loads bpf and converts it into a struct.
//
// The following types are suitable as obj argument:
//
// *bpfObjects
// *bpfPrograms
// *bpfMaps
//
// See ebpf.CollectionSpec.LoadAndAssign documentation for details.
func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error {
spec, err := loadBpf()
if err != nil {
return err
}
return spec.LoadAndAssign(obj, opts)
}
// bpfSpecs contains maps and programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfSpecs struct {
bpfProgramSpecs
bpfMapSpecs
bpfVariableSpecs
}
// bpfProgramSpecs contains programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfProgramSpecs struct {
NbXdpProg *ebpf.ProgramSpec `ebpf:"nb_xdp_prog"`
}
// bpfMapSpecs contains maps before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfMapSpecs struct {
NbFeatures *ebpf.MapSpec `ebpf:"nb_features"`
NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"`
}
// bpfVariableSpecs contains global variables before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type bpfVariableSpecs struct {
FlagFeatureWgProxy *ebpf.VariableSpec `ebpf:"flag_feature_wg_proxy"`
MapKeyFeatures *ebpf.VariableSpec `ebpf:"map_key_features"`
MapKeyProxyPort *ebpf.VariableSpec `ebpf:"map_key_proxy_port"`
MapKeyWgPort *ebpf.VariableSpec `ebpf:"map_key_wg_port"`
ProxyPort *ebpf.VariableSpec `ebpf:"proxy_port"`
WgPort *ebpf.VariableSpec `ebpf:"wg_port"`
}
// bpfObjects contains all objects after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfObjects struct {
bpfPrograms
bpfMaps
bpfVariables
}
func (o *bpfObjects) Close() error {
return _BpfClose(
&o.bpfPrograms,
&o.bpfMaps,
)
}
// bpfMaps contains all maps after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfMaps struct {
NbFeatures *ebpf.Map `ebpf:"nb_features"`
NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"`
}
func (m *bpfMaps) Close() error {
return _BpfClose(
m.NbFeatures,
m.NbWgProxySettingsMap,
)
}
// bpfVariables contains all global variables after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfVariables struct {
FlagFeatureWgProxy *ebpf.Variable `ebpf:"flag_feature_wg_proxy"`
MapKeyFeatures *ebpf.Variable `ebpf:"map_key_features"`
MapKeyProxyPort *ebpf.Variable `ebpf:"map_key_proxy_port"`
MapKeyWgPort *ebpf.Variable `ebpf:"map_key_wg_port"`
ProxyPort *ebpf.Variable `ebpf:"proxy_port"`
WgPort *ebpf.Variable `ebpf:"wg_port"`
}
// bpfPrograms contains all programs after they have been loaded into the kernel.
//
// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign.
type bpfPrograms struct {
NbXdpProg *ebpf.Program `ebpf:"nb_xdp_prog"`
}
func (p *bpfPrograms) Close() error {
return _BpfClose(
p.NbXdpProg,
)
}
func _BpfClose(closers ...io.Closer) error {
for _, closer := range closers {
if err := closer.Close(); err != nil {
return err
}
}
return nil
}
// Do not access this directly.
//
//go:embed bpf_bpfel.o
var _BpfBytes []byte
Binary file not shown.
-115
View File
@@ -1,115 +0,0 @@
package ebpf
import (
_ "embed"
"net"
"sync"
"github.com/cilium/ebpf/link"
"github.com/cilium/ebpf/rlimit"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/ebpf/manager"
)
const (
mapKeyFeatures uint32 = 0
featureFlagWGProxy = 0b00000001
)
var (
singleton manager.Manager
singletonLock = &sync.Mutex{}
)
// required packages libbpf-dev, libc6-dev-i386-amd64-cross
// GeneralManager is used to load multiple eBPF programs with a custom check (if then) done in prog.c
// The manager simply adds a feature (byte) of each program to a map that is shared between the userspace and kernel.
// When packet arrives, the C code checks for each feature (if it is set) and executes each enabled program (e.g., wg_proxy.c).
//
//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc clang-14 bpf src/prog.c -- -I /usr/x86_64-linux-gnu/include -include src/bpf_map_def.h
type GeneralManager struct {
lock sync.Mutex
link link.Link
featureFlags uint16
bpfObjs bpfObjects
}
// GetEbpfManagerInstance return a static eBpf Manager instance
func GetEbpfManagerInstance() manager.Manager {
singletonLock.Lock()
defer singletonLock.Unlock()
if singleton != nil {
return singleton
}
singleton = &GeneralManager{}
return singleton
}
func (tf *GeneralManager) setFeatureFlag(feature uint16) {
tf.featureFlags |= feature
}
func (tf *GeneralManager) loadXdp() error {
if tf.link != nil {
return nil
}
// it required for Docker
err := rlimit.RemoveMemlock()
if err != nil {
return err
}
iFace, err := net.InterfaceByName("lo")
if err != nil {
return err
}
// load pre-compiled programs into the kernel.
err = loadBpfObjects(&tf.bpfObjs, nil)
if err != nil {
return err
}
tf.link, err = link.AttachXDP(link.XDPOptions{
Program: tf.bpfObjs.NbXdpProg,
Interface: iFace.Index,
})
if err != nil {
_ = tf.bpfObjs.Close()
tf.link = nil
return err
}
return nil
}
func (tf *GeneralManager) unsetFeatureFlag(feature uint16) error {
tf.lock.Lock()
defer tf.lock.Unlock()
tf.featureFlags &^= feature
if tf.link == nil {
return nil
}
if tf.featureFlags == 0 {
return tf.close()
}
return tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags)
}
func (tf *GeneralManager) close() error {
log.Debugf("detach ebpf program ")
err := tf.bpfObjs.Close()
if err != nil {
log.Warnf("failed to close eBpf objects: %s", err)
}
err = tf.link.Close()
tf.link = nil
return err
}
@@ -1,31 +0,0 @@
package ebpf
import (
"testing"
)
func TestManager_setFeatureFlag(t *testing.T) {
mgr := GeneralManager{}
mgr.setFeatureFlag(featureFlagWGProxy)
if mgr.featureFlags != featureFlagWGProxy {
t.Errorf("invalid feature state")
}
mgr.setFeatureFlag(featureFlagWGProxy)
if mgr.featureFlags != featureFlagWGProxy {
t.Errorf("setting a flag twice must be idempotent, got: %d", mgr.featureFlags)
}
}
func TestManager_unsetFeatureFlag(t *testing.T) {
mgr := GeneralManager{}
mgr.setFeatureFlag(featureFlagWGProxy)
err := mgr.unsetFeatureFlag(featureFlagWGProxy)
if err != nil {
t.Errorf("unexpected error: %s", err)
}
if mgr.featureFlags != 0 {
t.Errorf("invalid feature state, expected: %d, got: %d", 0, mgr.featureFlags)
}
}
@@ -1,16 +0,0 @@
// libbpf 1.0 removed struct bpf_map_def, but the programs here keep the legacy
// map definitions: they load on kernels built without BTF, which BTF-style
// (SEC(".maps")) definitions do not. Define the struct ourselves so the
// programs compile against current libbpf headers.
#ifndef NB_BPF_MAP_DEF_H
#define NB_BPF_MAP_DEF_H
struct bpf_map_def {
unsigned int type;
unsigned int key_size;
unsigned int value_size;
unsigned int max_entries;
unsigned int map_flags;
};
#endif
-54
View File
@@ -1,54 +0,0 @@
#include <stdbool.h>
#include <linux/if_ether.h> // ETH_P_IP
#include <linux/udp.h>
#include <linux/ip.h>
#include <netinet/in.h>
#include <linux/bpf.h>
#include <bpf/bpf_helpers.h>
#include "wg_proxy.c"
const __u16 flag_feature_wg_proxy = 0b01;
const __u32 map_key_features = 0;
struct bpf_map_def SEC("maps") nb_features = {
.type = BPF_MAP_TYPE_ARRAY,
.key_size = sizeof(__u32),
.value_size = sizeof(__u16),
.max_entries = 10,
};
SEC("xdp")
int nb_xdp_prog(struct xdp_md *ctx) {
__u16 *features;
features = bpf_map_lookup_elem(&nb_features, &map_key_features);
if (!features) {
return XDP_PASS;
}
void *data = (void *)(long)ctx->data;
void *data_end = (void *)(long)ctx->data_end;
struct ethhdr *eth = data;
struct iphdr *ip = (data + sizeof(struct ethhdr));
struct udphdr *udp = (data + sizeof(struct ethhdr) + sizeof(struct iphdr));
// return early if not enough data
if (data + sizeof(struct ethhdr) + sizeof(struct iphdr) + sizeof(struct udphdr) > data_end){
return XDP_PASS;
}
// skip non IPv4 packages
if (eth->h_proto != htons(ETH_P_IP)) {
return XDP_PASS;
}
// skip non UPD packages
if (ip->protocol != IPPROTO_UDP) {
return XDP_PASS;
}
if (*features & flag_feature_wg_proxy) {
xdp_wg_proxy(ip, udp);
}
return XDP_PASS;
}
char _license[] SEC("license") = "GPL";
-27
View File
@@ -1,27 +0,0 @@
# XDP programs
`prog.c` is attached to the `lo` device and dispatches to the features enabled in the
`nb_features` map. The only feature is the WireGuard proxy (`wg_proxy.c`): it rewrites
loopback UDP sent from the WireGuard listen port so it reaches the userspace relay proxy
port instead, and swaps the peer endpoint port into the source so the proxy can tell
peers apart.
Maps use the legacy `struct bpf_map_def` form, defined in `bpf_map_def.h` because libbpf
1.0 removed it. They load on kernels built without BTF, which BTF-style (`SEC(".maps")`)
definitions do not.
Regenerate the objects with `go generate ./client/internal/ebpf/ebpf/`; it needs
`clang-14`. Loading a regenerated object needs root, attaching it needs `bpf_link`
(kernel >= 5.7), and only one XDP program can own `lo` at a time.
# Debug
The CONFIG_BPF_EVENTS kernel module is required for bpf_printk.
Apply this code to use bpf_printk
```
#define bpf_printk(fmt, ...) \
({ \
char ____fmt[] = fmt; \
bpf_trace_printk(____fmt, sizeof(____fmt), ##__VA_ARGS__); \
})
```
-60
View File
@@ -1,60 +0,0 @@
const __u32 map_key_proxy_port = 0;
const __u32 map_key_wg_port = 1;
struct bpf_map_def SEC("maps") nb_wg_proxy_settings_map = {
.type = BPF_MAP_TYPE_ARRAY,
.key_size = sizeof(__u32),
.value_size = sizeof(__u16),
.max_entries = 10,
};
__u16 proxy_port = 0;
__u16 wg_port = 0;
bool read_port_settings() {
__u16 *value;
value = bpf_map_lookup_elem(&nb_wg_proxy_settings_map, &map_key_proxy_port);
if (!value) {
return false;
}
proxy_port = *value;
value = bpf_map_lookup_elem(&nb_wg_proxy_settings_map, &map_key_wg_port);
if (!value) {
return false;
}
wg_port = htons(*value);
return true;
}
int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) {
if (proxy_port == 0 || wg_port == 0) {
if (!read_port_settings()){
return XDP_PASS;
}
// bpf_printk("proxy port: %d, wg port: %d", proxy_port, wg_port);
}
// 2130706433 = 127.0.0.1
if (ip->daddr != htonl(2130706433)) {
return XDP_PASS;
}
if (udp->source != wg_port){
return XDP_PASS;
}
__be16 new_src_port = udp->dest;
__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;
}
@@ -1,41 +0,0 @@
package ebpf
import log "github.com/sirupsen/logrus"
const (
mapKeyProxyPort uint32 = 0
mapKeyWgPort uint32 = 1
)
func (tf *GeneralManager) LoadWgProxy(proxyPort, wgPort int) error {
log.Debugf("load ebpf WG proxy")
tf.lock.Lock()
defer tf.lock.Unlock()
err := tf.loadXdp()
if err != nil {
return err
}
err = tf.bpfObjs.NbWgProxySettingsMap.Put(mapKeyProxyPort, uint16(proxyPort))
if err != nil {
return err
}
err = tf.bpfObjs.NbWgProxySettingsMap.Put(mapKeyWgPort, uint16(wgPort))
if err != nil {
return err
}
tf.setFeatureFlag(featureFlagWGProxy)
err = tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags)
if err != nil {
return err
}
return nil
}
func (tf *GeneralManager) FreeWGProxy() error {
log.Debugf("free ebpf WG proxy")
return tf.unsetFeatureFlag(featureFlagWGProxy)
}
@@ -1,15 +0,0 @@
//go:build !android
package ebpf
import (
"github.com/netbirdio/netbird/client/internal/ebpf/ebpf"
"github.com/netbirdio/netbird/client/internal/ebpf/manager"
)
// GetEbpfManagerInstance is a wrapper function. This encapsulation is required because if the code import the internal
// ebpf package the Go compiler will include the object files. But it is not supported on Android. It can cause instant
// panic on older Android version.
func GetEbpfManagerInstance() manager.Manager {
return ebpf.GetEbpfManagerInstance()
}
@@ -1,10 +0,0 @@
//go:build !linux || android
package ebpf
import "github.com/netbirdio/netbird/client/internal/ebpf/manager"
// GetEbpfManagerInstance return error because ebpf is not supported on all os
func GetEbpfManagerInstance() manager.Manager {
panic("unsupported os")
}
-7
View File
@@ -1,7 +0,0 @@
package manager
// Manager is used to load multiple eBPF programs. E.g., the WireGuard proxy
type Manager interface {
LoadWgProxy(proxyPort, wgPort int) error
FreeWGProxy() error
}
+11
View File
@@ -6,6 +6,17 @@ import (
"path/filepath"
)
// CheckOnlyOwnerWritable reports an error unless path, and every directory
// leading to it, is owned by an account that can already act with the privileges
// the caller holds, and is writable by nobody else.
//
// Exported for callers outside elevation that read a file while privileged and
// then act on what it says: the same question this package asks of an
// executable, asked of a configuration file.
func CheckOnlyOwnerWritable(path string) error {
return checkOnlyOwnerWritable(path)
}
// trustedSelf returns the path of this executable, provided it is one we are
// willing to have run as root.
//
+6 -26
View File
@@ -664,10 +664,6 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
}
e.wgDevice.Store(e.wgInterface.GetWGDevice())
// Set up notrack rules immediately after proxy is listening to prevent
// conntrack entries from being created before the rules are in place
e.setupWGProxyNoTrack()
// Start after interface is up since port may have been resolved from 0 or changed if occupied
e.shutdownWg.Add(1)
go func() {
@@ -805,23 +801,6 @@ func (e *Engine) initFirewall() error {
return nil
}
// setupWGProxyNoTrack configures connection tracking exclusion for WireGuard proxy traffic.
// This prevents conntrack/MASQUERADE from affecting loopback traffic between WireGuard and the eBPF proxy.
func (e *Engine) setupWGProxyNoTrack() {
if e.firewall == nil {
return
}
proxyPort := e.wgInterface.GetProxyPort()
if proxyPort == 0 {
return
}
if err := e.firewall.SetupEBPFProxyNoTrack(proxyPort, uint16(e.config.WgPort)); err != nil {
log.Warnf("failed to setup ebpf proxy notrack: %v", err)
}
}
func (e *Engine) blockLanAccess() {
if e.config.BlockInbound {
// no need to set up extra deny rules if inbound is already blocked in general
@@ -1064,7 +1043,11 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
// 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)
// With the firewall disabled there is no ACL manager to program, so
// RoutesFirewallRules would be built and then dropped. On a peer that
// routes many network resources that is the single most expensive
// step of the sync.
result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName, e.config.DisableFirewall)
if err != nil {
return fmt.Errorf("decode network map envelope: %w", err)
}
@@ -2208,10 +2191,7 @@ func (e *Engine) close() {
}
func (e *Engine) newWgIface() (*iface.WGIface, error) {
transportNet, err := e.newStdNet()
if err != nil {
log.Errorf("failed to create pion's stdnet: %s", err)
}
transportNet := e.newStdNet()
opts := iface.WGIFaceOpts{
IFaceName: e.config.WgIfaceName,
+15 -11
View File
@@ -12,12 +12,12 @@ import (
"testing"
"time"
"go.uber.org/mock/gomock"
"github.com/google/uuid"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel"
"go.uber.org/mock/gomock"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"google.golang.org/grpc"
"google.golang.org/grpc/keepalive"
@@ -27,6 +27,7 @@ import (
"github.com/netbirdio/netbird/client/iface/wgaddr"
"github.com/netbirdio/netbird/client/internal/dns"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
nbssh "github.com/netbirdio/netbird/client/ssh"
"github.com/netbirdio/netbird/client/system"
nbdns "github.com/netbirdio/netbird/dns"
@@ -81,6 +82,7 @@ func TestEngine_SSH(t *testing.T) {
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
WgPrivateKey: key,
WgPort: 33100,
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
ServerSSHAllowed: true,
MTU: iface.DefaultMTU,
SSHKey: sshKey,
@@ -204,11 +206,12 @@ func TestEngine_Sync(t *testing.T) {
}
relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU)
engine := NewEngine(ctx, cancel, &EngineConfig{
WgIfaceName: "utun103",
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
WgPrivateKey: key,
WgPort: 33100,
MTU: iface.DefaultMTU,
WgIfaceName: "utun103",
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
WgPrivateKey: key,
WgPort: 33100,
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
MTU: iface.DefaultMTU,
}, EngineServices{
SignalClient: &signal.MockClient{},
MgmClient: &mgmt.MockClient{SyncFunc: syncFunc},
@@ -412,11 +415,12 @@ func createEngine(ctx context.Context, cancel context.CancelFunc, setupKey strin
wgPort := 33100 + i
conf := &EngineConfig{
WgIfaceName: ifaceName,
WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address),
WgPrivateKey: key,
WgPort: wgPort,
MTU: iface.DefaultMTU,
WgIfaceName: ifaceName,
WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address),
WgPrivateKey: key,
WgPort: wgPort,
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
MTU: iface.DefaultMTU,
}
relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU)
+7 -5
View File
@@ -1,4 +1,4 @@
//go:build !js
//go:build !js && !android
package internal
@@ -7,10 +7,12 @@ import (
"github.com/netbirdio/netbird/client/internal/peer"
)
// newSessionWatcher returns the real SSO session expiry watcher for every
// non-wasm build. The js/wasm build gets a no-op stub from
// engine_sessionwatch_js.go so the sessionwatch package (and its timer
// machinery) never links into the wasm binary.
// newSessionWatcher returns the real SSO session expiry watcher. The js/wasm
// build gets a no-op stub from engine_sessionwatch_js.go so the sessionwatch
// package (and its timer machinery) never links into the wasm binary; the
// android build gets a deadline-only watcher from
// engine_sessionwatch_android.go because the app schedules the warnings
// itself.
func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
return sessionwatch.New(recorder)
}
@@ -0,0 +1,12 @@
//go:build android
package internal
import (
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
"github.com/netbirdio/netbird/client/internal/peer"
)
func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
return sessionwatch.NewDeadlineOnly(recorder)
}
+1 -1
View File
@@ -6,6 +6,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet"
)
func (e *Engine) newStdNet() (*stdnet.Net, error) {
func (e *Engine) newStdNet() *stdnet.Net {
return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList)
}
+1 -1
View File
@@ -2,6 +2,6 @@ package internal
import "github.com/netbirdio/netbird/client/internal/stdnet"
func (e *Engine) newStdNet() (*stdnet.Net, error) {
func (e *Engine) newStdNet() *stdnet.Net {
return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList)
}
+20 -16
View File
@@ -65,7 +65,6 @@ type MockWGIface struct {
GetStatsFunc func() (map[string]configurer.WGStats, error)
GetInterfaceGUIDStringFunc func() (string, error)
GetProxyFunc func() wgproxy.Proxy
GetProxyPortFunc func() uint16
GetNetFunc func() *netstack.Net
LastActivitiesFunc func() map[string]monotime.Time
}
@@ -162,13 +161,6 @@ func (m *MockWGIface) GetProxy() wgproxy.Proxy {
return m.GetProxyFunc()
}
func (m *MockWGIface) GetProxyPort() uint16 {
if m.GetProxyPortFunc != nil {
return m.GetProxyPortFunc()
}
return 0
}
func (m *MockWGIface) GetNet() *netstack.Net {
return m.GetNetFunc()
}
@@ -696,10 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) {
StatusRecorder: peer.NewRecorder("https://mgm"),
}, MobileDependency{})
engine.ctx = ctx
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
opts := iface.WGIFaceOpts{
IFaceName: wgIfaceName,
@@ -904,10 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) {
}, MobileDependency{})
engine.ctx = ctx
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
opts := iface.WGIFaceOpts{
IFaceName: wgIfaceName,
Address: wgaddr.MustParseWGAddress(wgAddr),
@@ -1533,3 +1519,21 @@ func TestOverlayAddrsFromAllowedIPs(t *testing.T) {
})
}
}
func TestEngine_SyncResponsePersistence(t *testing.T) {
e := &Engine{}
_, err := e.GetLatestSyncResponse()
require.Error(t, err, "persistence is disabled by default")
e.SetSyncResponsePersistence(true)
e.persistSyncResponse(&mgmtProto.SyncResponse{NetworkMap: &mgmtProto.NetworkMap{Serial: 7}})
got, err := e.GetLatestSyncResponse()
require.NoError(t, err)
assert.Equal(t, uint64(7), got.GetNetworkMap().GetSerial())
e.SetSyncResponsePersistence(false)
_, err = e.GetLatestSyncResponse()
require.Error(t, err)
}
-1
View File
@@ -28,7 +28,6 @@ type wgIfaceBase interface {
Up() (*udpmux.UniversalUDPMuxDefault, error)
UpdateAddr(newAddr wgaddr.Address) error
GetProxy() wgproxy.Proxy
GetProxyPort() uint16
UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error
RemoveEndpointAddress(key string) error
RemovePeer(peerKey string) error
+1 -1
View File
@@ -38,7 +38,7 @@ func asDaemon(t *testing.T, id Identity) {
prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate
t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate })
selfIdentity, selfKnown = id, true
selfMayDelegate = !id.IsPrivileged()
selfMayDelegate = mayDelegate(id)
}
func TestCallerIdentity_DirectConnections(t *testing.T) {
+6 -6
View File
@@ -18,7 +18,8 @@ import (
"google.golang.org/grpc/peer"
)
// Well-known Windows SIDs that identify a fully privileged principal.
// Well-known Windows SIDs. Only LocalSystem and BUILTIN\Administrators identify a
// privileged principal; the service accounts are shared by unrelated services.
const (
sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM
sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE
@@ -67,9 +68,9 @@ func (i Identity) IsWindows() bool {
// 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
// account names or group RIDs: an elevated token, the LocalSystem SID, or a
// token with BUILTIN\Administrators enabled. LocalService and NetworkService
// are not privileged by SID. 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
@@ -83,8 +84,7 @@ func (i Identity) IsPrivileged() bool {
return true
}
switch i.SID {
case sidLocalSystem, sidLocalService, sidNetworkService:
if i.SID == sidLocalSystem {
return true
}
@@ -64,3 +64,58 @@ func TestIdentitySameUser(t *testing.T) {
})
}
}
func TestIdentityIsPrivileged(t *testing.T) {
tests := []struct {
name string
id Identity
want bool
}{
{
name: "Root",
id: Identity{UID: 0, GID: 0},
want: true,
},
{
name: "Non-root",
id: Identity{UID: 1000, GID: 1000},
want: false,
},
{
name: "Local system windows",
id: Identity{SID: sidLocalSystem},
want: true,
},
{
name: "Windows elevated",
id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001", Elevated: true},
want: true,
},
{
name: "Admin group windows",
id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001", Groups: []string{sidAdministrators}},
want: true,
},
{
name: "Regular user windows",
id: Identity{SID: "S-1-5-21-1927267129-3959769253-3036563910-1001"},
want: false,
},
{
name: "Network service windows",
id: Identity{SID: sidNetworkService},
want: false,
},
{
name: "Local service windows",
id: Identity{SID: sidLocalService},
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assert.Equal(t, tt.want, tt.id.IsPrivileged())
})
}
}
+9 -1
View File
@@ -45,7 +45,15 @@ func init() {
// 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()
selfMayDelegate = mayDelegate(id)
}
// mayDelegate reports whether a daemon running as id may extend its authority to
// callers sharing its identity. The shared service accounts are excluded: their
// SID is held by unrelated services, so matching on it would grant them the
// daemon's authority.
func mayDelegate(id Identity) bool {
return !id.IsPrivileged() && id.SID != sidLocalService && id.SID != sidNetworkService
}
// IsDaemonSelf reports whether an identity is this very process. The JSON gateway
+26 -1
View File
@@ -98,7 +98,7 @@ func TestIsPrivilegedCaller_SelfRule(t *testing.T) {
t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate })
selfIdentity, selfKnown = tt.self, tt.selfKnown
selfMayDelegate = tt.selfKnown && !tt.self.IsPrivileged()
selfMayDelegate = tt.selfKnown && mayDelegate(tt.self)
if got := IsPrivilegedCaller(tt.caller); got != tt.want {
t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t",
@@ -132,3 +132,28 @@ func TestIsPrivilegedCaller_ThisProcess(t *testing.T) {
t.Errorf("an unrelated identity %v was treated as privileged", other)
}
}
// The shared service accounts are held by unrelated services, so a daemon running
// as one of them must not extend its authority to every process with that SID.
func TestMayDelegate(t *testing.T) {
tests := []struct {
name string
self Identity
want bool
}{
{name: "unprivileged unix user", self: Identity{UID: 1000}, want: true},
{name: "root", self: Identity{UID: 0}, want: false},
{name: "unprivileged windows user", self: Identity{SID: "S-1-5-21-1-2-3-1001"}, want: true},
{name: "elevated windows user", self: Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}, want: false},
{name: "local system", self: Identity{SID: sidLocalSystem}, want: false},
{name: "local service", self: Identity{SID: sidLocalService}, want: false},
{name: "network service", self: Identity{SID: sidNetworkService}, want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := mayDelegate(tt.self); got != tt.want {
t.Errorf("mayDelegate(%+v) = %v, want %v", tt.self, got, tt.want)
}
})
}
}
+41 -21
View File
@@ -135,9 +135,10 @@ type Conn struct {
// used to store the remote Rosenpass key for Relayed connection in case of connection update from ice
rosenpassRemoteKey []byte
wgProxyICE wgproxy.Proxy
wgProxyRelay wgproxy.Proxy
handshaker *Handshaker
wgProxyICE wgproxy.Proxy
wgProxyRelay wgproxy.Proxy
relayedConnRef *relayClient.Conn
handshaker *Handshaker
guard *guard.Guard
wg sync.WaitGroup
@@ -560,7 +561,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.mu.Lock()
defer conn.mu.Unlock()
if conn.ctx.Err() != nil {
if conn.ctx.Err() != nil || rci.relayedConn.Context().Err() != nil {
if err := rci.relayedConn.Close(); err != nil {
conn.Log.Warnf("failed to close unnecessary relayed connection: %v", err)
}
@@ -575,7 +576,9 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.Log.Errorf("failed to add relayed net.Conn to local proxy: %v", err)
return
}
wgProxy.SetDisconnectListener(conn.onRelayDisconnected)
wgProxy.SetDisconnectListener(func() {
conn.onRelayDisconnected(rci.relayedConn)
})
conn.dumpState.NewLocalProxy()
@@ -583,7 +586,7 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
if conn.isICEActive() {
conn.Log.Debugf("do not switch to relay because current priority is: %s", conn.currentConnPriority.String())
conn.setRelayedProxy(wgProxy)
conn.setRelayedProxy(wgProxy, rci.relayedConn)
conn.statusRelay.SetConnected()
conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, time.Now())
return
@@ -614,15 +617,26 @@ func (conn *Conn) onRelayConnectionIsReady(rci RelayConnInfo) {
conn.rosenpassRemoteKey = rci.rosenpassPubKey
conn.currentConnPriority = conntype.Relay
conn.statusRelay.SetConnected()
conn.setRelayedProxy(wgProxy)
conn.setRelayedProxy(wgProxy, rci.relayedConn)
conn.updateRelayStatus(rci.relayedConn.RemoteAddr().String(), rci.rosenpassPubKey, updateTime)
conn.Log.Infof("start to communicate with peer via relay")
conn.doOnConnected(rci.rosenpassPubKey, rci.rosenpassAddr, updateTime)
}
func (conn *Conn) onRelayDisconnected() {
// onRelayDisconnected reports the teardown of a relayed connection. relayedConn
// names the connection the signal belongs to, so a signal that arrives after
// its connection was replaced is ignored instead of tearing down its successor.
// A nil relayedConn means the caller does not track generations and the current
// connection is always torn down.
func (conn *Conn) onRelayDisconnected(relayedConn *relayClient.Conn) {
conn.mu.Lock()
defer conn.mu.Unlock()
if relayedConn != nil && conn.relayedConnRef != relayedConn {
conn.Log.Debugf("ignoring relay disconnect of a superseded connection")
return
}
conn.handleRelayDisconnectedLocked()
}
@@ -646,6 +660,7 @@ func (conn *Conn) handleRelayDisconnectedLocked() {
_ = conn.wgProxyRelay.CloseConn()
conn.wgProxyRelay = nil
}
conn.relayedConnRef = nil
changed := conn.statusRelay.Get() != worker.StatusDisconnected
if changed {
@@ -813,7 +828,8 @@ func (conn *Conn) evalStatus() ConnStatus {
//
// The result is a tri-state:
// - ConnStatusConnected: all available transports are up
// - ConnStatusPartiallyConnected: relay is up but ICE is still pending/reconnecting
// - ConnStatusPartiallyConnected: one transport carries the traffic and the other does
// not: relay up with ICE down, or ICE up with the shared relay transport down
// - ConnStatusDisconnected: no working transport
func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
defer func() {
@@ -830,13 +846,14 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
}
return evalConnStatus(connStatusInputs{
forceRelay: IsForceRelayed(),
peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(),
relayConnected: conn.statusRelay.Get() == worker.StatusConnected,
remoteSupportsICE: conn.handshaker.RemoteICESupported(),
iceWorkerCreated: iceWorkerCreated,
iceStatusConnecting: conn.statusICE.Get() != worker.StatusDisconnected,
iceInProgress: iceInProgress,
forceRelay: IsForceRelayed(),
peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(),
relayConnected: conn.statusRelay.Get() == worker.StatusConnected,
relayTransportConnected: conn.workerRelay.IsTransportConnected(),
remoteSupportsICE: conn.handshaker.RemoteICESupported(),
iceWorkerCreated: iceWorkerCreated,
iceStatusConnected: conn.statusICE.Get() == worker.StatusConnected,
iceInProgress: iceInProgress,
})
}
@@ -930,13 +947,14 @@ func (conn *Conn) logTraceConnState() {
}
}
func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy) {
func (conn *Conn) setRelayedProxy(proxy wgproxy.Proxy, relayedConn *relayClient.Conn) {
if conn.wgProxyRelay != nil {
if err := conn.wgProxyRelay.CloseConn(); err != nil {
conn.Log.Warnf("failed to close deprecated wg proxy conn: %v", err)
}
}
conn.wgProxyRelay = proxy
conn.relayedConnRef = relayedConn
}
// onWGHandshakeSuccess is called when the first WireGuard handshake is detected
@@ -1044,19 +1062,21 @@ func evalConnStatus(in connStatusInputs) guard.ConnStatus {
return boolToConnStatus(relayUsedAndUp)
}
// ICE counts as "up" when the status is anything other than Disconnected, OR
// when a negotiation is currently in progress (so we don't spam offers while one is in flight).
iceUp := in.iceStatusConnecting || in.iceInProgress
// ICE counts as "running" when either connected or attempting to connect.
iceRunning := in.iceStatusConnected || in.iceInProgress
// Relay side is acceptable if the peer doesn't rely on relay, or relay is connected.
relayOK := !in.peerUsesRelay || in.relayConnected
switch {
case iceUp && relayOK:
case iceRunning && relayOK:
return guard.ConnStatusConnected
case relayUsedAndUp:
// Relay is up but ICE is down — partially connected.
return guard.ConnStatusPartiallyConnected
case in.iceStatusConnected && !in.relayTransportConnected:
// ICE is up and the shared relay transport is down — offers cannot restore it.
return guard.ConnStatusPartiallyConnected
default:
return guard.ConnStatusDisconnected
}
+8 -7
View File
@@ -17,13 +17,14 @@ const (
// tri-state connection classification. Extracted so the decision logic can be unit-tested
// without constructing full Worker/Handshaker objects.
type connStatusInputs struct {
forceRelay bool // NB_FORCE_RELAY or JS/WASM
peerUsesRelay bool // remote peer advertises relay support AND local has relay
relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay)
remoteSupportsICE bool // remote peer sent ICE credentials
iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode)
iceStatusConnecting bool // statusICE is anything other than Disconnected
iceInProgress bool // a negotiation is currently in flight
forceRelay bool // NB_FORCE_RELAY or JS/WASM
peerUsesRelay bool // remote peer advertises relay support AND local has relay
relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay)
relayTransportConnected bool // the relay transport shared by all peers on that server is up
remoteSupportsICE bool // remote peer sent ICE credentials
iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode)
iceStatusConnected bool // statusICE reports Connected
iceInProgress bool // a negotiation is currently in flight
}
// ConnStatus describe the status of a peer's connection
+71 -12
View File
@@ -30,6 +30,21 @@ func TestEvalConnStatus_ForceRelay(t *testing.T) {
},
want: guard.ConnStatusDisconnected,
},
{
name: "force relay, relay up but the shared transport reports down",
in: connStatusInputs{
forceRelay: true,
peerUsesRelay: true,
relayConnected: true,
relayTransportConnected: false,
// The ICE inputs are set so that the force-relay return is the only branch
// that can produce Connected here: without it the peer would fall through to
// relayUsedAndUp and report PartiallyConnected.
remoteSupportsICE: true,
iceWorkerCreated: true,
},
want: guard.ConnStatusConnected,
},
{
name: "force relay, peer does NOT use relay - disconnected forever",
in: connStatusInputs{
@@ -123,24 +138,28 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = true
in.iceStatusConnecting = true
in.relayTransportConnected = true
in.iceStatusConnected = true
},
want: guard.ConnStatusConnected,
},
{
name: "ICE connected, peer does NOT use relay",
name: "ICE connected, peer does NOT use relay, shared transport down",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false
in.relayConnected = false
in.iceStatusConnecting = true
in.relayTransportConnected = false
in.iceStatusConnected = true
},
// A peer that does not rely on relay is unaffected by the shared transport:
// relayOK is true, so the first arm matches before the transport is considered.
want: guard.ConnStatusConnected,
},
{
name: "ICE InProgress only, peer does NOT use relay",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false
in.iceStatusConnecting = false
in.iceStatusConnected = false
in.iceInProgress = true
},
want: guard.ConnStatusConnected,
@@ -150,7 +169,8 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = true
in.iceStatusConnecting = false
in.relayTransportConnected = true
in.iceStatusConnected = false
in.iceInProgress = false
},
want: guard.ConnStatusPartiallyConnected,
@@ -160,21 +180,60 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false
in.relayConnected = false
in.iceStatusConnecting = false
in.iceStatusConnected = false
in.iceInProgress = false
},
want: guard.ConnStatusDisconnected,
},
{
name: "ICE up, peer uses relay but relay down -> partial (relay required, ICE ignored)",
name: "ICE connected, relay down for this peer but the shared transport is up -> disconnected",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.iceStatusConnecting = true
in.relayTransportConnected = true
in.iceStatusConnected = true
},
// The transport is fine, so the peer itself is unreachable over relay: it may have
// moved to another server, and only an offer carries its new relay address.
want: guard.ConnStatusDisconnected,
},
{
name: "ICE connected, the shared relay transport is down -> partial",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = true
},
// ICE carries the traffic and the relay transport is restored by the relay client's
// own guard, not by offers, so this must not trigger the aggressive retry.
want: guard.ConnStatusPartiallyConnected,
},
{
name: "ICE only negotiating while the shared relay transport is down -> disconnected",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = false
in.iceInProgress = true
},
// A negotiation in flight is not a working transport, so this peer has no path at
// all and must keep the aggressive retry. Calling it partially connected spends the
// ICE retry budget and parks the guard on the hourly ticker, and nothing wakes it
// when the negotiation then fails: onICEStateDisconnected is only reached once ICE
// has reached Connected (worker_ice.go onConnectionStateChange).
want: guard.ConnStatusDisconnected,
},
{
name: "ICE down and the shared relay transport is down -> disconnected",
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true
in.relayConnected = false
in.relayTransportConnected = false
in.iceStatusConnected = false
in.iceInProgress = false
},
// relayOK = false (peer uses relay but it's down), iceUp = true
// first switch arm fails (relayOK false), relayUsedAndUp = false (relay down),
// falls into default: Disconnected.
want: guard.ConnStatusDisconnected,
},
{
@@ -182,7 +241,7 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false
in.relayConnected = true // not actually used since peer doesn't rely on it
in.iceStatusConnecting = false
in.iceStatusConnected = false
in.iceInProgress = false
},
want: guard.ConnStatusDisconnected,
+5 -3
View File
@@ -14,7 +14,8 @@ type ConnStatus int
const (
// ConnStatusDisconnected means neither ICE nor Relay is connected.
ConnStatusDisconnected ConnStatus = iota
// ConnStatusPartiallyConnected means Relay is connected but ICE is not.
// ConnStatusPartiallyConnected means one transport is usable and the other is not:
// relay connected with ICE down, or ICE connected with the shared relay transport down.
ConnStatusPartiallyConnected
// ConnStatusConnected means all required connections are established.
ConnStatusConnected
@@ -87,8 +88,9 @@ func (g *Guard) SetICEConnDisconnected() {
// - Connected: no action, the peer is fully reachable.
// - Disconnected (neither ICE nor Relay): retries aggressively with exponential backoff (800ms doubling
// up to timeout), never gives up. This ensures rapid recovery when the peer has no connectivity at all.
// - PartiallyConnected (Relay up, ICE not): retries up to 3 times with exponential backoff, then switches
// to one attempt per hour. This limits signaling traffic when relay already provides connectivity.
// - PartiallyConnected (one transport usable, the other not): retries up to 3 times
// with exponential backoff, then switches to one attempt per hour. This limits
// signaling traffic while the peer still has a working path.
//
// External events (relay/ICE disconnect, signal/relay reconnect, candidate changes) reset the retry
// counter and backoff ticker, giving ICE a fresh chance after network conditions change.
+1 -4
View File
@@ -39,10 +39,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c
iceFailedTimeout := iceFailedTimeout()
iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait()
transportNet, err := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList)
if err != nil {
log.Errorf("failed to create pion's stdnet: %s", err)
}
transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList)
fac := logging.NewDefaultLoggerFactory()
+1 -1
View File
@@ -8,6 +8,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet"
)
func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) {
func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net {
return stdnet.NewNet(ctx, ifaceBlacklist)
}
+1 -1
View File
@@ -6,6 +6,6 @@ import (
"github.com/netbirdio/netbird/client/internal/stdnet"
)
func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) {
func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net {
return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist)
}
+20
View File
@@ -196,6 +196,7 @@ type Status struct {
muxRelays sync.RWMutex
peers map[string]State
ipToKey map[string]string
activeRoutePeers map[route.HAUniqueID]string
changeNotify map[string]map[string]*StatusChangeSubscription // map[peerID]map[subscriptionID]*StatusChangeSubscription
signalState bool
signalError error
@@ -257,6 +258,7 @@ func NewRecorder(mgmAddress string) *Status {
return &Status{
peers: make(map[string]State),
ipToKey: make(map[string]string),
activeRoutePeers: make(map[route.HAUniqueID]string),
changeNotify: make(map[string]map[string]*StatusChangeSubscription),
eventStreams: make(map[string]chan *proto.SystemEvent),
eventQueue: NewEventQueue(eventQueueSize),
@@ -481,6 +483,24 @@ func (d *Status) RemovePeerStateRoute(peer string, route string) error {
return nil
}
func (d *Status) AddActiveRoutePeer(haID route.HAUniqueID, peer string) {
d.mux.Lock()
defer d.mux.Unlock()
d.activeRoutePeers[haID] = peer
}
func (d *Status) RemoveActiveRoutePeer(haID route.HAUniqueID) {
d.mux.Lock()
defer d.mux.Unlock()
delete(d.activeRoutePeers, haID)
}
func (d *Status) GetActiveRoutePeers() map[route.HAUniqueID]string {
d.mux.RLock()
defer d.mux.RUnlock()
return maps.Clone(d.activeRoutePeers)
}
// CheckRoutes checks if the source and destination addresses are within the same route
// and returns the resource ID of the route that contains the addresses
func (d *Status) CheckRoutes(ip netip.Addr) ([]byte, bool) {
+23
View File
@@ -9,6 +9,8 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/route"
)
func TestAddPeer(t *testing.T) {
@@ -372,3 +374,24 @@ func TestMarkServerStateDoesNotNotifyWhenUnchanged(t *testing.T) {
status.MarkManagementDisconnected(err)
assert.False(t, notified(ch), "redundant disconnect should not notify")
}
func TestActiveRoutePeers(t *testing.T) {
status := NewRecorder("https://mgm")
netA := route.HAUniqueID("net-a-10.0.0.0/24")
netB := route.HAUniqueID("net-b-10.0.0.0/24")
status.AddActiveRoutePeer(netA, "peerA")
status.AddActiveRoutePeer(netB, "peerB")
active := status.GetActiveRoutePeers()
assert.Equal(t, "peerA", active[netA])
assert.Equal(t, "peerB", active[netB])
status.RemoveActiveRoutePeer(netA)
delete(active, netB)
active = status.GetActiveRoutePeers()
_, ok := active[netA]
assert.False(t, ok)
assert.Equal(t, "peerB", active[netB])
}
+17 -10
View File
@@ -121,11 +121,8 @@ func (w *WorkerICE) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
}
}
sessionID, err := NewICESessionID()
if err != nil {
w.log.Errorf("failed to create new session ID: %s", err)
}
w.sessionID = sessionID
// Keep the ID already advertised to the remote. Answers do not get a
// reply, so changing it here makes the next offer restart both sides.
w.abandonNegotiation()
}
@@ -205,6 +202,9 @@ func (w *WorkerICE) Close() {
w.muxAgent.Lock()
defer w.muxAgent.Unlock()
if w.agent != nil || w.agentConnecting {
w.renewSessionID()
}
if w.agent != nil {
w.agentDialerCancel()
if err := w.agent.Close(); err != nil {
@@ -366,16 +366,23 @@ func (w *WorkerICE) closeAgent(agent *icemaker.ThreadSafeAgent, cancel context.C
// Only the owner of the current session may reset its state: a stale dial
// goroutine waking after a newer attempt must not clobber it.
if w.agent == agent {
sessionID, err := NewICESessionID()
if err != nil {
w.log.Errorf("failed to create new session ID: %s", err)
}
w.sessionID = sessionID
w.renewSessionID()
w.abandonNegotiation()
}
return sessionChanged
}
// renewSessionID starts a new local session, so the remote treats our next offer
// or answer as a restart. Caller holds muxAgent.
func (w *WorkerICE) renewSessionID() {
sessionID, err := NewICESessionID()
if err != nil {
w.log.Errorf("failed to create new session ID: %s", err)
return
}
w.sessionID = sessionID
}
// abandonNegotiation drops all recorded ICE session state so the worker treats the
// next offer as a fresh start instead of a duplicate of a dead negotiation. The
// agent and agentConnecting flags must change together: leaving one stale wedges
@@ -0,0 +1,375 @@
package peer
import (
"context"
"fmt"
"net"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
icemaker "github.com/netbirdio/netbird/client/internal/peer/ice"
)
func TestWorkerICE_RemoteRestartPreservesAdvertisedSession(t *testing.T) {
w := newTestWorkerICE(t)
t.Cleanup(w.Close)
w.dialFunc = parkDial
advertised := w.SessionID()
remoteSession := ICESessionID("remote-first")
offer := OfferAnswer{
IceCredentials: IceCredentials{UFrag: "remoteufrag", Pwd: "remote-password-long-enough"},
SessionID: &remoteSession,
}
w.OnNewOffer(&offer)
require.True(t, w.InProgress(), "the first remote session must start ICE")
w.muxAgent.Lock()
firstAgent := w.agent
w.muxAgent.Unlock()
// The same callback handles answers. A changed remote ID must not create
// an unannounced local ID that makes the remote restart on our next offer.
secondSession := ICESessionID("remote-restarted")
answer := offer
answer.SessionID = &secondSession
w.OnNewOffer(&answer)
assert.Equal(t, advertised, w.SessionID(), "following a remote restart must keep our advertised ID")
w.muxAgent.Lock()
secondAgent := w.agent
w.muxAgent.Unlock()
assert.NotSame(t, firstAgent, secondAgent, "the changed remote session must still rebuild ICE")
w.OnNewOffer(&answer)
w.muxAgent.Lock()
defer w.muxAgent.Unlock()
assert.Same(t, secondAgent, w.agent, "a repeated answer must keep the replacement agent")
}
func TestWorkerICE_LocalCloseChangesAdvertisedSession(t *testing.T) {
w := newTestWorkerICE(t)
dialStarted := make(chan struct{})
dialDone := make(chan struct{})
w.dialFunc = func(ctx context.Context, _ *icemaker.ThreadSafeAgent, _ *OfferAnswer) (net.Conn, error) {
close(dialStarted)
defer close(dialDone)
<-ctx.Done()
return nil, ctx.Err()
}
session := ICESessionID("remote-session")
w.OnNewOffer(&OfferAnswer{
IceCredentials: IceCredentials{UFrag: "remoteufrag", Pwd: "remote-password-long-enough"},
SessionID: &session,
})
<-dialStarted
advertised := w.SessionID()
w.Close()
assert.NotEqual(t, advertised, w.SessionID(), "a local teardown must tell the remote to restart")
closedSession := w.SessionID()
// The abandoned dial goroutine cleans up after Close returned.
<-dialDone
assert.Never(t, func() bool { return w.SessionID() != closedSession }, 200*time.Millisecond, 10*time.Millisecond,
"the late cleanup of a closed negotiation must not restart again")
w.Close()
assert.Equal(t, closedSession, w.SessionID(), "closing an idle worker must not restart again")
}
// parkDial stands in for the ICE dial. It never connects and returns once the
// negotiation is abandoned, so a test decides when a negotiation fails.
func parkDial(ctx context.Context, _ *icemaker.ThreadSafeAgent, _ *OfferAnswer) (net.Conn, error) {
<-ctx.Done()
return nil, ctx.Err()
}
func newTestSessionID(t *testing.T) ICESessionID {
t.Helper()
sid, err := NewICESessionID()
require.NoError(t, err)
return sid
}
// handshakeSide is one end of a simulated signaling exchange.
type handshakeSide interface {
// message builds the offer or answer the side would send now.
message() OfferAnswer
// receive hands a remote offer or answer to the side's ICE logic.
receive(msg OfferAnswer)
// teardowns counts negotiations the side tore down to follow a remote restart.
teardowns() int
// failAgent ends the side's current negotiation as an ICE failure does.
failAgent()
}
// workerSide drives a real WorkerICE.
type workerSide struct {
t *testing.T
w *WorkerICE
replaced int
}
func newWorkerSide(t *testing.T) *workerSide {
t.Helper()
w := newTestWorkerICE(t)
w.dialFunc = parkDial
t.Cleanup(w.Close)
return &workerSide{t: t, w: w}
}
func (s *workerSide) message() OfferAnswer {
sid := s.w.SessionID()
ufrag, pwd := s.w.GetLocalUserCredentials()
return OfferAnswer{IceCredentials: IceCredentials{UFrag: ufrag, Pwd: pwd}, SessionID: &sid}
}
func (s *workerSide) receive(msg OfferAnswer) {
before := s.agent()
s.w.OnNewOffer(&msg)
if after := s.agent(); before != nil && after != before {
s.replaced++
}
}
func (s *workerSide) teardowns() int { return s.replaced }
func (s *workerSide) agent() *icemaker.ThreadSafeAgent {
s.w.muxAgent.Lock()
defer s.w.muxAgent.Unlock()
return s.w.agent
}
// failAgent runs the cleanup the dial goroutine or the Failed state callback
// performs when the current negotiation dies.
func (s *workerSide) failAgent() {
s.t.Helper()
s.w.muxAgent.Lock()
agent, cancel := s.w.agent, s.w.agentDialerCancel
s.w.muxAgent.Unlock()
require.NotNil(s.t, agent, "failing requires a running negotiation")
s.w.closeAgent(agent, cancel)
}
// legacySide models a remote peer running a release from before this change:
// when it follows a remote restart it also picks a new session ID of its own,
// which it announces only with its next offer or answer.
type legacySide struct {
t *testing.T
sessionID ICESessionID
remoteID ICESessionID
hasAgent bool
replaced int
}
func newLegacySide(t *testing.T) *legacySide {
return &legacySide{t: t, sessionID: newTestSessionID(t)}
}
func (s *legacySide) message() OfferAnswer {
sid := s.sessionID
return OfferAnswer{
IceCredentials: IceCredentials{UFrag: "legacyufrag", Pwd: "legacy-password-long-enough"},
SessionID: &sid,
}
}
func (s *legacySide) receive(msg OfferAnswer) {
if msg.SessionID == nil {
s.hasAgent = true
return
}
if s.hasAgent {
if *msg.SessionID == s.remoteID {
return
}
s.replaced++
s.sessionID = newTestSessionID(s.t)
}
s.hasAgent = true
s.remoteID = *msg.SessionID
}
func (s *legacySide) teardowns() int { return s.replaced }
func (s *legacySide) failAgent() {
s.hasAgent = false
s.remoteID = ""
s.sessionID = newTestSessionID(s.t)
}
// exchange runs one guard-driven round in the order Handshaker.Listen uses: the
// answerer handles the offer and answers with the session ID it holds
// afterwards, and the offerer handles the answer without replying.
func exchange(offerer, answerer handshakeSide) {
answerer.receive(offerer.message())
offerer.receive(answerer.message())
}
// offerPattern decides which side's guard sends the offer in a round.
type offerPattern struct {
name string
picker func(round int, local, remote handshakeSide) (offerer, answerer handshakeSide)
}
var offerPatterns = []offerPattern{
{
// A routing peer whose relay is down keeps offering on its own.
name: "local peer offers",
picker: func(_ int, local, remote handshakeSide) (handshakeSide, handshakeSide) {
return local, remote
},
},
{
name: "both peers offer",
picker: func(round int, local, remote handshakeSide) (handshakeSide, handshakeSide) {
if round%2 == 0 {
return local, remote
}
return remote, local
},
},
}
// assertSettles runs guard rounds and requires the pair to stop restarting
// each other: at most maxTeardowns in total, and none once half the rounds ran.
func assertSettles(t *testing.T, pattern offerPattern, local, remote handshakeSide, maxTeardowns int) {
t.Helper()
const rounds = 10
total := func() int { return local.teardowns() + remote.teardowns() }
start := total()
var halfway int
for round := range rounds {
if round == rounds/2 {
halfway = total()
}
offerer, answerer := pattern.picker(round, local, remote)
exchange(offerer, answerer)
}
assert.LessOrEqual(t, total()-start, maxTeardowns, "the peers must not keep restarting each other")
assert.Equal(t, halfway, total(), "the negotiation must be stable in the later rounds")
}
// establish runs the first offer and answer, so both sides negotiate.
func establish(t *testing.T, local, remote handshakeSide) {
t.Helper()
exchange(local, remote)
require.Zero(t, local.teardowns()+remote.teardowns(), "the first exchange must not restart anything")
}
func TestICESession_SettlesAfterAgentFailure(t *testing.T) {
sides := []struct {
name string
remote func(t *testing.T) handshakeSide
}{
{name: "current remote", remote: func(t *testing.T) handshakeSide { return newWorkerSide(t) }},
{name: "legacy remote", remote: func(t *testing.T) handshakeSide { return newLegacySide(t) }},
}
failures := []struct {
name string
fail func(local, remote handshakeSide)
}{
{name: "remote agent fails", fail: func(_, remote handshakeSide) { remote.failAgent() }},
{name: "local agent fails", fail: func(local, _ handshakeSide) { local.failAgent() }},
{name: "both agents fail", fail: func(local, remote handshakeSide) {
local.failAgent()
remote.failAgent()
}},
}
for _, side := range sides {
for _, failure := range failures {
for _, pattern := range offerPatterns {
t.Run(fmt.Sprintf("%s/%s/%s", side.name, failure.name, pattern.name), func(t *testing.T) {
local := newWorkerSide(t)
remote := side.remote(t)
establish(t, local, remote)
failure.fail(local, remote)
assertSettles(t, pattern, local, remote, 2)
})
}
}
}
}
// TestICESession_LocalCloseRestartsRemote covers an explicit teardown, as on a
// WireGuard handshake timeout. The remote must start over as well, or it keeps
// answering from the negotiation this side just abandoned.
func TestICESession_LocalCloseRestartsRemote(t *testing.T) {
for _, pattern := range offerPatterns {
t.Run(pattern.name, func(t *testing.T) {
local := newWorkerSide(t)
remote := newWorkerSide(t)
establish(t, local, remote)
local.w.Close()
assertSettles(t, pattern, local, remote, 1)
assert.Equal(t, 1, remote.teardowns(), "the remote must restart its negotiation exactly once")
})
}
}
func TestICESession_DuplicateMessagesKeepNegotiation(t *testing.T) {
local := newWorkerSide(t)
remote := newWorkerSide(t)
offer := local.message()
remote.receive(offer)
answer := remote.message()
local.receive(answer)
// Signaling may deliver the same message again, and a peer answers every
// offer, including repeats of one it already handled.
remote.receive(offer)
local.receive(answer)
local.receive(remote.message())
assert.Zero(t, local.teardowns(), "a repeated answer must not restart the negotiation")
assert.Zero(t, remote.teardowns(), "a repeated offer must not restart the negotiation")
}
// TestICESession_RemoteWithoutSessionIDKeepsNegotiation covers remote peers
// too old to send session IDs: once negotiating, their messages cannot tell a
// restart from a repeat, so they must not tear anything down.
func TestICESession_RemoteWithoutSessionIDKeepsNegotiation(t *testing.T) {
local := newWorkerSide(t)
unversioned := OfferAnswer{IceCredentials: IceCredentials{UFrag: "oldufrag", Pwd: "old-password-long-enough"}}
local.receive(unversioned)
require.NotNil(t, local.agent(), "a message without a session ID must still start ICE")
advertised := local.w.SessionID()
for range 3 {
local.receive(unversioned)
}
assert.Zero(t, local.teardowns(), "messages without a session ID must not restart the negotiation")
assert.Equal(t, advertised, local.w.SessionID(), "the advertised session must not change")
}
// TestWorkerICE_StaleCleanupKeepsAdvertisedSession covers the cleanup of a
// replaced negotiation finishing late, from its dial goroutine or its Closed
// state callback. It must neither pick a new session ID, an unannounced local
// restart, nor disturb the negotiation that replaced it.
func TestWorkerICE_StaleCleanupKeepsAdvertisedSession(t *testing.T) {
local := newWorkerSide(t)
remote := newWorkerSide(t)
establish(t, local, remote)
local.w.muxAgent.Lock()
oldAgent, oldCancel := local.w.agent, local.w.agentDialerCancel
local.w.muxAgent.Unlock()
remote.failAgent()
exchange(local, remote)
require.Equal(t, 1, local.teardowns(), "the local side must follow the remote restart")
advertised := local.w.SessionID()
current := local.agent()
local.w.closeAgent(oldAgent, oldCancel)
assert.Equal(t, advertised, local.w.SessionID(), "a stale cleanup must not change the advertised session")
assert.Same(t, current, local.agent(), "a stale cleanup must keep the current negotiation")
assertSettles(t, offerPatterns[1], local, remote, 0)
}
+17 -14
View File
@@ -3,7 +3,6 @@ package peer
import (
"context"
"errors"
"net"
"net/netip"
"sync"
"sync/atomic"
@@ -14,7 +13,7 @@ import (
)
type RelayConnInfo struct {
relayedConn net.Conn
relayedConn *relayClient.Conn
rosenpassPubKey []byte
rosenpassAddr string
}
@@ -27,7 +26,7 @@ type WorkerRelay struct {
conn *Conn
relayManager *relayClient.Manager
relayedConn net.Conn
relayedConn *relayClient.Conn
relayLock sync.Mutex
relaySupportedOnRemotePeer atomic.Bool
@@ -80,12 +79,7 @@ func (w *WorkerRelay) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
w.relayedConn = relayedConn
w.relayLock.Unlock()
err = w.relayManager.AddCloseListener(srv, w.onRelayClientDisconnected)
if err != nil {
log.Errorf("failed to add close listener: %s", err)
_ = relayedConn.Close()
return
}
go w.watchRelayedConn(relayedConn)
w.log.Debugf("peer conn opened via Relay: %s", srv)
go w.conn.onRelayConnectionIsReady(RelayConnInfo{
@@ -107,14 +101,21 @@ func (w *WorkerRelay) RelayIsSupportedLocally() bool {
return w.relayManager.HasRelayAddress()
}
func (w *WorkerRelay) IsTransportConnected() bool {
return w.relayManager.Ready()
}
func (w *WorkerRelay) CloseConn() {
w.relayLock.Lock()
defer w.relayLock.Unlock()
if w.relayedConn == nil {
conn := w.relayedConn
w.relayedConn = nil
w.relayLock.Unlock()
if conn == nil {
return
}
if err := w.relayedConn.Close(); err != nil {
if err := conn.Close(); err != nil {
w.log.Warnf("failed to close relay connection: %v", err)
}
}
@@ -133,6 +134,8 @@ func (w *WorkerRelay) preferredRelayServer(myRelayAddress, remoteRelayAddress st
return remoteRelayAddress
}
func (w *WorkerRelay) onRelayClientDisconnected() {
go w.conn.onRelayDisconnected()
func (w *WorkerRelay) watchRelayedConn(relayedConn *relayClient.Conn) {
<-relayedConn.Context().Done()
w.conn.onRelayDisconnected(relayedConn)
}
@@ -0,0 +1,76 @@
package profilemanager
import (
"fmt"
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// Regression test: a concurrent Get and Set of the ActiveProfileState will
// fail on Windows since the write is a temp file renamed over an open file.
// Windows will refuse to replace a file another handle holds open by default.
func TestActiveProfileState_ReadsDoNotBreakAConcurrentWrite(t *testing.T) {
withTempConfigDir(t, func(configDir string) {
withPatchedGlobals(t, configDir, func() {
sm := &ServiceManager{}
require.NoError(t, sm.CreateDefaultProfile())
require.NoError(t, sm.SetActiveProfileStateToDefault())
const switched = ID("0123456789abcdef0123456789abcdef")
const rounds = 50
var wg sync.WaitGroup
errs := make(chan error, 128)
for i := 0; i < 8; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for r := 0; r < rounds; r++ {
state, err := sm.GetActiveProfileState()
if err != nil {
errs <- fmt.Errorf("read: %w", err)
return
}
if state.ID != defaultProfileName && state.ID != switched {
errs <- fmt.Errorf("read: active profile is %q, which no writer wrote", state.ID)
return
}
}
}()
}
for i := 0; i < 2; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for r := 0; r < rounds; r++ {
id := switched
if r%2 == 0 {
id = defaultProfileName
}
if err := sm.SetActiveProfileState(&ActiveProfileState{ID: id, Username: "testuser"}); err != nil {
errs <- fmt.Errorf("switch: %w", err)
return
}
}
}()
}
wg.Wait()
close(errs)
for err := range errs {
assert.NoError(t, err, "a switch and a read of the active profile state must not collide")
}
state, err := sm.GetActiveProfileState()
require.NoError(t, err)
assert.Contains(t, []ID{defaultProfileName, switched}, state.ID,
"the file holds whichever switch landed last, not a mix of the two")
})
})
}
+2 -10
View File
@@ -201,11 +201,7 @@ func (p *StunTurnProbe) probeSTUN(ctx context.Context, uri *stun.URI) (addr stri
}
}()
net, err := stdnet.NewNet(ctx, nil)
if err != nil {
probeErr = fmt.Errorf("new net: %w", err)
return
}
net := stdnet.NewNet(ctx, nil)
client, err := stun.DialURI(uri, &stun.DialConfig{
Net: net,
@@ -290,11 +286,7 @@ func (p *StunTurnProbe) probeTURN(ctx context.Context, uri *stun.URI) (addr stri
}
}()
net, err := stdnet.NewNet(ctx, nil)
if err != nil {
probeErr = fmt.Errorf("new net: %w", err)
return
}
net := stdnet.NewNet(ctx, nil)
cfg := &turn.ClientConfig{
STUNServerAddr: turnServerAddr,
TURNServerAddr: turnServerAddr,
@@ -294,6 +294,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error {
return fmt.Errorf("add allowed IPs for peer %s: %w", route.Peer, err)
}
w.statusRecorder.AddActiveRoutePeer(route.GetHAUniqueID(), route.Peer)
if err := w.statusRecorder.AddPeerStateRoute(route.Peer, w.handler.String(), route.GetResourceID()); err != nil {
log.Warnf("Failed to update peer state: %v", err)
}
@@ -303,6 +304,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error {
}
func (w *Watcher) removeAllowedIPs(route *route.Route, rsn reason) error {
w.statusRecorder.RemoveActiveRoutePeer(route.GetHAUniqueID())
if err := w.statusRecorder.RemovePeerStateRoute(route.Peer, w.handler.String()); err != nil {
log.Warnf("Failed to update peer state: %v", err)
}
+2 -4
View File
@@ -8,6 +8,7 @@ import (
"net/netip"
"testing"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/internal/stdnet"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
@@ -406,10 +407,7 @@ func TestManagerUpdateRoutes(t *testing.T) {
for n, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
peerPrivateKey, _ := wgtypes.GeneratePrivateKey()
newNet, err := stdnet.NewNet(context.Background(), nil)
if err != nil {
t.Fatal(err)
}
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
opts := iface.WGIFaceOpts{
IFaceName: fmt.Sprintf("utun43%d", n),
Address: wgaddr.MustParseWGAddress("100.65.65.2/24"),
@@ -15,6 +15,7 @@ import (
"syscall"
"testing"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/internal/stdnet"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -436,8 +437,7 @@ func createWGInterface(t *testing.T, interfaceName, ipAddressCIDR string, listen
peerPrivateKey, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err)
newNet, err := stdnet.NewNet(context.Background(), nil)
require.NoError(t, err)
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
opts := iface.WGIFaceOpts{
IFaceName: interfaceName,
+47 -38
View File
@@ -45,7 +45,7 @@ type Net struct {
}
// NewNetWithDiscover creates a new StdNet instance.
func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) (*Net, error) {
func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) *Net {
if ctx == nil {
ctx = context.Background()
}
@@ -60,20 +60,19 @@ func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover
} else {
n.iFaceDiscover = newMobileIFaceDiscover(iFaceDiscover)
}
return n, n.UpdateInterfaces()
return n
}
// NewNet creates a new StdNet instance.
func NewNet(ctx context.Context, disallowList []string) (*Net, error) {
func NewNet(ctx context.Context, disallowList []string) *Net {
if ctx == nil {
ctx = context.Background()
}
n := &Net{
return &Net{
iFaceDiscover: pionDiscover{},
interfaceFilter: InterfaceFilter(disallowList),
ctx: ctx,
}
return n, n.UpdateInterfaces()
}
// resolveAddr performs DNS resolution with context support and timeout.
@@ -122,45 +121,18 @@ func (n *Net) resolveAddr(network, address string) (netip.AddrPort, error) {
return netip.AddrPortFrom(addrs[0], uint16(port)), nil
}
// UpdateInterfaces updates the internal list of network interfaces
// and associated addresses filtering them by name.
// The interfaces are discovered by an external iFaceDiscover function or by a default discoverer if the external one
// wasn't specified.
func (n *Net) UpdateInterfaces() (err error) {
n.mu.Lock()
defer n.mu.Unlock()
return n.updateInterfaces()
}
func (n *Net) updateInterfaces() (err error) {
allIfaces, err := n.iFaceDiscover.iFaces()
if err != nil {
return err
}
n.interfaces = n.filterInterfaces(allIfaces)
n.lastUpdate = time.Now()
return nil
}
// Interfaces returns a slice of interfaces which are available on the
// system
func (n *Net) Interfaces() ([]*transport.Interface, error) {
n.mu.Lock()
defer n.mu.Unlock()
if time.Since(n.lastUpdate) < updateInterval {
return slices.Clone(n.interfaces), nil
iFaces, err := n.freshInterfacesLocked()
if err != nil {
return nil, err
}
if err := n.updateInterfaces(); err != nil {
return nil, fmt.Errorf("update interfaces: %w", err)
}
return slices.Clone(n.interfaces), nil
return slices.Clone(iFaces), nil
}
// InterfaceByIndex returns the interface specified by index.
@@ -171,7 +143,13 @@ func (n *Net) Interfaces() ([]*transport.Interface, error) {
func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) {
n.mu.Lock()
defer n.mu.Unlock()
for _, ifc := range n.interfaces {
iFaces, err := n.freshInterfacesLocked()
if err != nil {
return nil, err
}
for _, ifc := range iFaces {
if ifc.Index == index {
return ifc, nil
}
@@ -184,7 +162,13 @@ func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) {
func (n *Net) InterfaceByName(name string) (*transport.Interface, error) {
n.mu.Lock()
defer n.mu.Unlock()
for _, ifc := range n.interfaces {
iFaces, err := n.freshInterfacesLocked()
if err != nil {
return nil, err
}
for _, ifc := range iFaces {
if ifc.Name == name {
return ifc, nil
}
@@ -193,6 +177,31 @@ func (n *Net) InterfaceByName(name string) (*transport.Interface, error) {
return nil, fmt.Errorf("%w: %s", transport.ErrInterfaceNotFound, name)
}
func (n *Net) freshInterfacesLocked() ([]*transport.Interface, error) {
if time.Since(n.lastUpdate) < updateInterval {
return n.interfaces, nil
}
if err := n.updateInterfacesLocked(); err != nil {
return nil, fmt.Errorf("update interfaces: %w", err)
}
return n.interfaces, nil
}
func (n *Net) updateInterfacesLocked() error {
allIFaces, err := n.iFaceDiscover.iFaces()
if err != nil {
return err
}
n.interfaces = n.filterInterfaces(allIFaces)
n.lastUpdate = time.Now()
return nil
}
func (n *Net) filterInterfaces(interfaces []*transport.Interface) []*transport.Interface {
if n.interfaceFilter == nil {
return interfaces
+136
View File
@@ -0,0 +1,136 @@
package stdnet
import (
"context"
"errors"
"net"
"testing"
"github.com/pion/transport/v3"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type countingDiscover struct {
calls int
list []*transport.Interface
err error
}
func (d *countingDiscover) iFaces() ([]*transport.Interface, error) {
d.calls++
if d.err != nil {
return nil, d.err
}
return d.list, nil
}
func newTestNet(t *testing.T, d iFaceDiscover) *Net {
t.Helper()
return &Net{
iFaceDiscover: d,
ctx: context.Background(),
}
}
func testIFace(index int, name string) *transport.Interface {
return transport.NewInterface(net.Interface{Index: index, Name: name})
}
func TestNet_InterfacesDiscoversLazilyAndCaches(t *testing.T) {
d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}}
n := newTestNet(t, d)
require.Zero(t, d.calls, "construction must not discover interfaces")
iFaces, err := n.Interfaces()
require.NoError(t, err)
require.Len(t, iFaces, 1)
assert.Equal(t, 1, d.calls)
_, err = n.Interfaces()
require.NoError(t, err)
assert.Equal(t, 1, d.calls)
}
func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) {
n := NewNet(context.Background(), nil)
require.NotNil(t, n)
assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold")
}
func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) {
n := NewNetWithDiscover(context.Background(), nil, nil)
require.NotNil(t, n)
assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold")
}
func TestNet_InterfacesRetryAfterDiscoveryFailure(t *testing.T) {
discoverErr := errors.New("discover failed")
d := &countingDiscover{err: discoverErr}
n := newTestNet(t, d)
_, err := n.Interfaces()
require.ErrorIs(t, err, discoverErr)
d.err = nil
d.list = []*transport.Interface{testIFace(1, "eth0")}
iFaces, err := n.Interfaces()
require.NoError(t, err)
require.Len(t, iFaces, 1)
assert.Equal(t, 2, d.calls)
}
func TestNet_InterfaceByNameRefreshes(t *testing.T) {
d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}}
n := newTestNet(t, d)
ifc, err := n.InterfaceByName("eth0")
require.NoError(t, err)
assert.Equal(t, "eth0", ifc.Name)
assert.Equal(t, 1, d.calls)
_, err = n.InterfaceByName("nope")
require.ErrorIs(t, err, transport.ErrInterfaceNotFound)
}
func TestNet_InterfaceByIndexRefreshes(t *testing.T) {
d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}}
n := newTestNet(t, d)
ifc, err := n.InterfaceByIndex(3)
require.NoError(t, err)
assert.Equal(t, "eth0", ifc.Name)
assert.Equal(t, 1, d.calls)
_, err = n.InterfaceByIndex(99)
require.ErrorIs(t, err, transport.ErrInterfaceNotFound)
}
func TestNet_InterfaceLookupPropagatesDiscoveryError(t *testing.T) {
discoverErr := errors.New("discover failed")
n := newTestNet(t, &countingDiscover{err: discoverErr})
_, err := n.InterfaceByName("eth0")
require.ErrorIs(t, err, discoverErr)
_, err = n.InterfaceByIndex(1)
require.ErrorIs(t, err, discoverErr)
}
func TestNet_InterfacesReturnsCopy(t *testing.T) {
d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}}
n := newTestNet(t, d)
iFaces, err := n.Interfaces()
require.NoError(t, err)
require.Len(t, iFaces, 1)
iFaces[0] = testIFace(2, "tampered")
iFaces, err = n.Interfaces()
require.NoError(t, err)
require.Len(t, iFaces, 1)
assert.Equal(t, "eth0", iFaces[0].Name)
}
@@ -0,0 +1,30 @@
// Package wincmd locates the Windows utilities the client shells out to.
package wincmd
import (
"path/filepath"
log "github.com/sirupsen/logrus"
"golang.org/x/sys/windows"
)
// defaultSystem32Dir is where the system directory is on every supported
// install, used only when the API that reports it fails.
const defaultSystem32Dir = `C:\Windows\System32`
// System32 returns the full path of a Windows utility under the system
// directory.
//
// PATH is deliberately not consulted. The daemon runs as LocalSystem with an
// environment of its own, so whoever can place an entry in that PATH chooses
// which binary runs with those privileges. The system directory is read from
// the API rather than from %SystemRoot% for the same reason.
func System32(command string) string {
sysDir, err := windows.GetSystemDirectory()
if err != nil {
log.Warnf("Failed to locate the Windows system directory, falling back to %s: %v", defaultSystem32Dir, err)
sysDir = defaultSystem32Dir
}
return filepath.Join(sysDir, command+".exe")
}
@@ -0,0 +1,31 @@
package wincmd
import (
"os"
"path/filepath"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestSystem32IgnoresPATH(t *testing.T) {
// A directory holding something that would win a PATH lookup, in front of
// everything else: the daemon runs as LocalSystem, so a PATH entry must not
// be able to decide what it executes.
planted := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(planted, "netsh.exe"), []byte("not really netsh"), 0o600))
t.Setenv("PATH", planted+string(os.PathListSeparator)+os.Getenv("PATH"))
got := System32("netsh")
assert.True(t, filepath.IsAbs(got), "the path must be absolute, got %q", got)
assert.NotContains(t, got, planted, "a PATH entry must not be consulted")
assert.True(t, strings.EqualFold(filepath.Base(got), "netsh.exe"), "unexpected file name in %q", got)
// The system directory is what Windows reports it to be, not %SystemRoot%,
// which the same caller could have set alongside PATH.
t.Setenv("SystemRoot", planted)
assert.Equal(t, got, System32("netsh"), "%SystemRoot% must not move the lookup")
}