mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-06 21:49:08 +02:00
Merge remote-tracking branch 'origin/main' into feat-post_quantum_ml_kem
This commit is contained in:
@@ -380,9 +380,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)
|
||||
}
|
||||
@@ -1729,3 +1758,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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -2249,10 +2249,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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -161,7 +161,6 @@ func (m *MockWGIface) GetProxy() wgproxy.Proxy {
|
||||
return m.GetProxyFunc()
|
||||
}
|
||||
|
||||
|
||||
func (m *MockWGIface) GetNet() *netstack.Net {
|
||||
return m.GetNetFunc()
|
||||
}
|
||||
@@ -689,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,
|
||||
@@ -897,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),
|
||||
@@ -1500,3 +1493,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)
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -205,6 +205,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
|
||||
@@ -268,6 +269,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),
|
||||
@@ -492,6 +494,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) {
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user