Merge remote-tracking branch 'origin/main' into feat-post_quantum_ml_kem

This commit is contained in:
riccardom
2026-10-05 16:03:28 +02:00
85 changed files with 3037 additions and 624 deletions
+53 -1
View File
@@ -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)
}
}
+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")
}
+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()
+1 -4
View File
@@ -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,
+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)
+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 -9
View File
@@ -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)
}
+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
@@ -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) {
+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])
}
+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)
}