mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-07 05:59:06 +02:00
Merge main into poc/certificate-posture
This commit is contained in:
+31
-11
@@ -104,8 +104,7 @@ type Client struct {
|
||||
|
||||
stateChangeMu sync.Mutex
|
||||
stateChangeSubID string
|
||||
eventSub *peer.EventSubscription
|
||||
// Closed to stop the watch goroutines from delivering buffered items to a
|
||||
// Closed to stop the watch goroutine from delivering buffered ticks to a
|
||||
// listener that has been removed or replaced. See stopStateChangeWatchLocked.
|
||||
stateChangeDone chan struct{}
|
||||
|
||||
@@ -213,6 +212,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
|
||||
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
|
||||
internal.WithNetEvents(c.netMgr))
|
||||
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
||||
connectClient.SetSyncResponsePersistence(true)
|
||||
// This path runs the interactive SSO flow, so reaching here means the peer
|
||||
// is authenticated again — release the latch Status() reports from. Clear
|
||||
// only once the fresh connect client is installed: until then Status()
|
||||
@@ -256,6 +256,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
|
||||
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
|
||||
internal.WithNetEvents(c.netMgr))
|
||||
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
||||
connectClient.SetSyncResponsePersistence(true)
|
||||
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
|
||||
}
|
||||
|
||||
@@ -327,6 +328,19 @@ func (c *Client) NotifyNetworkChange() {
|
||||
// or "strict"; strict also anonymizes internal IP ranges, peer names, and
|
||||
// WireGuard public keys, and implies anonymize.
|
||||
func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
|
||||
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, true)
|
||||
}
|
||||
|
||||
// DebugBundleFile generates a debug bundle and returns the path of the zip in
|
||||
// the cache directory instead of uploading it, so the app can hand the file to
|
||||
// the user for inspection. The caller owns the file and removes it once done;
|
||||
// the stale-bundle cleanup of later runs removes it only after a day.
|
||||
// anonymize and anonymizeLevel behave as in DebugBundle.
|
||||
func (c *Client) DebugBundleFile(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
|
||||
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, false)
|
||||
}
|
||||
|
||||
func (c *Client) debugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string, upload bool) (string, error) {
|
||||
cfg, cacheDir, cc := c.stateSnapshot()
|
||||
|
||||
// If the engine hasn't been started, load config from disk
|
||||
@@ -342,6 +356,11 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
|
||||
cacheDir = platformFiles.CacheDir()
|
||||
}
|
||||
|
||||
// Clear what an interrupted earlier run may have left in the cache before
|
||||
// adding to it. Remote debug jobs write to the same directory, so anything
|
||||
// younger than an hour is treated as possibly still in use.
|
||||
debug.RemoveStaleBundles(cacheDir, time.Hour)
|
||||
|
||||
deps := debug.GeneratorDependencies{
|
||||
InternalConfig: cfg,
|
||||
StatusRecorder: c.recorder,
|
||||
@@ -379,6 +398,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("generate debug bundle: %w", err)
|
||||
}
|
||||
if !upload {
|
||||
return debug.ExportBundle(path)
|
||||
}
|
||||
defer func() {
|
||||
if err := os.Remove(path); err != nil {
|
||||
log.Errorf("failed to remove debug bundle file: %v", err)
|
||||
@@ -475,6 +497,7 @@ func (c *Client) Networks() *NetworkArray {
|
||||
routesMap := routeManager.GetClientRoutesWithNetID()
|
||||
v6Merged := route.V6ExitMergeSet(routesMap)
|
||||
resolvedDomains := c.recorder.GetResolvedDomainsStates()
|
||||
activeRoutePeers := c.recorder.GetActiveRoutePeers()
|
||||
|
||||
networkArray := &NetworkArray{
|
||||
items: make([]Network, 0),
|
||||
@@ -488,7 +511,7 @@ func (c *Client) Networks() *NetworkArray {
|
||||
continue
|
||||
}
|
||||
|
||||
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged)
|
||||
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged, activeRoutePeers)
|
||||
if network == nil {
|
||||
continue
|
||||
}
|
||||
@@ -497,14 +520,14 @@ func (c *Client) Networks() *NetworkArray {
|
||||
return networkArray
|
||||
}
|
||||
|
||||
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}) *Network {
|
||||
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}, activeRoutePeers map[route.HAUniqueID]string) *Network {
|
||||
r := routes[0]
|
||||
netStr := r.Network.String()
|
||||
if r.IsDynamic() {
|
||||
netStr = r.Domains.SafeString()
|
||||
}
|
||||
|
||||
routePeer, err := c.findBestRoutePeer(routes)
|
||||
routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers)
|
||||
if err != nil {
|
||||
log.Errorf("could not get peer info for route %s: %v", id, err)
|
||||
return nil
|
||||
@@ -528,12 +551,9 @@ func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bo
|
||||
|
||||
// findBestRoutePeer returns the peer actively routing traffic for the given
|
||||
// HA route group. Falls back to the first connected peer, then the first peer.
|
||||
func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) {
|
||||
netStr := routes[0].Network.String()
|
||||
|
||||
fullStatus := c.recorder.GetFullStatus()
|
||||
for _, p := range fullStatus.Peers {
|
||||
if _, ok := p.GetRoutes()[netStr]; ok {
|
||||
func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) {
|
||||
if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok {
|
||||
if p, err := c.recorder.GetPeer(peerKey); err == nil {
|
||||
return p, nil
|
||||
}
|
||||
}
|
||||
|
||||
+10
-81
@@ -6,13 +6,8 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal"
|
||||
"github.com/netbirdio/netbird/client/internal/auth"
|
||||
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
|
||||
"github.com/netbirdio/netbird/client/internal/peer"
|
||||
cProto "github.com/netbirdio/netbird/client/proto"
|
||||
)
|
||||
|
||||
// StateChangeListener receives client state notifications.
|
||||
@@ -21,16 +16,11 @@ import (
|
||||
// changed: connection state, the run-loop status label (e.g. NeedsLogin) or
|
||||
// the session deadline. It mirrors the daemon's SubscribeStatus stream
|
||||
// trigger — on each signal the consumer pulls the fresh values via
|
||||
// Status() / SessionExpiresAtUnix().
|
||||
//
|
||||
// OnSessionExpiring forwards the engine's session-expiry warnings, fired at
|
||||
// sessionwatch.WarningLead before the deadline and again at FinalWarningLead
|
||||
// (finalWarning true). The second one is suppressed when the user dismissed
|
||||
// the first via DismissSessionWarning. The daemon turns the same events into
|
||||
// its tray notification.
|
||||
// Status() / SessionExpiresAtUnix(). The engine arms no expiry-warning
|
||||
// timers on Android; the app schedules the warnings from the deadline it
|
||||
// reads here.
|
||||
type StateChangeListener interface {
|
||||
OnStateChanged()
|
||||
OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool)
|
||||
}
|
||||
|
||||
// Status returns the connect run-loop's status label — the same value the
|
||||
@@ -110,11 +100,11 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
|
||||
return
|
||||
}
|
||||
|
||||
// Both subscriptions are buffered (one pending tick, ten pending events),
|
||||
// so unsubscribing is not enough to stop callbacks: the loops would drain
|
||||
// what is already queued and deliver it to a listener the caller has
|
||||
// already removed or replaced. Gate every callback on this registration's
|
||||
// own signal, which is closed before unsubscribing.
|
||||
// The subscription is buffered (one pending tick), so unsubscribing is
|
||||
// not enough to stop callbacks: the loop would drain what is already
|
||||
// queued and deliver it to a listener the caller has already removed or
|
||||
// replaced. Gate every callback on this registration's own signal, which
|
||||
// is closed before unsubscribing.
|
||||
done := make(chan struct{})
|
||||
c.stateChangeDone = done
|
||||
|
||||
@@ -133,9 +123,6 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
|
||||
listener.OnStateChanged()
|
||||
}
|
||||
}()
|
||||
|
||||
c.eventSub = c.recorder.SubscribeToEvents()
|
||||
go watchSessionWarnings(c.eventSub, listener, done)
|
||||
}
|
||||
|
||||
// RemoveStateChangeListener unregisters the state notification listener.
|
||||
@@ -145,21 +132,6 @@ func (c *Client) RemoveStateChangeListener() {
|
||||
c.stopStateChangeWatchLocked()
|
||||
}
|
||||
|
||||
// DismissSessionWarning records the user's "Dismiss" on the first expiry
|
||||
// warning and suppresses the final one for the current deadline. A refreshed
|
||||
// deadline re-arms both. No-op while the engine is not running.
|
||||
func (c *Client) DismissSessionWarning() {
|
||||
cc := c.getConnectClient()
|
||||
if cc == nil {
|
||||
return
|
||||
}
|
||||
engine := cc.Engine()
|
||||
if engine == nil {
|
||||
return
|
||||
}
|
||||
engine.DismissSessionWarning()
|
||||
}
|
||||
|
||||
// ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and
|
||||
// asks the management server to extend the session deadline. The tunnel is
|
||||
// untouched: no resync, no reconnect. Async; the result arrives on the
|
||||
@@ -201,8 +173,8 @@ func (c *Client) CancelExtendAuthSession() {
|
||||
}
|
||||
|
||||
func (c *Client) stopStateChangeWatchLocked() {
|
||||
// Signal first, unsubscribe second: closing the channels only stops new
|
||||
// items, and the loops would still hand whatever is buffered to a listener
|
||||
// Signal first, unsubscribe second: closing the channel only stops new
|
||||
// items, and the loop would still hand whatever is buffered to a listener
|
||||
// that is no longer registered.
|
||||
if c.stateChangeDone != nil {
|
||||
close(c.stateChangeDone)
|
||||
@@ -212,49 +184,6 @@ func (c *Client) stopStateChangeWatchLocked() {
|
||||
c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID)
|
||||
c.stateChangeSubID = ""
|
||||
}
|
||||
if c.eventSub != nil {
|
||||
// Closes the channel, which ends watchSessionWarnings.
|
||||
c.recorder.UnsubscribeFromEvents(c.eventSub)
|
||||
c.eventSub = nil
|
||||
}
|
||||
}
|
||||
|
||||
// watchSessionWarnings forwards the engine's session-expiry warnings to the
|
||||
// listener. The event stream also carries unrelated traffic — network-map
|
||||
// updates on every sync, DNS and route errors — so everything but an
|
||||
// AUTHENTICATION event carrying the session-warning marker is dropped. Exits
|
||||
// when the subscription is closed by UnsubscribeFromEvents, or earlier when
|
||||
// done is closed — the stream buffers up to ten events, and a deregistered
|
||||
// listener must not receive the ones already queued.
|
||||
func watchSessionWarnings(sub *peer.EventSubscription, listener StateChangeListener, done <-chan struct{}) {
|
||||
for ev := range sub.Events() {
|
||||
select {
|
||||
case <-done:
|
||||
return
|
||||
default:
|
||||
}
|
||||
if ev.GetCategory() != cProto.SystemEvent_AUTHENTICATION {
|
||||
continue
|
||||
}
|
||||
meta := ev.GetMetadata()
|
||||
if meta[sessionwatch.MetaSessionWarning] != "true" {
|
||||
// Other AUTHENTICATION events exist (e.g. a deadline rejected as
|
||||
// out of range); they carry no warning marker.
|
||||
continue
|
||||
}
|
||||
deadline, err := sessionwatch.ParseExpiresAt(meta[sessionwatch.MetaSessionExpiresAt])
|
||||
if err != nil {
|
||||
log.Warnf("session warning event with unparsable deadline: %v", err)
|
||||
continue
|
||||
}
|
||||
lead, err := sessionwatch.ParseLeadMinutes(meta[sessionwatch.MetaSessionLeadMinutes])
|
||||
if err != nil {
|
||||
// Informational only — the deadline above is what drives the UI.
|
||||
lead = 0
|
||||
}
|
||||
listener.OnSessionExpiring(deadline.Unix(), int64(lead),
|
||||
meta[sessionwatch.MetaSessionFinal] == "true")
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) beginExtend() (context.Context, error) {
|
||||
|
||||
@@ -31,6 +31,8 @@ const (
|
||||
// PasswordRequiredMarker tells Java to prompt for a password and retry. It is
|
||||
// a string because gomobile flattens errors to their message, so a sentinel
|
||||
// value would not survive the binding.
|
||||
//
|
||||
//nolint:gosec // G101 false positive: a sentinel marker, not a credential
|
||||
const PasswordRequiredMarker = "netbird-ssh-password-required"
|
||||
|
||||
// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation,
|
||||
|
||||
+58
-26
@@ -23,7 +23,10 @@ import (
|
||||
"github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
const errCloseConnection = "Failed to close connection: %v"
|
||||
const (
|
||||
errCloseConnection = "Failed to close connection: %v"
|
||||
noUpDownFlag = "no-updown"
|
||||
)
|
||||
|
||||
var (
|
||||
logFileCount uint32
|
||||
@@ -257,13 +260,14 @@ func runForDuration(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
|
||||
stateWasDown := stat.Status != string(internal.StatusConnected) && stat.Status != string(internal.StatusConnecting)
|
||||
noUpDown, _ := cmd.Flags().GetBool(noUpDownFlag)
|
||||
|
||||
initialLogLevel, err := client.GetLogLevel(cmd.Context(), &proto.GetLogLevelRequest{})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get log level: %v", status.Convert(err).Message())
|
||||
}
|
||||
|
||||
if stateWasDown {
|
||||
if stateWasDown && !noUpDown {
|
||||
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
|
||||
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
|
||||
} else {
|
||||
@@ -284,34 +288,20 @@ func runForDuration(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
|
||||
needsRestoreUp := false
|
||||
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
|
||||
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message())
|
||||
if noUpDown {
|
||||
enableSyncResponsePersistence(cmd, client)
|
||||
} else {
|
||||
needsRestoreUp = !stateWasDown
|
||||
cmd.Println("netbird down")
|
||||
needsRestoreUp = restartDaemon(cmd, client, stateWasDown)
|
||||
}
|
||||
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
// Enable sync response persistence before bringing the service up
|
||||
if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message())
|
||||
}
|
||||
|
||||
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
|
||||
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
|
||||
} else {
|
||||
needsRestoreUp = false
|
||||
cmd.Println("netbird up")
|
||||
}
|
||||
|
||||
time.Sleep(3 * time.Second)
|
||||
|
||||
cpuProfilingStarted := false
|
||||
if _, err := client.StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil {
|
||||
cmd.PrintErrf("Failed to start CPU profiling: %v\n", err)
|
||||
if msg := status.Convert(err).Message(); strings.Contains(msg, "already in progress") {
|
||||
cmd.PrintErrln("CPU profiling is already running (started with `netbird debug cpu start`). " +
|
||||
"It is left running and is included in a bundle created after `netbird debug cpu stop`.")
|
||||
} else {
|
||||
cmd.PrintErrf("Failed to start CPU profiling: %v\n", msg)
|
||||
}
|
||||
} else {
|
||||
cpuProfilingStarted = true
|
||||
defer func() {
|
||||
@@ -401,7 +391,7 @@ func runForDuration(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
}
|
||||
|
||||
if stateWasDown {
|
||||
if stateWasDown && !noUpDown {
|
||||
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
|
||||
cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message())
|
||||
} else {
|
||||
@@ -458,6 +448,47 @@ func setSyncResponsePersistence(cmd *cobra.Command, args []string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// enableSyncResponsePersistence asks the daemon to keep the latest sync
|
||||
// response so the bundle carries the network map. With a running daemon only
|
||||
// syncs received after the call are kept.
|
||||
func enableSyncResponsePersistence(cmd *cobra.Command, client proto.DaemonServiceClient) {
|
||||
if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{
|
||||
Enabled: true,
|
||||
}); err != nil {
|
||||
cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message())
|
||||
}
|
||||
}
|
||||
|
||||
// restartDaemon cycles the daemon down and up with sync response persistence
|
||||
// enabled so the bundle carries the network map. It reports whether the
|
||||
// daemon was left down although it was running before, so the caller can
|
||||
// bring it back up.
|
||||
func restartDaemon(cmd *cobra.Command, client proto.DaemonServiceClient, stateWasDown bool) bool {
|
||||
needsRestoreUp := false
|
||||
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
|
||||
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message())
|
||||
} else {
|
||||
needsRestoreUp = !stateWasDown
|
||||
cmd.Println("netbird down")
|
||||
}
|
||||
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
// Enable sync response persistence before bringing the service up
|
||||
enableSyncResponsePersistence(cmd, client)
|
||||
|
||||
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
|
||||
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
|
||||
} else {
|
||||
needsRestoreUp = false
|
||||
cmd.Println("netbird up")
|
||||
}
|
||||
|
||||
time.Sleep(3 * time.Second)
|
||||
|
||||
return needsRestoreUp
|
||||
}
|
||||
|
||||
func waitForDurationOrCancel(ctx context.Context, duration time.Duration, cmd *cobra.Command) error {
|
||||
ticker := time.NewTicker(1 * time.Second)
|
||||
defer ticker.Stop()
|
||||
@@ -546,4 +577,5 @@ func init() {
|
||||
forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle")
|
||||
forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root")
|
||||
forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle")
|
||||
forCmd.Flags().Bool(noUpDownFlag, false, "Keep the daemon running instead of bringing it down and up before collecting. The bundle only includes the network map if a sync arrives during the run")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/spf13/cobra"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
)
|
||||
|
||||
var debugCPUCmd = &cobra.Command{
|
||||
Use: "cpu",
|
||||
Short: "Profile the daemon's CPU usage",
|
||||
Long: `Starts and stops CPU profiling in the running daemon without restarting it.
|
||||
The profile is included in the next debug bundle as cpu.prof.
|
||||
|
||||
Profiling is not time limited: it keeps running, and keeps costing CPU, until
|
||||
"netbird debug cpu stop" is run.`,
|
||||
}
|
||||
|
||||
var debugCPUStartCmd = &cobra.Command{
|
||||
Use: "start",
|
||||
Short: "Start CPU profiling in the daemon",
|
||||
Example: " netbird debug cpu start",
|
||||
Args: cobra.NoArgs,
|
||||
RunE: debugCPUStart,
|
||||
}
|
||||
|
||||
var debugCPUStopCmd = &cobra.Command{
|
||||
Use: "stop",
|
||||
Short: "Stop CPU profiling in the daemon",
|
||||
Long: `Stops CPU profiling. The captured profile stays in the daemon until the next
|
||||
debug bundle is created, which includes it as cpu.prof.`,
|
||||
Example: " netbird debug cpu stop && netbird debug bundle",
|
||||
Args: cobra.NoArgs,
|
||||
RunE: debugCPUStop,
|
||||
}
|
||||
|
||||
func debugCPUStart(cmd *cobra.Command, _ []string) error {
|
||||
conn, err := getClient(cmd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
if err := conn.Close(); err != nil {
|
||||
log.Errorf(errCloseConnection, err)
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := proto.NewDaemonServiceClient(conn).StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil {
|
||||
return fmt.Errorf("start CPU profiling: %v", status.Convert(err).Message())
|
||||
}
|
||||
|
||||
cmd.Println("CPU profiling started and runs until stopped. Run `netbird debug cpu stop` and then `netbird debug bundle` to collect it.")
|
||||
return nil
|
||||
}
|
||||
|
||||
func debugCPUStop(cmd *cobra.Command, _ []string) error {
|
||||
conn, err := getClient(cmd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
if err := conn.Close(); err != nil {
|
||||
log.Errorf(errCloseConnection, err)
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := proto.NewDaemonServiceClient(conn).StopCPUProfile(cmd.Context(), &proto.StopCPUProfileRequest{}); err != nil {
|
||||
return fmt.Errorf("stop CPU profiling: %v", status.Convert(err).Message())
|
||||
}
|
||||
|
||||
cmd.Println("CPU profiling stopped. Run `netbird debug bundle` to include cpu.prof.")
|
||||
return nil
|
||||
}
|
||||
|
||||
func init() {
|
||||
debugCPUCmd.AddCommand(debugCPUStartCmd)
|
||||
debugCPUCmd.AddCommand(debugCPUStopCmd)
|
||||
debugCmd.AddCommand(debugCPUCmd)
|
||||
}
|
||||
@@ -0,0 +1,164 @@
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"os/user"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/spf13/pflag"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal"
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
)
|
||||
|
||||
// startDebugTestDaemon starts an in-process daemon with an isolated profile
|
||||
// directory and returns the address the CLI should dial.
|
||||
func startDebugTestDaemon(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
origDefaultProfileDir := profilemanager.DefaultConfigPathDir
|
||||
origActiveProfileStatePath := profilemanager.ActiveProfileStatePath
|
||||
origConfigDirOverride := profilemanager.ConfigDirOverride
|
||||
origDaemonAddr := daemonAddr
|
||||
t.Cleanup(func() {
|
||||
profilemanager.DefaultConfigPathDir = origDefaultProfileDir
|
||||
profilemanager.ActiveProfileStatePath = origActiveProfileStatePath
|
||||
profilemanager.ConfigDirOverride = origConfigDirOverride
|
||||
daemonAddr = origDaemonAddr
|
||||
})
|
||||
|
||||
profilemanager.DefaultConfigPathDir = tempDir
|
||||
profilemanager.ActiveProfileStatePath = tempDir + "/active_profile.json"
|
||||
profilemanager.ConfigDirOverride = tempDir
|
||||
|
||||
currUser, err := user.Current()
|
||||
require.NoError(t, err)
|
||||
sm := profilemanager.ServiceManager{}
|
||||
created, err := sm.AddProfile("test1", currUser.Username)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, sm.SetActiveProfileState(&profilemanager.ActiveProfileState{
|
||||
ID: created.ID,
|
||||
Username: currUser.Username,
|
||||
}))
|
||||
|
||||
ctx, cancel := context.WithCancel(internal.CtxInitState(context.Background()))
|
||||
srv, lis := startClientDaemon(t, ctx, "", tempDir+"/config.json")
|
||||
t.Cleanup(func() {
|
||||
cancel()
|
||||
srv.Stop()
|
||||
})
|
||||
|
||||
return "tcp://" + lis.Addr().String()
|
||||
}
|
||||
|
||||
// runDebugCmd runs `netbird debug <args>` against the daemon at addr and
|
||||
// returns everything the command printed.
|
||||
func runDebugCmd(addr string, args ...string) (string, error) {
|
||||
daemonAddr = addr
|
||||
var out bytes.Buffer
|
||||
rootCmd.SetOut(&out)
|
||||
rootCmd.SetErr(&out)
|
||||
rootCmd.SetArgs(append(append([]string{"debug"}, args...), "--daemon-addr", addr, "--log-file", ""))
|
||||
err := rootCmd.Execute()
|
||||
rootCmd.SetOut(nil)
|
||||
rootCmd.SetErr(nil)
|
||||
rootCmd.SetArgs(nil)
|
||||
resetFlags(rootCmd)
|
||||
return out.String(), err
|
||||
}
|
||||
|
||||
// resetFlags puts every flag of the command and its subcommands back to its
|
||||
// default so a value parsed in one run does not leak into the next in-process
|
||||
// execution.
|
||||
func resetFlags(cmd *cobra.Command) {
|
||||
reset := func(f *pflag.Flag) {
|
||||
// Set appends to a slice flag and would parse the "[a,b]" default
|
||||
// text as elements, so slices are replaced instead.
|
||||
if sv, ok := f.Value.(pflag.SliceValue); ok {
|
||||
var def []string
|
||||
if trimmed := strings.Trim(f.DefValue, "[]"); trimmed != "" {
|
||||
def = strings.Split(trimmed, ",")
|
||||
}
|
||||
_ = sv.Replace(def)
|
||||
} else {
|
||||
_ = f.Value.Set(f.DefValue)
|
||||
}
|
||||
f.Changed = false
|
||||
}
|
||||
cmd.Flags().VisitAll(reset)
|
||||
cmd.PersistentFlags().VisitAll(reset)
|
||||
// Commands pin their writers to the buffer of the run that first used
|
||||
// them, so a later run would print into the old buffer.
|
||||
cmd.SetOut(nil)
|
||||
cmd.SetErr(nil)
|
||||
for _, sub := range cmd.Commands() {
|
||||
resetFlags(sub)
|
||||
}
|
||||
}
|
||||
|
||||
// TestResetFlagsSliceDefault guards against Set("[]") on slice flags, which
|
||||
// stores a literal "[]" element instead of the empty default.
|
||||
func TestResetFlagsSliceDefault(t *testing.T) {
|
||||
cmd := &cobra.Command{Use: "x"}
|
||||
var env, withDefault []string
|
||||
cmd.Flags().StringSliceVar(&env, "env", nil, "")
|
||||
cmd.Flags().StringSliceVar(&withDefault, "names", []string{"a", "b"}, "")
|
||||
require.NoError(t, cmd.Flags().Parse([]string{"--env", "K=V", "--names", "c"}))
|
||||
|
||||
resetFlags(cmd)
|
||||
|
||||
assert.Empty(t, env, "slice flag with no default must reset to empty")
|
||||
assert.Equal(t, []string{"a", "b"}, withDefault, "slice flag must reset to its default")
|
||||
}
|
||||
|
||||
func TestDebugCPUStartStop(t *testing.T) {
|
||||
addr := startDebugTestDaemon(t)
|
||||
|
||||
run := func(args ...string) error {
|
||||
_, err := runDebugCmd(addr, append([]string{"cpu"}, args...)...)
|
||||
return err
|
||||
}
|
||||
|
||||
require.Error(t, run("stop"), "stop without a running profile must fail")
|
||||
require.NoError(t, run("start"))
|
||||
assert.Error(t, run("start"), "second start must be rejected while profiling")
|
||||
require.NoError(t, run("stop"))
|
||||
assert.Error(t, run("stop"), "second stop must be rejected")
|
||||
assert.NoError(t, run("start"), "profiling can be started again after a stop")
|
||||
assert.NoError(t, run("stop"))
|
||||
}
|
||||
|
||||
// TestDebugForKeepsRunningCPUProfile covers `debug for` started while a
|
||||
// profile from `debug cpu start` is running: it must say so, leave the
|
||||
// profile alone, and still create the bundle.
|
||||
func TestDebugForKeepsRunningCPUProfile(t *testing.T) {
|
||||
addr := startDebugTestDaemon(t)
|
||||
|
||||
_, err := runDebugCmd(addr, "cpu", "start")
|
||||
require.NoError(t, err)
|
||||
|
||||
out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown")
|
||||
require.NoError(t, err, "output: %s", out)
|
||||
assert.Contains(t, out, "CPU profiling is already running", "the conflict must be explained")
|
||||
assert.NotContains(t, out, "rpc error", "the raw RPC error must not reach the user")
|
||||
assert.Contains(t, out, "Local file:", "the bundle must still be created")
|
||||
|
||||
_, err = runDebugCmd(addr, "cpu", "stop")
|
||||
assert.NoError(t, err, "the profile started by the user must still be running")
|
||||
}
|
||||
|
||||
func TestDebugForNoUpDown(t *testing.T) {
|
||||
addr := startDebugTestDaemon(t)
|
||||
|
||||
out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown")
|
||||
require.NoError(t, err, "output: %s", out)
|
||||
assert.NotContains(t, out, "netbird down", "--no-updown must not bring the daemon down")
|
||||
assert.NotContains(t, out, "netbird up", "--no-updown must not bring the daemon up")
|
||||
assert.Contains(t, out, "Local file:", "the bundle must still be created")
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
@@ -25,6 +26,30 @@ var serviceCmd = &cobra.Command{
|
||||
|
||||
const defaultJSONSocket = "unix:///var/run/netbird-http.sock"
|
||||
|
||||
// forbiddenServiceEnvVars are the environment variables the service is never
|
||||
// registered with, keyed in upper case since these are Windows names. Each one
|
||||
// decides where the daemon resolves something it then uses with the privileges
|
||||
// of the account it runs under — LocalSystem on Windows, root elsewhere: the
|
||||
// executables it runs (PATH, PATHEXT, COMSPEC, SystemRoot, windir) or the
|
||||
// directory it writes temporary files in (TEMP, TMP). The daemon needs none of
|
||||
// them, and the utilities it shells out to are resolved by absolute path.
|
||||
var forbiddenServiceEnvVars = map[string]struct{}{
|
||||
"PATH": {},
|
||||
"PATHEXT": {},
|
||||
"SYSTEMROOT": {},
|
||||
"WINDIR": {},
|
||||
"COMSPEC": {},
|
||||
"TEMP": {},
|
||||
"TMP": {},
|
||||
}
|
||||
|
||||
// forbiddenServiceEnvPrefixes are the dynamic-loader families, refused whole
|
||||
// rather than by name: LD_PRELOAD, DYLD_INSERT_LIBRARIES and their siblings all
|
||||
// reach the loader of the process, the set differs per platform and libc, and
|
||||
// new members arrive with new OS releases. Listing them one by one is a list
|
||||
// that is wrong the moment it is written.
|
||||
var forbiddenServiceEnvPrefixes = []string{"LD_", "DYLD_"}
|
||||
|
||||
var (
|
||||
serviceName string
|
||||
serviceEnvVars []string
|
||||
@@ -127,8 +152,33 @@ func parseServiceEnvVars(envVars []string) (map[string]string, error) {
|
||||
return nil, fmt.Errorf("empty environment variable key in: %s", env)
|
||||
}
|
||||
|
||||
if isForbiddenServiceEnvVar(key) {
|
||||
return nil, fmt.Errorf("environment variable %s cannot be set on the service: it decides where the service resolves the executables, libraries or temporary files it uses", key)
|
||||
}
|
||||
|
||||
envMap[key] = value
|
||||
}
|
||||
|
||||
return envMap, nil
|
||||
}
|
||||
|
||||
// isForbiddenServiceEnvVar reports whether name is one the service must not be
|
||||
// registered with.
|
||||
//
|
||||
// The names are matched case-insensitively only on Windows, where they are the
|
||||
// same variable however they are spelled. Elsewhere the environment is
|
||||
// case-sensitive, so Path and PATH are two different variables and only the
|
||||
// exact spelling is the one the loader reads.
|
||||
func isForbiddenServiceEnvVar(name string) bool {
|
||||
if runtime.GOOS == "windows" {
|
||||
name = strings.ToUpper(name)
|
||||
}
|
||||
|
||||
if _, forbidden := forbiddenServiceEnvVars[name]; forbidden {
|
||||
return true
|
||||
}
|
||||
|
||||
return slices.ContainsFunc(forbiddenServiceEnvPrefixes, func(prefix string) bool {
|
||||
return strings.HasPrefix(name, prefix)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/client/configs"
|
||||
"github.com/netbirdio/netbird/client/internal/daemonaddr"
|
||||
"github.com/netbirdio/netbird/client/internal/elevate"
|
||||
"github.com/netbirdio/netbird/util"
|
||||
)
|
||||
|
||||
@@ -43,10 +44,33 @@ func serviceParamsPath() string {
|
||||
|
||||
// loadServiceParams reads saved service parameters from disk.
|
||||
// Returns nil with no error if the file does not exist.
|
||||
//
|
||||
// The file is read by an elevated install and decides the arguments and the
|
||||
// environment of the service it then registers, so it is used only when its
|
||||
// ownership and permissions are the ones saveServiceParams leaves behind. That
|
||||
// restricted ACL is applied when the file is written, which is not necessarily
|
||||
// before it is first read, so this is checked rather than assumed. A file that
|
||||
// fails the check is treated as absent, and the install proceeds with its
|
||||
// defaults.
|
||||
func loadServiceParams() (*serviceParams, error) {
|
||||
path := serviceParamsPath()
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
// Resolve links first so the checks apply to the file that is actually read.
|
||||
// Since the check covers every directory above it as well, nobody who fails
|
||||
// it can swap the file between here and the read below.
|
||||
resolved, err := filepath.EvalSymlinks(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, nil //nolint:nilnil
|
||||
}
|
||||
return nil, fmt.Errorf("resolve service params %s: %w", path, err)
|
||||
}
|
||||
|
||||
if err := elevate.CheckOnlyOwnerWritable(resolved); err != nil {
|
||||
return nil, fmt.Errorf("refusing to read service params from %s: %w", resolved, err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(resolved)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, nil //nolint:nilnil
|
||||
@@ -182,10 +206,16 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
|
||||
// If --service-env was explicitly set to empty, all saved env vars are cleared.
|
||||
// If --service-env was not set, saved env vars are used entirely.
|
||||
func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
|
||||
// A forbidden name explicitly passed on the command line is an error the
|
||||
// operator is told about, but one restored from a file written by an older
|
||||
// version is dropped: an install that refuses to run would leave the host
|
||||
// without a daemon over a variable nobody is asking for any more.
|
||||
saved := dropForbiddenServiceEnvVars(cmd, params.ServiceEnvVars)
|
||||
|
||||
if !cmd.Flags().Changed("service-env") {
|
||||
if len(params.ServiceEnvVars) > 0 {
|
||||
if len(saved) > 0 {
|
||||
// No explicit env vars: rebuild serviceEnvVars from saved params.
|
||||
serviceEnvVars = envMapToSlice(params.ServiceEnvVars)
|
||||
serviceEnvVars = envMapToSlice(saved)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -204,13 +234,13 @@ func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
|
||||
return
|
||||
}
|
||||
|
||||
if len(params.ServiceEnvVars) == 0 {
|
||||
if len(saved) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// Merge saved values underneath explicit ones.
|
||||
merged := make(map[string]string, len(params.ServiceEnvVars)+len(explicit))
|
||||
maps.Copy(merged, params.ServiceEnvVars)
|
||||
merged := make(map[string]string, len(saved)+len(explicit))
|
||||
maps.Copy(merged, saved)
|
||||
maps.Copy(merged, explicit) // explicit wins on conflict
|
||||
serviceEnvVars = envMapToSlice(merged)
|
||||
}
|
||||
@@ -233,6 +263,20 @@ var resetParamsCmd = &cobra.Command{
|
||||
},
|
||||
}
|
||||
|
||||
// dropForbiddenServiceEnvVars returns the saved entries that may still be
|
||||
// registered on the service, reporting every one it leaves behind.
|
||||
func dropForbiddenServiceEnvVars(cmd *cobra.Command, saved map[string]string) map[string]string {
|
||||
kept := make(map[string]string, len(saved))
|
||||
for key, value := range saved {
|
||||
if isForbiddenServiceEnvVar(key) {
|
||||
cmd.PrintErrf("Warning: ignoring saved service environment variable %s: it decides where the service resolves the executables, libraries or temporary files it uses\n", key)
|
||||
continue
|
||||
}
|
||||
kept[key] = value
|
||||
}
|
||||
return kept
|
||||
}
|
||||
|
||||
// envMapToSlice converts a map of env vars to a KEY=VALUE slice.
|
||||
func envMapToSlice(m map[string]string) []string {
|
||||
s := make([]string, 0, len(m))
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"go/token"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -353,6 +354,59 @@ func TestApplyServiceEnvParams_NotChanged(t *testing.T) {
|
||||
assert.Equal(t, map[string]string{"FROM_SAVED": "val"}, result)
|
||||
}
|
||||
|
||||
func TestParseServiceEnvVars_RejectsForbiddenNames(t *testing.T) {
|
||||
for _, env := range []string{"PATH=C:\\somewhere", "LD_PRELOAD=/tmp/lib.so", "DYLD_FALLBACK_LIBRARY_PATH=/tmp"} {
|
||||
_, err := parseServiceEnvVars([]string{"KEEP=me", env})
|
||||
require.Errorf(t, err, "%s selects what the service resolves and must be refused", env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsForbiddenServiceEnvVar(t *testing.T) {
|
||||
// The loader families are matched by prefix, so a name nobody has heard of
|
||||
// yet is refused too.
|
||||
for _, name := range []string{
|
||||
"PATH", "PATHEXT", "COMSPEC", "SYSTEMROOT", "WINDIR", "TEMP", "TMP",
|
||||
"LD_PRELOAD", "LD_AUDIT", "DYLD_INSERT_LIBRARIES", "DYLD_FALLBACK_FRAMEWORK_PATH",
|
||||
} {
|
||||
assert.Truef(t, isForbiddenServiceEnvVar(name), "%s must be refused", name)
|
||||
}
|
||||
|
||||
// The prefix must not swallow names that merely start with the same letters.
|
||||
for _, name := range []string{"NB_LOG_LEVEL", "NB_WG_DEBUG", "HTTPS_PROXY", "LDAP_URL", "DYLDX"} {
|
||||
assert.Falsef(t, isForbiddenServiceEnvVar(name), "%s has no reason to be refused", name)
|
||||
}
|
||||
|
||||
// On Windows a variable is the same one however it is spelled; elsewhere
|
||||
// Path and PATH are two variables and only the exact one is read.
|
||||
if runtime.GOOS == "windows" {
|
||||
assert.True(t, isForbiddenServiceEnvVar("Path"))
|
||||
assert.True(t, isForbiddenServiceEnvVar("ld_preload"))
|
||||
} else {
|
||||
assert.False(t, isForbiddenServiceEnvVar("Path"))
|
||||
assert.False(t, isForbiddenServiceEnvVar("ld_preload"))
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyServiceEnvParams_DropsForbiddenSavedNames(t *testing.T) {
|
||||
origServiceEnvVars := serviceEnvVars
|
||||
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
|
||||
|
||||
serviceEnvVars = nil
|
||||
|
||||
cmd := &cobra.Command{}
|
||||
cmd.Flags().StringSlice("service-env", nil, "")
|
||||
|
||||
saved := &serviceParams{
|
||||
ServiceEnvVars: map[string]string{"PATH": "C:\\attacker", "NB_LOG_FORMAT": "json"},
|
||||
}
|
||||
|
||||
applyServiceEnvParams(cmd, saved)
|
||||
|
||||
result, err := parseServiceEnvVars(serviceEnvVars)
|
||||
require.NoError(t, err, "a saved PATH must be dropped rather than fail the install")
|
||||
assert.Equal(t, map[string]string{"NB_LOG_FORMAT": "json"}, result)
|
||||
}
|
||||
|
||||
func TestApplyServiceEnvParams_ExplicitEmptyClears(t *testing.T) {
|
||||
origServiceEnvVars := serviceEnvVars
|
||||
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
//go:build !windows && !ios && !android
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/client/configs"
|
||||
)
|
||||
|
||||
// The Windows equivalent of this is the ACL check in
|
||||
// elevate.CheckOnlyOwnerWritable, covered by that package's own tests; here the
|
||||
// point is that loadServiceParams asks the question at all.
|
||||
func TestLoadServiceParams_RefusesWorldWritableFile(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
original := configs.StateDir
|
||||
t.Cleanup(func() { configs.StateDir = original })
|
||||
configs.StateDir = tmpDir
|
||||
|
||||
path := filepath.Join(tmpDir, serviceParamsFile)
|
||||
require.NoError(t, os.WriteFile(path, []byte(`{"log_level":"debug"}`), 0o666))
|
||||
// WriteFile is subject to the umask, so set the bits that matter explicitly.
|
||||
require.NoError(t, os.Chmod(path, 0o666))
|
||||
|
||||
params, err := loadServiceParams()
|
||||
require.Error(t, err, "a service.json anyone can rewrite must not be trusted")
|
||||
assert.Nil(t, params)
|
||||
|
||||
require.NoError(t, os.Chmod(path, 0o600))
|
||||
params, err = loadServiceParams()
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, params)
|
||||
assert.Equal(t, "debug", params.LogLevel)
|
||||
}
|
||||
|
||||
func TestLoadServiceParams_RefusesWorldWritableDirectory(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
stateDir := filepath.Join(tmpDir, "state")
|
||||
require.NoError(t, os.Mkdir(stateDir, 0o777))
|
||||
require.NoError(t, os.Chmod(stateDir, 0o777))
|
||||
|
||||
original := configs.StateDir
|
||||
t.Cleanup(func() { configs.StateDir = original })
|
||||
configs.StateDir = stateDir
|
||||
|
||||
require.NoError(t, os.WriteFile(filepath.Join(stateDir, serviceParamsFile), []byte(`{}`), 0o600))
|
||||
|
||||
params, err := loadServiceParams()
|
||||
require.Error(t, err, "a service.json in a directory anyone can replace entries in must not be trusted")
|
||||
assert.Nil(t, params)
|
||||
}
|
||||
@@ -24,6 +24,7 @@ const (
|
||||
tableFilter = "filter"
|
||||
tableNat = "nat"
|
||||
tableMangle = "mangle"
|
||||
tableRaw = "raw"
|
||||
|
||||
// chainACLInput is the peer ACL chain that holds installed
|
||||
// peer-filtering rules.
|
||||
@@ -34,6 +35,7 @@ const (
|
||||
mangleForwardKey chainKey = "MANGLE-FORWARD"
|
||||
|
||||
chainInput = "INPUT"
|
||||
chainOutput = "OUTPUT"
|
||||
chainPostrouting = "POSTROUTING"
|
||||
chainPrerouting = "PREROUTING"
|
||||
chainForward = "FORWARD"
|
||||
|
||||
@@ -25,9 +25,8 @@ type Manager struct {
|
||||
|
||||
wgIface iFaceMapper
|
||||
|
||||
ipv4Client *iptables.IPTables
|
||||
family4 *family
|
||||
rawSupported bool
|
||||
ipv4Client *iptables.IPTables
|
||||
family4 *family
|
||||
|
||||
// IPv6 counterparts, nil when no v6 overlay
|
||||
ipv6Client *iptables.IPTables
|
||||
@@ -108,10 +107,6 @@ func (m *Manager) Init(stateManager *statemanager.Manager) error {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := m.initNoTrackChain(); err != nil {
|
||||
log.Warnf("raw table not available, notrack rules will be disabled: %v", err)
|
||||
}
|
||||
|
||||
// Trust after all fatal init steps so a later failure doesn't leave the
|
||||
// interface in firewalld's trusted zone without a corresponding Close.
|
||||
if err := firewalld.TrustInterface(m.wgIface.Name()); err != nil {
|
||||
@@ -285,10 +280,6 @@ func (m *Manager) Close(stateManager *statemanager.Manager) error {
|
||||
|
||||
var merr *multierror.Error
|
||||
|
||||
if err := m.cleanupNoTrackChain(); err != nil {
|
||||
merr = multierror.Append(merr, fmt.Errorf("cleanup notrack chain: %w", err))
|
||||
}
|
||||
|
||||
if m.hasIPv6() {
|
||||
if err := m.family6.Reset(); err != nil {
|
||||
merr = multierror.Append(merr, fmt.Errorf("reset v6 family: %w", err))
|
||||
@@ -440,134 +431,6 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto
|
||||
return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||
}
|
||||
|
||||
const (
|
||||
chainNameRaw = "NETBIRD-RAW"
|
||||
chainOutput = "OUTPUT"
|
||||
tableRaw = "raw"
|
||||
)
|
||||
|
||||
// SetupEBPFProxyNoTrack creates notrack rules for eBPF proxy loopback traffic.
|
||||
// This prevents conntrack from tracking WireGuard proxy traffic on loopback, which
|
||||
// can interfere with MASQUERADE rules (e.g., from container runtimes like Podman/netavark).
|
||||
//
|
||||
// Traffic flows that need NOTRACK:
|
||||
//
|
||||
// 1. Egress: WireGuard -> fake endpoint (before eBPF rewrite)
|
||||
// src=127.0.0.1:wgPort -> dst=127.0.0.1:fakePort
|
||||
// Matched by: sport=wgPort
|
||||
//
|
||||
// 2. Egress: Proxy -> WireGuard (via raw socket)
|
||||
// src=127.0.0.1:fakePort -> dst=127.0.0.1:wgPort
|
||||
// Matched by: dport=wgPort
|
||||
//
|
||||
// 3. Ingress: Packets to WireGuard
|
||||
// dst=127.0.0.1:wgPort
|
||||
// Matched by: dport=wgPort
|
||||
//
|
||||
// 4. Ingress: Packets to proxy (after eBPF rewrite)
|
||||
// dst=127.0.0.1:proxyPort
|
||||
// Matched by: dport=proxyPort
|
||||
//
|
||||
// Rules are cleaned up when the firewall manager is closed.
|
||||
func (m *Manager) SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error {
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
|
||||
if !m.rawSupported {
|
||||
return fmt.Errorf("raw table not available")
|
||||
}
|
||||
|
||||
wgPortStr := fmt.Sprintf("%d", wgPort)
|
||||
proxyPortStr := fmt.Sprintf("%d", proxyPort)
|
||||
|
||||
// Egress rules: match outgoing loopback UDP packets
|
||||
outputRuleSport := []string{"-o", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--sport", wgPortStr, "-j", "NOTRACK"}
|
||||
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, outputRuleSport...); err != nil {
|
||||
return fmt.Errorf("add output sport notrack rule: %w", err)
|
||||
}
|
||||
|
||||
outputRuleDport := []string{"-o", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", wgPortStr, "-j", "NOTRACK"}
|
||||
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, outputRuleDport...); err != nil {
|
||||
return fmt.Errorf("add output dport notrack rule: %w", err)
|
||||
}
|
||||
|
||||
// Ingress rules: match incoming loopback UDP packets
|
||||
preroutingRuleWg := []string{"-i", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", wgPortStr, "-j", "NOTRACK"}
|
||||
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, preroutingRuleWg...); err != nil {
|
||||
return fmt.Errorf("add prerouting wg notrack rule: %w", err)
|
||||
}
|
||||
|
||||
preroutingRuleProxy := []string{"-i", "lo", "-s", "127.0.0.1", "-d", "127.0.0.1", "-p", "udp", "--dport", proxyPortStr, "-j", "NOTRACK"}
|
||||
if err := m.ipv4Client.AppendUnique(tableRaw, chainNameRaw, preroutingRuleProxy...); err != nil {
|
||||
return fmt.Errorf("add prerouting proxy notrack rule: %w", err)
|
||||
}
|
||||
|
||||
log.Debugf("set up ebpf proxy notrack rules for ports %d,%d", proxyPort, wgPort)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) initNoTrackChain() error {
|
||||
if err := m.cleanupNoTrackChain(); err != nil {
|
||||
log.Debugf("cleanup notrack chain: %v", err)
|
||||
}
|
||||
|
||||
if err := m.ipv4Client.NewChain(tableRaw, chainNameRaw); err != nil {
|
||||
return fmt.Errorf("create chain: %w", err)
|
||||
}
|
||||
|
||||
jumpRule := []string{"-j", chainNameRaw}
|
||||
|
||||
if err := m.ipv4Client.InsertUnique(tableRaw, chainOutput, 1, jumpRule...); err != nil {
|
||||
if delErr := m.ipv4Client.DeleteChain(tableRaw, chainNameRaw); delErr != nil {
|
||||
log.Debugf("delete orphan chain: %v", delErr)
|
||||
}
|
||||
return fmt.Errorf("add output jump rule: %w", err)
|
||||
}
|
||||
|
||||
if err := m.ipv4Client.InsertUnique(tableRaw, chainPrerouting, 1, jumpRule...); err != nil {
|
||||
if delErr := m.ipv4Client.DeleteIfExists(tableRaw, chainOutput, jumpRule...); delErr != nil {
|
||||
log.Debugf("delete output jump rule: %v", delErr)
|
||||
}
|
||||
if delErr := m.ipv4Client.DeleteChain(tableRaw, chainNameRaw); delErr != nil {
|
||||
log.Debugf("delete orphan chain: %v", delErr)
|
||||
}
|
||||
return fmt.Errorf("add prerouting jump rule: %w", err)
|
||||
}
|
||||
|
||||
m.rawSupported = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) cleanupNoTrackChain() error {
|
||||
exists, err := m.ipv4Client.ChainExists(tableRaw, chainNameRaw)
|
||||
if err != nil {
|
||||
if !m.rawSupported {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("check chain exists: %w", err)
|
||||
}
|
||||
if !exists {
|
||||
return nil
|
||||
}
|
||||
|
||||
jumpRule := []string{"-j", chainNameRaw}
|
||||
|
||||
if err := m.ipv4Client.DeleteIfExists(tableRaw, chainOutput, jumpRule...); err != nil {
|
||||
return fmt.Errorf("remove output jump rule: %w", err)
|
||||
}
|
||||
|
||||
if err := m.ipv4Client.DeleteIfExists(tableRaw, chainPrerouting, jumpRule...); err != nil {
|
||||
return fmt.Errorf("remove prerouting jump rule: %w", err)
|
||||
}
|
||||
|
||||
if err := m.ipv4Client.ClearAndDeleteChain(tableRaw, chainNameRaw); err != nil {
|
||||
return fmt.Errorf("clear and delete chain: %w", err)
|
||||
}
|
||||
|
||||
m.rawSupported = false
|
||||
return nil
|
||||
}
|
||||
|
||||
func getConntrackEstablished() []string {
|
||||
return []string{"-m", "conntrack", "--ctstate", "RELATED,ESTABLISHED", "-j", "ACCEPT"}
|
||||
}
|
||||
|
||||
@@ -192,10 +192,6 @@ type Manager interface {
|
||||
|
||||
// RemoveOutputDNAT removes an OUTPUT chain DNAT rule.
|
||||
RemoveOutputDNAT(localAddr netip.Addr, protocol Protocol, originalPort, translatedPort uint16) error
|
||||
|
||||
// SetupEBPFProxyNoTrack creates static notrack rules for eBPF proxy loopback traffic.
|
||||
// This prevents conntrack from interfering with WireGuard proxy communication.
|
||||
SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error
|
||||
}
|
||||
|
||||
// GenKey builds the rule id for this pair from the given format.
|
||||
|
||||
@@ -12,7 +12,6 @@ import (
|
||||
"github.com/google/nftables/expr"
|
||||
"github.com/hashicorp/go-multierror"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||
firewall "github.com/netbirdio/netbird/client/firewall/manager"
|
||||
@@ -55,9 +54,6 @@ type Manager struct {
|
||||
// IPv6 counterpart, nil when no v6 overlay.
|
||||
family6 *family
|
||||
|
||||
notrackOutputChain *nftables.Chain
|
||||
notrackPreroutingChain *nftables.Chain
|
||||
|
||||
extMonitor *externalChainMonitor
|
||||
}
|
||||
|
||||
@@ -170,10 +166,6 @@ func (m *Manager) initFirewall() (err error) {
|
||||
}
|
||||
}
|
||||
|
||||
if err := m.initNoTrackChains(workTable); err != nil {
|
||||
log.Warnf("raw priority chains not available, notrack rules will be disabled: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -455,10 +447,6 @@ func (m *Manager) Flush() error {
|
||||
}
|
||||
}
|
||||
|
||||
if err := m.refreshNoTrackChains(); err != nil {
|
||||
log.Errorf("failed to refresh notrack chains: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -571,176 +559,6 @@ func (m *Manager) RemoveOutputDNAT(localAddr netip.Addr, protocol firewall.Proto
|
||||
return m.family4.RemoveOutputDNAT(localAddr, protocol, originalPort, translatedPort)
|
||||
}
|
||||
|
||||
const (
|
||||
chainNameRawOutput = "netbird-raw-out"
|
||||
chainNameRawPrerouting = "netbird-raw-pre"
|
||||
)
|
||||
|
||||
// SetupEBPFProxyNoTrack creates notrack rules for eBPF proxy loopback traffic.
|
||||
// This prevents conntrack from tracking WireGuard proxy traffic on loopback, which
|
||||
// can interfere with MASQUERADE rules (e.g., from container runtimes like Podman/netavark).
|
||||
//
|
||||
// Traffic flows that need NOTRACK:
|
||||
//
|
||||
// 1. Egress: WireGuard -> fake endpoint (before eBPF rewrite)
|
||||
// src=127.0.0.1:wgPort -> dst=127.0.0.1:fakePort
|
||||
// Matched by: sport=wgPort
|
||||
//
|
||||
// 2. Egress: Proxy -> WireGuard (via raw socket)
|
||||
// src=127.0.0.1:fakePort -> dst=127.0.0.1:wgPort
|
||||
// Matched by: dport=wgPort
|
||||
//
|
||||
// 3. Ingress: Packets to WireGuard
|
||||
// dst=127.0.0.1:wgPort
|
||||
// Matched by: dport=wgPort
|
||||
//
|
||||
// 4. Ingress: Packets to proxy (after eBPF rewrite)
|
||||
// dst=127.0.0.1:proxyPort
|
||||
// Matched by: dport=proxyPort
|
||||
//
|
||||
// Rules are cleaned up when the firewall manager is closed.
|
||||
func (m *Manager) SetupEBPFProxyNoTrack(proxyPort, wgPort uint16) error {
|
||||
m.mutex.Lock()
|
||||
defer m.mutex.Unlock()
|
||||
|
||||
if m.notrackOutputChain == nil || m.notrackPreroutingChain == nil {
|
||||
return fmt.Errorf("notrack chains not initialized")
|
||||
}
|
||||
|
||||
proxyPortBytes := binaryutil.BigEndian.PutUint16(proxyPort)
|
||||
wgPortBytes := binaryutil.BigEndian.PutUint16(wgPort)
|
||||
loopback := []byte{127, 0, 0, 1}
|
||||
|
||||
// Egress rules: match outgoing loopback UDP packets
|
||||
m.rConn.AddRule(&nftables.Rule{
|
||||
Table: m.notrackOutputChain.Table,
|
||||
Chain: m.notrackOutputChain,
|
||||
Exprs: []expr.Any{
|
||||
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 2},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // sport=wgPort
|
||||
&expr.Counter{},
|
||||
&expr.Notrack{},
|
||||
},
|
||||
})
|
||||
m.rConn.AddRule(&nftables.Rule{
|
||||
Table: m.notrackOutputChain.Table,
|
||||
Chain: m.notrackOutputChain,
|
||||
Exprs: []expr.Any{
|
||||
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // dport=wgPort
|
||||
&expr.Counter{},
|
||||
&expr.Notrack{},
|
||||
},
|
||||
})
|
||||
|
||||
// Ingress rules: match incoming loopback UDP packets
|
||||
m.rConn.AddRule(&nftables.Rule{
|
||||
Table: m.notrackPreroutingChain.Table,
|
||||
Chain: m.notrackPreroutingChain,
|
||||
Exprs: []expr.Any{
|
||||
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: wgPortBytes}, // dport=wgPort
|
||||
&expr.Counter{},
|
||||
&expr.Notrack{},
|
||||
},
|
||||
})
|
||||
m.rConn.AddRule(&nftables.Rule{
|
||||
Table: m.notrackPreroutingChain.Table,
|
||||
Chain: m.notrackPreroutingChain,
|
||||
Exprs: []expr.Any{
|
||||
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname("lo")},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4}, // saddr
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 16, Len: 4}, // daddr
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: loopback},
|
||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_UDP}},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: proxyPortBytes}, // dport=proxyPort
|
||||
&expr.Counter{},
|
||||
&expr.Notrack{},
|
||||
},
|
||||
})
|
||||
|
||||
if err := m.rConn.Flush(); err != nil {
|
||||
return fmt.Errorf("flush notrack rules: %w", err)
|
||||
}
|
||||
|
||||
log.Debugf("set up ebpf proxy notrack rules for ports %d,%d", proxyPort, wgPort)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) initNoTrackChains(table *nftables.Table) error {
|
||||
m.notrackOutputChain = m.rConn.AddChain(&nftables.Chain{
|
||||
Name: chainNameRawOutput,
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeFilter,
|
||||
Hooknum: nftables.ChainHookOutput,
|
||||
Priority: nftables.ChainPriorityRaw,
|
||||
})
|
||||
|
||||
m.notrackPreroutingChain = m.rConn.AddChain(&nftables.Chain{
|
||||
Name: chainNameRawPrerouting,
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeFilter,
|
||||
Hooknum: nftables.ChainHookPrerouting,
|
||||
Priority: nftables.ChainPriorityRaw,
|
||||
})
|
||||
|
||||
if err := m.rConn.Flush(); err != nil {
|
||||
return fmt.Errorf("flush chain creation: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) refreshNoTrackChains() error {
|
||||
chains, err := m.rConn.ListChainsOfTableFamily(nftables.TableFamilyIPv4)
|
||||
if err != nil {
|
||||
return fmt.Errorf("list chains: %w", err)
|
||||
}
|
||||
|
||||
tableName := getTableName()
|
||||
for _, c := range chains {
|
||||
if c.Table.Name != tableName {
|
||||
continue
|
||||
}
|
||||
switch c.Name {
|
||||
case chainNameRawOutput:
|
||||
m.notrackOutputChain = c
|
||||
case chainNameRawPrerouting:
|
||||
m.notrackPreroutingChain = c
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) createWorkTable() (*nftables.Table, error) {
|
||||
return m.createWorkTableFamily(nftables.TableFamilyIPv4)
|
||||
}
|
||||
|
||||
@@ -192,7 +192,7 @@ func (r *family) addPostroutingRules() {
|
||||
Data: binaryutil.NativeEndian.PutUint32(nbnet.PreroutingFwmarkMasquerade),
|
||||
},
|
||||
|
||||
// We need to exclude the loopback interface as this changes the ebpf proxy port
|
||||
// We need to exclude the loopback interface as this changes the wg proxy port
|
||||
&expr.Meta{
|
||||
Key: expr.MetaKeyOIFNAME,
|
||||
Register: 1,
|
||||
|
||||
@@ -879,12 +879,6 @@ func (m *Manager) resetState() {
|
||||
}
|
||||
}
|
||||
|
||||
// SetupEBPFProxyNoTrack is not supported by the userspace firewall: eBPF isn't
|
||||
// used in userspace mode, so this should never be called.
|
||||
func (m *Manager) SetupEBPFProxyNoTrack(uint16, uint16) error {
|
||||
return errNotSupported
|
||||
}
|
||||
|
||||
// UpdateSet updates the rule destinations associated with the given set
|
||||
// by merging the existing prefixes with the new ones, then deduplicating.
|
||||
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||
"github.com/netbirdio/netbird/client/internal/wincmd"
|
||||
)
|
||||
|
||||
type action string
|
||||
@@ -91,7 +92,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
|
||||
if action == addRule {
|
||||
args = append(args, extraArgs...)
|
||||
}
|
||||
netshCmd := GetSystem32Command("netsh")
|
||||
netshCmd := wincmd.System32("netsh")
|
||||
cmd := exec.Command(netshCmd, args...)
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
||||
return cmd.Run()
|
||||
@@ -100,7 +101,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
|
||||
func isWindowsFirewallReachable() bool {
|
||||
args := []string{"advfirewall", "show", "allprofiles", "state"}
|
||||
|
||||
netshCmd := GetSystem32Command("netsh")
|
||||
netshCmd := wincmd.System32("netsh")
|
||||
|
||||
cmd := exec.Command(netshCmd, args...)
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
||||
@@ -117,23 +118,10 @@ func isWindowsFirewallReachable() bool {
|
||||
func isFirewallRuleActive(ruleName string) bool {
|
||||
args := []string{"advfirewall", "firewall", "show", "rule", "name=" + ruleName}
|
||||
|
||||
netshCmd := GetSystem32Command("netsh")
|
||||
netshCmd := wincmd.System32("netsh")
|
||||
|
||||
cmd := exec.Command(netshCmd, args...)
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
||||
_, err := cmd.Output()
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
|
||||
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
|
||||
func GetSystem32Command(command string) string {
|
||||
_, err := exec.LookPath(command)
|
||||
if err == nil {
|
||||
return command
|
||||
}
|
||||
|
||||
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
|
||||
|
||||
return "C:\\windows\\system32\\" + command + ".exe"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
package configurer
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"sync"
|
||||
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
)
|
||||
|
||||
// allowedIPStore mirrors the allowed IPs configured on each peer of a device.
|
||||
//
|
||||
// A configurer is the only writer of its device's peer set, so the mirror is authoritative
|
||||
// by construction. It spares the paths that have to rewrite one peer's allowed IPs a full
|
||||
// device dump just to recover prefixes the process already configured itself.
|
||||
//
|
||||
// An allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away
|
||||
// from whichever peer held it before, and the configurer leaves that handover to the device
|
||||
// rather than removing the prefix from the previous holder itself. The store tracks the
|
||||
// owner of each prefix and performs the same handover, so rewriting one peer's list never
|
||||
// takes a prefix back from the peer that owns it now.
|
||||
//
|
||||
// Its own lock guards the map alone, not the device write it accompanies. Consistency
|
||||
// between the two rests on the caller serializing every configurer call, which WGIface
|
||||
// does with its mutex; two unserialized writers would interleave a device write with the
|
||||
// record of a different one.
|
||||
//
|
||||
// An operator reconfiguring the device out of band, through `wg set` or the UAPI socket,
|
||||
// is the one way the mirror can still go stale. A peer missing from it falls back to the
|
||||
// device, which reseats that peer's prefixes and their ownership; a peer that is present
|
||||
// does not, so one recorded from empty while the device already held prefixes keeps only
|
||||
// what was recorded, and the next endpoint removal drops the rest.
|
||||
type allowedIPStore struct {
|
||||
mu sync.RWMutex
|
||||
peers map[wgtypes.Key][]netip.Prefix
|
||||
owners map[netip.Prefix]wgtypes.Key
|
||||
}
|
||||
|
||||
func newAllowedIPStore() *allowedIPStore {
|
||||
return &allowedIPStore{
|
||||
peers: make(map[wgtypes.Key][]netip.Prefix),
|
||||
owners: make(map[netip.Prefix]wgtypes.Key),
|
||||
}
|
||||
}
|
||||
|
||||
// get returns the prefixes recorded for a peer, and whether the peer is known at all.
|
||||
// The caller receives a copy and may retain or modify it freely.
|
||||
func (s *allowedIPStore) get(key wgtypes.Key) ([]netip.Prefix, bool) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
prefixes, ok := s.peers[key]
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
return slices.Clone(prefixes), true
|
||||
}
|
||||
|
||||
// set replaces the prefixes recorded for a peer.
|
||||
func (s *allowedIPStore) set(key wgtypes.Key, prefixes []netip.Prefix) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
k := key
|
||||
s.releaseLocked(k)
|
||||
|
||||
normalized := normalizePrefixes(prefixes)
|
||||
for _, prefix := range normalized {
|
||||
s.claimLocked(k, prefix)
|
||||
}
|
||||
s.peers[k] = normalized
|
||||
}
|
||||
|
||||
// add records prefixes on a peer without dropping the ones already there, matching the
|
||||
// union semantics of a peer update that does not replace its allowed IPs. It records the
|
||||
// peer if it is not known yet, so it belongs to the operations that create a peer on the
|
||||
// device rather than to the update-only ones.
|
||||
func (s *allowedIPStore) add(key wgtypes.Key, prefixes []netip.Prefix) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.mergeLocked(key, prefixes)
|
||||
}
|
||||
|
||||
// addExisting is add for an update-only device operation. Such an operation is a silent
|
||||
// no-op when the peer is absent, so recording a peer here would leave the store claiming
|
||||
// prefixes the device never took, and the peer would then be recreated by the next endpoint
|
||||
// removal, stealing those allowed IPs from the peer that legitimately holds them.
|
||||
func (s *allowedIPStore) addExisting(key wgtypes.Key, prefixes []netip.Prefix) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
k := key
|
||||
if _, ok := s.peers[k]; !ok {
|
||||
return
|
||||
}
|
||||
s.mergeLocked(k, prefixes)
|
||||
}
|
||||
|
||||
// ensure records a peer with no prefixes unless it is already known. A device operation
|
||||
// that is not update-only creates the peer when it is absent, so it has to be recorded even
|
||||
// when it configures nothing else; otherwise the peer exists on the device while the store
|
||||
// treats it as unknown, and a prefix later handed over to it is not accounted for.
|
||||
func (s *allowedIPStore) ensure(key wgtypes.Key) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
k := key
|
||||
if _, ok := s.peers[k]; !ok {
|
||||
s.peers[k] = nil
|
||||
}
|
||||
}
|
||||
|
||||
// forget drops every prefix recorded for a peer.
|
||||
func (s *allowedIPStore) forget(key wgtypes.Key) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
k := key
|
||||
s.releaseLocked(k)
|
||||
delete(s.peers, k)
|
||||
}
|
||||
|
||||
// reset drops every peer, mirroring a device reconfiguration that replaces the peer set.
|
||||
func (s *allowedIPStore) reset() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.peers = make(map[wgtypes.Key][]netip.Prefix)
|
||||
s.owners = make(map[netip.Prefix]wgtypes.Key)
|
||||
}
|
||||
|
||||
// mergeLocked unions normalized prefixes into a peer and transfers their ownership.
|
||||
// The caller must hold s.mu for writing.
|
||||
func (s *allowedIPStore) mergeLocked(k wgtypes.Key, prefixes []netip.Prefix) {
|
||||
merged := s.peers[k]
|
||||
for _, prefix := range prefixes {
|
||||
prefix = normalizePrefix(prefix)
|
||||
s.claimLocked(k, prefix)
|
||||
if !slices.Contains(merged, prefix) {
|
||||
merged = append(merged, prefix)
|
||||
}
|
||||
}
|
||||
s.peers[k] = merged
|
||||
}
|
||||
|
||||
// claimLocked hands a prefix over to a peer, taking it from its previous owner the way the
|
||||
// device does when the same prefix is configured on a second peer.
|
||||
func (s *allowedIPStore) claimLocked(k wgtypes.Key, prefix netip.Prefix) {
|
||||
if owner, ok := s.owners[prefix]; ok && owner != k {
|
||||
s.peers[owner] = slices.DeleteFunc(s.peers[owner], func(p netip.Prefix) bool {
|
||||
return p == prefix
|
||||
})
|
||||
}
|
||||
s.owners[prefix] = k
|
||||
}
|
||||
|
||||
// releaseLocked drops a peer's claim on every prefix it currently holds.
|
||||
func (s *allowedIPStore) releaseLocked(k wgtypes.Key) {
|
||||
for _, prefix := range s.peers[k] {
|
||||
if s.owners[prefix] == k {
|
||||
delete(s.owners, prefix)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// normalizePrefix puts a prefix into the form the store recognises it by. It clears the
|
||||
// host bits, which a device does on its own, so a caller passing 10.20.0.1/16 still matches
|
||||
// the 10.20.0.0/16 read back from the device; and it unmaps a v4-mapped prefix so that it
|
||||
// compares equal to, and marshals like, the plain v4 prefix for the same network.
|
||||
//
|
||||
// Masking comes first because it also decides the address family: only a prefix at least 96
|
||||
// bits long keeps the mapped marker through the mask, so a shorter prefix inside the mapped
|
||||
// range is a genuine v6 prefix and unmapping it would yield an invalid v4 prefix.
|
||||
func normalizePrefix(prefix netip.Prefix) netip.Prefix {
|
||||
masked := prefix.Masked()
|
||||
|
||||
addr := masked.Addr()
|
||||
if !addr.Is4In6() {
|
||||
return masked
|
||||
}
|
||||
return netip.PrefixFrom(addr.Unmap(), masked.Bits()-96)
|
||||
}
|
||||
|
||||
// normalizePrefixes returns a normalized copy without changing the caller's slice.
|
||||
func normalizePrefixes(prefixes []netip.Prefix) []netip.Prefix {
|
||||
normalized := make([]netip.Prefix, len(prefixes))
|
||||
for i, prefix := range prefixes {
|
||||
normalized[i] = normalizePrefix(prefix)
|
||||
}
|
||||
return normalized
|
||||
}
|
||||
|
||||
// ipNetsToPrefixes converts addresses read back from a device. Unmap keeps a v4-mapped v6
|
||||
// address comparable to the plain v4 prefix the configurer was given.
|
||||
func ipNetsToPrefixes(ipNets []net.IPNet) []netip.Prefix {
|
||||
prefixes := make([]netip.Prefix, 0, len(ipNets))
|
||||
for _, ipNet := range ipNets {
|
||||
addr, ok := netip.AddrFromSlice(ipNet.IP)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
ones, maskBits := ipNet.Mask.Size()
|
||||
// A device may report a v4 prefix as a v4-mapped address. Align the address form with
|
||||
// the mask rather than unmapping on sight: a 32 bit mask always describes v4, while a
|
||||
// 128 bit mask describes v4 only when it covers the mapped prefix, so a genuine v6
|
||||
// prefix inside the mapped range stays v6 instead of being dropped as invalid.
|
||||
if addr.Is4In6() {
|
||||
switch {
|
||||
case maskBits == 32:
|
||||
addr = addr.Unmap()
|
||||
case maskBits == 128 && ones >= 96:
|
||||
addr, ones = addr.Unmap(), ones-96
|
||||
}
|
||||
}
|
||||
|
||||
prefix := netip.PrefixFrom(addr, ones)
|
||||
if !prefix.IsValid() {
|
||||
continue
|
||||
}
|
||||
prefixes = append(prefixes, prefix.Masked())
|
||||
}
|
||||
return prefixes
|
||||
}
|
||||
@@ -0,0 +1,263 @@
|
||||
package configurer
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
)
|
||||
|
||||
// The store keys on the parsed key, so the tests use two distinct ones rather than names.
|
||||
var (
|
||||
testPeer = wgtypes.Key{1}
|
||||
otherPeer = wgtypes.Key{2}
|
||||
)
|
||||
|
||||
func TestAllowedIPStoreUnknownPeer(t *testing.T) {
|
||||
s := newAllowedIPStore()
|
||||
|
||||
prefixes, ok := s.get(testPeer)
|
||||
assert.False(t, ok, "an unconfigured peer must be reported as unknown, not as one without prefixes")
|
||||
assert.Nil(t, prefixes, "an unknown peer has no prefixes")
|
||||
}
|
||||
|
||||
func TestAllowedIPStoreAddUnions(t *testing.T) {
|
||||
s := newAllowedIPStore()
|
||||
overlay := netip.MustParsePrefix("100.64.0.1/32")
|
||||
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||
|
||||
s.set(testPeer, []netip.Prefix{overlay})
|
||||
// A peer update does not replace allowed IPs, and a repeated prefix must not be doubled.
|
||||
s.add(testPeer, []netip.Prefix{overlay, routed})
|
||||
|
||||
prefixes, ok := s.get(testPeer)
|
||||
require.True(t, ok, "peer must be known after set")
|
||||
assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "add must union rather than replace")
|
||||
}
|
||||
|
||||
func TestAllowedIPStoreGetReturnsCopy(t *testing.T) {
|
||||
s := newAllowedIPStore()
|
||||
overlay := netip.MustParsePrefix("100.64.0.1/32")
|
||||
s.set(testPeer, []netip.Prefix{overlay})
|
||||
|
||||
prefixes, ok := s.get(testPeer)
|
||||
require.True(t, ok, "peer must be known after set")
|
||||
prefixes[0] = netip.MustParsePrefix("0.0.0.0/0")
|
||||
|
||||
stored, _ := s.get(testPeer)
|
||||
assert.Equal(t, []netip.Prefix{overlay}, stored, "a caller mutating the returned slice must not corrupt the store")
|
||||
}
|
||||
|
||||
func TestAllowedIPStoreForgetAndReset(t *testing.T) {
|
||||
s := newAllowedIPStore()
|
||||
s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")})
|
||||
s.set(otherPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")})
|
||||
|
||||
s.forget(testPeer)
|
||||
_, ok := s.get(testPeer)
|
||||
assert.False(t, ok, "a forgotten peer must be unknown")
|
||||
_, ok = s.get(otherPeer)
|
||||
assert.True(t, ok, "forgetting one peer must not touch the others")
|
||||
|
||||
s.reset()
|
||||
_, ok = s.get(otherPeer)
|
||||
assert.False(t, ok, "reset must drop every peer")
|
||||
}
|
||||
|
||||
func TestIPNetsToPrefixes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ipNet net.IPNet
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "v4",
|
||||
ipNet: net.IPNet{IP: net.IP{10, 20, 0, 0}, Mask: net.CIDRMask(16, 32)},
|
||||
want: "10.20.0.0/16",
|
||||
},
|
||||
{
|
||||
name: "v4 mapped under a 128 bit mask",
|
||||
ipNet: net.IPNet{IP: net.ParseIP("10.20.0.0"), Mask: net.CIDRMask(112, 128)},
|
||||
want: "10.20.0.0/16",
|
||||
},
|
||||
{
|
||||
name: "v6",
|
||||
ipNet: net.IPNet{IP: net.ParseIP("fd00::"), Mask: net.CIDRMask(64, 128)},
|
||||
want: "fd00::/64",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := ipNetsToPrefixes([]net.IPNet{tc.ipNet})
|
||||
require.Len(t, got, 1, "the address must be converted, not dropped")
|
||||
assert.Equal(t, tc.want, got[0].String(), "converted prefix")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPNetsToPrefixesRoundTrip(t *testing.T) {
|
||||
prefixes := []netip.Prefix{
|
||||
netip.MustParsePrefix("100.64.0.1/32"),
|
||||
netip.MustParsePrefix("10.20.0.0/16"),
|
||||
netip.MustParsePrefix("fd00::/64"),
|
||||
}
|
||||
|
||||
assert.Equal(t, prefixes, ipNetsToPrefixes(prefixesToIPNets(prefixes)),
|
||||
"prefixes handed to a device must come back unchanged")
|
||||
}
|
||||
|
||||
func TestAllowedIPStoreNormalizesMappedPrefixes(t *testing.T) {
|
||||
s := newAllowedIPStore()
|
||||
v4 := netip.MustParsePrefix("10.20.0.0/16")
|
||||
mapped := netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112)
|
||||
|
||||
s.set(testPeer, []netip.Prefix{mapped})
|
||||
// A v4 rule only matches a v4-mapped address once it has been unmapped, so the store must
|
||||
// hold the plain form and recognise the two spellings as the same prefix.
|
||||
s.add(testPeer, []netip.Prefix{v4})
|
||||
|
||||
prefixes, ok := s.get(testPeer)
|
||||
require.True(t, ok, "peer must be known after set")
|
||||
assert.Equal(t, []netip.Prefix{v4}, prefixes, "a mapped prefix must be stored unmapped and not duplicated")
|
||||
}
|
||||
|
||||
func TestNormalizePrefix(t *testing.T) {
|
||||
v4 := netip.MustParsePrefix("10.20.0.0/16")
|
||||
v6 := netip.MustParsePrefix("fd00::/64")
|
||||
|
||||
assert.Equal(t, v4, normalizePrefix(v4), "a plain v4 prefix is unchanged")
|
||||
assert.Equal(t, v6, normalizePrefix(v6), "a real v6 prefix is unchanged")
|
||||
assert.Equal(t, v4, normalizePrefix(netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112)),
|
||||
"a mapped prefix under a 128 bit mask becomes plain v4")
|
||||
// A prefix shorter than /96 inside the mapped range is a genuine v6 prefix. Unmapping it
|
||||
// would pair a v4 address with a v6 sized mask, which is invalid, and the store would then
|
||||
// record a zero prefix that can never recreate the allowed IP.
|
||||
for _, tc := range []string{"::ffff:0:0/64", "::ffff:1.2.3.4/80", "::ffff:1.2.3.4/95"} {
|
||||
got := normalizePrefix(netip.MustParsePrefix(tc))
|
||||
assert.True(t, got.IsValid(), "%s must normalize to a valid prefix", tc)
|
||||
assert.False(t, got.Addr().Is4(), "%s must stay v6", tc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllowedIPStoreAddExistingDoesNotCreate(t *testing.T) {
|
||||
s := newAllowedIPStore()
|
||||
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||
|
||||
// An update-only device operation on an absent peer is a silent no-op, so nothing may be
|
||||
// recorded for a peer the store does not already know.
|
||||
s.addExisting(testPeer, []netip.Prefix{routed})
|
||||
_, ok := s.get(testPeer)
|
||||
assert.False(t, ok, "addExisting must not record an unknown peer")
|
||||
|
||||
overlay := netip.MustParsePrefix("100.64.0.1/32")
|
||||
s.set(testPeer, []netip.Prefix{overlay})
|
||||
s.addExisting(testPeer, []netip.Prefix{routed})
|
||||
|
||||
prefixes, _ := s.get(testPeer)
|
||||
assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "addExisting must union onto a known peer")
|
||||
}
|
||||
|
||||
func TestAllowedIPStoreHandsPrefixOverToTheNewOwner(t *testing.T) {
|
||||
s := newAllowedIPStore()
|
||||
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||
other := otherPeer
|
||||
|
||||
s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32"), routed})
|
||||
s.set(other, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")})
|
||||
|
||||
// The device takes an allowed IP away from its previous holder when it is configured on
|
||||
// another peer, so the store must do the same rather than list it under both.
|
||||
s.addExisting(other, []netip.Prefix{routed})
|
||||
|
||||
previous, _ := s.get(testPeer)
|
||||
assert.NotContains(t, previous, routed, "the previous owner must lose the prefix")
|
||||
current, _ := s.get(other)
|
||||
assert.Contains(t, current, routed, "the new owner must hold the prefix")
|
||||
}
|
||||
|
||||
func TestAllowedIPStoreForgetReleasesOwnership(t *testing.T) {
|
||||
s := newAllowedIPStore()
|
||||
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||
|
||||
s.set(testPeer, []netip.Prefix{routed})
|
||||
s.forget(testPeer)
|
||||
s.set(otherPeer, []netip.Prefix{routed})
|
||||
|
||||
// A forgotten peer must not be resurrected as a key in the peer map by a later claim.
|
||||
_, ok := s.get(testPeer)
|
||||
assert.False(t, ok, "the forgotten peer must stay unknown")
|
||||
current, _ := s.get(otherPeer)
|
||||
assert.Equal(t, []netip.Prefix{routed}, current, "the new owner must hold the prefix")
|
||||
}
|
||||
|
||||
func TestNormalizePrefixClearsHostBits(t *testing.T) {
|
||||
// A device stores a prefix masked, so a caller passing host bits must still match what a
|
||||
// device fallback seeded, otherwise that prefix could never be removed by value.
|
||||
assert.Equal(t, netip.MustParsePrefix("10.20.0.0/16"),
|
||||
normalizePrefix(netip.MustParsePrefix("10.20.0.1/16")), "host bits must be cleared")
|
||||
assert.Equal(t, netip.MustParsePrefix("fd00::/64"),
|
||||
normalizePrefix(netip.MustParsePrefix("fd00::1/64")), "host bits must be cleared for v6")
|
||||
}
|
||||
|
||||
func TestIPNetsToPrefixesKeepsV6InTheMappedRange(t *testing.T) {
|
||||
// ::ffff:0:0/64 reads as v4-mapped but is a genuine v6 prefix: unmapping it would leave a
|
||||
// v4 address under a 64 bit mask, which is invalid, and the allowed IP would be dropped.
|
||||
got := ipNetsToPrefixes([]net.IPNet{{
|
||||
IP: net.ParseIP("::ffff:0:0"),
|
||||
Mask: net.CIDRMask(64, 128),
|
||||
}})
|
||||
|
||||
require.Len(t, got, 1, "the prefix must be converted, not dropped")
|
||||
assert.False(t, got[0].Addr().Is4(), "a v6 prefix in the mapped range must not become v4")
|
||||
assert.Equal(t, 64, got[0].Bits(), "the prefix length must survive the conversion")
|
||||
}
|
||||
|
||||
func TestPrefixesToIPNetsNormalizes(t *testing.T) {
|
||||
// net.IPNet prints a v4-mapped address as v4 but takes the length from its 16 byte
|
||||
// mask, so an unnormalized ::ffff:10.1.2.3/64 reaches a userspace device as 10.1.2.3/0,
|
||||
// an allowed IP that matches every v4 address.
|
||||
tests := []struct {
|
||||
name string
|
||||
given string
|
||||
want string
|
||||
}{
|
||||
{name: "mapped below /96", given: "::ffff:10.1.2.3/64", want: "::/64"},
|
||||
{name: "mapped at /112", given: "::ffff:10.1.2.3/112", want: "10.1.0.0/16"},
|
||||
{name: "host bits are cleared", given: "10.20.0.1/16", want: "10.20.0.0/16"},
|
||||
{name: "v6 is untouched", given: "fd00::1/64", want: "fd00::/64"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := prefixesToIPNets([]netip.Prefix{netip.MustParsePrefix(tc.given)})
|
||||
require.Len(t, got, 1, "the prefix must be converted, not dropped")
|
||||
assert.Equal(t, tc.want, got[0].String(), "what the device is given")
|
||||
assert.NotEqual(t, 0, mustOnes(t, got[0]), "a device must never be given a zero length allowed IP")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func mustOnes(t *testing.T, ipNet net.IPNet) int {
|
||||
t.Helper()
|
||||
|
||||
ones, _ := ipNet.Mask.Size()
|
||||
return ones
|
||||
}
|
||||
|
||||
// TestPrefixesToIPNetsAgreesWithTheStore pins the property the store depends on: what a
|
||||
// device is given and what is recorded for it are the same prefix.
|
||||
func TestPrefixesToIPNetsAgreesWithTheStore(t *testing.T) {
|
||||
for _, given := range []string{"::ffff:10.1.2.3/64", "::ffff:10.1.2.3/112", "10.20.0.1/16", "fd00::1/64"} {
|
||||
prefix := netip.MustParsePrefix(given)
|
||||
|
||||
toDevice := prefixesToIPNets([]netip.Prefix{prefix})
|
||||
recorded := normalizePrefix(prefix)
|
||||
|
||||
assert.Equal(t, recorded.String(), toDevice[0].String(),
|
||||
"%s must reach the device in the form the store records", given)
|
||||
}
|
||||
}
|
||||
@@ -19,12 +19,18 @@ func buildPresharedKeyConfig(peerKey wgtypes.Key, psk wgtypes.Key, updateOnly bo
|
||||
}
|
||||
}
|
||||
|
||||
// prefixesToIPNets converts prefixes on their way to a device. It is the only place that
|
||||
// conversion happens, so it also normalizes: the device is then given the same form the
|
||||
// store records, and a v4-mapped prefix cannot reach net.IPNet, which prints such an
|
||||
// address as v4 while taking the length from its 16 byte mask and so turns
|
||||
// ::ffff:10.1.2.3/64 into 10.1.2.3/0 — an allowed IP matching every v4 address.
|
||||
func prefixesToIPNets(prefixes []netip.Prefix) []net.IPNet {
|
||||
ipNets := make([]net.IPNet, len(prefixes))
|
||||
for i, prefix := range prefixes {
|
||||
normalized := normalizePrefix(prefix)
|
||||
ipNets[i] = net.IPNet{
|
||||
IP: prefix.Addr().AsSlice(), // Convert netip.Addr to net.IP
|
||||
Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), // Create subnet mask
|
||||
IP: normalized.Addr().AsSlice(),
|
||||
Mask: net.CIDRMask(normalized.Bits(), normalized.Addr().BitLen()),
|
||||
}
|
||||
}
|
||||
return ipNets
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
@@ -18,16 +19,22 @@ import (
|
||||
type KernelConfigurer struct {
|
||||
deviceName string
|
||||
statsCache *statsCache
|
||||
allowedIPs *allowedIPStore
|
||||
}
|
||||
|
||||
// NewKernelConfigurer creates a configurer with an empty allowed IP mirror
|
||||
// and a statistics cache for the named kernel device.
|
||||
func NewKernelConfigurer(deviceName string) *KernelConfigurer {
|
||||
c := &KernelConfigurer{
|
||||
deviceName: deviceName,
|
||||
allowedIPs: newAllowedIPStore(),
|
||||
}
|
||||
c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats)
|
||||
return c
|
||||
}
|
||||
|
||||
// ConfigureInterface sets the device key, port and firewall mark, replacing all peers.
|
||||
// The allowed IP mirror is reset only after the device accepts the configuration.
|
||||
func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error {
|
||||
log.Debugf("adding Wireguard private key")
|
||||
key, err := wgtypes.ParseKey(privateKey)
|
||||
@@ -46,6 +53,8 @@ func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error
|
||||
if err != nil {
|
||||
return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port)
|
||||
}
|
||||
|
||||
c.allowedIPs.reset()
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -58,9 +67,20 @@ func (c *KernelConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, upda
|
||||
}
|
||||
|
||||
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
|
||||
return c.configure(cfg)
|
||||
if err := c.configure(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Without updateOnly this creates the peer when it is absent, so the store has to
|
||||
// know about it even though no allowed IP was configured.
|
||||
if !updateOnly {
|
||||
c.allowedIPs.ensure(parsedPeerKey)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set.
|
||||
// Prefixes assigned to this peer are transferred from their previous owners.
|
||||
func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
@@ -83,19 +103,23 @@ func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
|
||||
if err != nil {
|
||||
return fmt.Errorf(`received error "%w" while updating peer on interface %s with settings: allowed ips %s, endpoint %s`, err, c.deviceName, allowedIps, endpoint.String())
|
||||
}
|
||||
|
||||
c.allowedIPs.add(peerKeyParsed, allowedIps)
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured.
|
||||
// Neither the netlink API nor the userspace one can clear an endpoint in place, so the peer
|
||||
// is removed and re-added with the allowed IPs it already had.
|
||||
func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Get the existing peer to preserve its allowed IPs
|
||||
existingPeer, err := c.getPeer(c.deviceName, peerKey)
|
||||
allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get peer: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
removePeerCfg := wgtypes.PeerConfig{
|
||||
@@ -104,26 +128,27 @@ func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
|
||||
}
|
||||
|
||||
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{removePeerCfg}}); err != nil {
|
||||
return fmt.Errorf(`error removing peer %s from interface %s: %w`, peerKey, c.deviceName, err)
|
||||
return fmt.Errorf("remove peer %s from interface %s: %w", peerKey, c.deviceName, err)
|
||||
}
|
||||
|
||||
//Re-add the peer without the endpoint but same AllowedIPs
|
||||
reAddPeerCfg := wgtypes.PeerConfig{
|
||||
PublicKey: peerKeyParsed,
|
||||
AllowedIPs: existingPeer.AllowedIPs,
|
||||
AllowedIPs: prefixesToIPNets(allowedIPs),
|
||||
ReplaceAllowedIPs: true,
|
||||
}
|
||||
|
||||
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil {
|
||||
c.allowedIPs.forget(peerKeyParsed)
|
||||
return fmt.Errorf(
|
||||
`error re-adding peer %s to interface %s with allowed IPs %v: %w`,
|
||||
peerKey, c.deviceName, existingPeer.AllowedIPs, err,
|
||||
"re-add peer %s to interface %s with allowed IPs %v: %w",
|
||||
peerKey, c.deviceName, allowedIPs, err,
|
||||
)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemovePeer removes a peer and forgets its allowed IPs after a successful device write.
|
||||
func (c *KernelConfigurer) RemovePeer(peerKey string) error {
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
@@ -142,15 +167,13 @@ func (c *KernelConfigurer) RemovePeer(peerKey string) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName)
|
||||
}
|
||||
|
||||
c.allowedIPs.forget(peerKeyParsed)
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op.
|
||||
func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
||||
ipNet := net.IPNet{
|
||||
IP: allowedIP.Addr().AsSlice(),
|
||||
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
|
||||
}
|
||||
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -159,7 +182,7 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
|
||||
PublicKey: peerKeyParsed,
|
||||
UpdateOnly: true,
|
||||
ReplaceAllowedIPs: false,
|
||||
AllowedIPs: []net.IPNet{ipNet},
|
||||
AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
|
||||
}
|
||||
|
||||
config := wgtypes.Config{
|
||||
@@ -169,52 +192,69 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
|
||||
if err != nil {
|
||||
return fmt.Errorf(`received error "%w" while adding allowed Ip to peer on interface %s with settings: allowed ips %s`, err, c.deviceName, allowedIP)
|
||||
}
|
||||
|
||||
c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP})
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs.
|
||||
// A prefix not assigned to the peer is a no-op.
|
||||
func (c *KernelConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
||||
ipNet := net.IPNet{
|
||||
IP: allowedIP.Addr().AsSlice(),
|
||||
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
|
||||
}
|
||||
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse peer key: %w", err)
|
||||
}
|
||||
|
||||
existingPeer, err := c.getPeer(c.deviceName, peerKey)
|
||||
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get peer: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
newAllowedIPs := existingPeer.AllowedIPs
|
||||
|
||||
for i, existingAllowedIP := range existingPeer.AllowedIPs {
|
||||
if existingAllowedIP.String() == ipNet.String() {
|
||||
newAllowedIPs = append(existingPeer.AllowedIPs[:i], existingPeer.AllowedIPs[i+1:]...) //nolint:gocritic
|
||||
break
|
||||
}
|
||||
idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
|
||||
if idx < 0 {
|
||||
return nil
|
||||
}
|
||||
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
|
||||
|
||||
peer := wgtypes.PeerConfig{
|
||||
PublicKey: peerKeyParsed,
|
||||
UpdateOnly: true,
|
||||
ReplaceAllowedIPs: true,
|
||||
AllowedIPs: newAllowedIPs,
|
||||
AllowedIPs: prefixesToIPNets(newAllowedIPs),
|
||||
}
|
||||
|
||||
config := wgtypes.Config{
|
||||
Peers: []wgtypes.PeerConfig{peer},
|
||||
}
|
||||
err = c.configure(config)
|
||||
if err != nil {
|
||||
if err := c.configure(config); err != nil {
|
||||
return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err)
|
||||
}
|
||||
|
||||
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
|
||||
// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device
|
||||
// only for a peer the store has not seen. Dumping the device costs a netlink round trip
|
||||
// proportional to the whole network map, and this runs on every relay and ICE transition.
|
||||
func (c *KernelConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) {
|
||||
if prefixes, ok := c.allowedIPs.get(peerKey); ok {
|
||||
return prefixes, nil
|
||||
}
|
||||
|
||||
existingPeer, err := c.getPeer(c.deviceName, peerKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get peer: %w", err)
|
||||
}
|
||||
|
||||
prefixes := ipNetsToPrefixes(existingPeer.AllowedIPs)
|
||||
c.allowedIPs.set(peerKey, prefixes)
|
||||
return prefixes, nil
|
||||
}
|
||||
|
||||
// getPeer scans the device for one peer. wgtypes.Key is an array, so the comparison is a
|
||||
// plain equality: Key.String would base64 encode into a fresh allocation for every peer.
|
||||
func (c *KernelConfigurer) getPeer(ifaceName string, peerPubKey wgtypes.Key) (wgtypes.Peer, error) {
|
||||
wg, err := wgctrl.New()
|
||||
if err != nil {
|
||||
return wgtypes.Peer{}, fmt.Errorf("wgctl: %w", err)
|
||||
@@ -231,7 +271,7 @@ func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer,
|
||||
return wgtypes.Peer{}, fmt.Errorf("get device %s: %w", ifaceName, err)
|
||||
}
|
||||
for _, peer := range wgDevice.Peers {
|
||||
if peer.PublicKey.String() == peerPubKey {
|
||||
if peer.PublicKey == peerPubKey {
|
||||
return peer, nil
|
||||
}
|
||||
}
|
||||
|
||||
+120
-92
@@ -8,6 +8,7 @@ import (
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -41,31 +42,38 @@ type WGUSPConfigurer struct {
|
||||
deviceName string
|
||||
activityRecorder *bind.ActivityRecorder
|
||||
statsCache *statsCache
|
||||
allowedIPs *allowedIPStore
|
||||
|
||||
uapiListener net.Listener
|
||||
}
|
||||
|
||||
// NewUSPConfigurer creates a userspace configurer and starts its UAPI listener.
|
||||
func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
|
||||
wgCfg := &WGUSPConfigurer{
|
||||
device: device,
|
||||
deviceName: deviceName,
|
||||
activityRecorder: activityRecorder,
|
||||
allowedIPs: newAllowedIPStore(),
|
||||
}
|
||||
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
|
||||
wgCfg.startUAPI()
|
||||
return wgCfg
|
||||
}
|
||||
|
||||
// NewUSPConfigurerNoUAPI creates a userspace configurer without a UAPI listener.
|
||||
func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
|
||||
wgCfg := &WGUSPConfigurer{
|
||||
device: device,
|
||||
deviceName: deviceName,
|
||||
activityRecorder: activityRecorder,
|
||||
allowedIPs: newAllowedIPStore(),
|
||||
}
|
||||
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
|
||||
return wgCfg
|
||||
}
|
||||
|
||||
// ConfigureInterface sets the device key, port and firewall mark, replacing all peers.
|
||||
// The allowed IP mirror is reset only after the device accepts the configuration.
|
||||
func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error {
|
||||
log.Debugf("adding Wireguard private key")
|
||||
key, err := wgtypes.ParseKey(privateKey)
|
||||
@@ -80,7 +88,12 @@ func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error
|
||||
ListenPort: &port,
|
||||
}
|
||||
|
||||
return c.device.IpcSet(toWgUserspaceString(config))
|
||||
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
c.allowedIPs.reset()
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetPresharedKey sets the preshared key for a peer.
|
||||
@@ -92,14 +105,38 @@ func (c *WGUSPConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, updat
|
||||
}
|
||||
|
||||
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
|
||||
return c.device.IpcSet(toWgUserspaceString(cfg))
|
||||
if err := c.device.IpcSet(toWgUserspaceString(cfg)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Without updateOnly this creates the peer when it is absent, so the store has to
|
||||
// know about it even though no allowed IP was configured.
|
||||
if !updateOnly {
|
||||
c.allowedIPs.ensure(parsedPeerKey)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set.
|
||||
// It validates the endpoint before writing and records changes after a successful write.
|
||||
func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Everything that can fail is done before the device is touched, so a failure here
|
||||
// cannot leave the device holding a peer that the activity recorder and the allowed
|
||||
// IP store never learned about.
|
||||
var addrPort netip.AddrPort
|
||||
if endpoint != nil {
|
||||
addr, err := netip.ParseAddr(endpoint.IP.String())
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse endpoint address: %w", err)
|
||||
}
|
||||
addrPort = netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port))
|
||||
}
|
||||
|
||||
peer := wgtypes.PeerConfig{
|
||||
PublicKey: peerKeyParsed,
|
||||
ReplaceAllowedIPs: false,
|
||||
@@ -119,47 +156,27 @@ func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
|
||||
}
|
||||
|
||||
if endpoint != nil {
|
||||
addr, err := netip.ParseAddr(endpoint.IP.String())
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to parse endpoint address: %w", err)
|
||||
}
|
||||
addrPort := netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port))
|
||||
c.activityRecorder.UpsertAddress(peerKey, addrPort)
|
||||
}
|
||||
|
||||
c.allowedIPs.add(peerKeyParsed, allowedIps)
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured.
|
||||
// The UAPI cannot clear an endpoint in place, so the peer is removed and re-added with the
|
||||
// allowed IPs it already had.
|
||||
func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse peer key: %w", err)
|
||||
}
|
||||
|
||||
ipcStr, err := c.device.IpcGet()
|
||||
allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get IPC config: %w", err)
|
||||
return err
|
||||
}
|
||||
|
||||
// Parse current status to get allowed IPs for the peer
|
||||
stats, err := parseStatus(c.deviceName, ipcStr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse IPC config: %w", err)
|
||||
}
|
||||
|
||||
var allowedIPs []net.IPNet
|
||||
found := false
|
||||
for _, peer := range stats.Peers {
|
||||
if peer.PublicKey == peerKey {
|
||||
allowedIPs = peer.AllowedIPs
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return fmt.Errorf("peer %s not found", peerKey)
|
||||
}
|
||||
|
||||
// remove the peer from the WireGuard configuration
|
||||
peer := wgtypes.PeerConfig{
|
||||
PublicKey: peerKeyParsed,
|
||||
Remove: true,
|
||||
@@ -169,14 +186,13 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
|
||||
Peers: []wgtypes.PeerConfig{peer},
|
||||
}
|
||||
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
|
||||
return fmt.Errorf("failed to remove peer: %s", ipcErr)
|
||||
return fmt.Errorf("remove peer: %w", ipcErr)
|
||||
}
|
||||
|
||||
// Build the peer config
|
||||
peer = wgtypes.PeerConfig{
|
||||
PublicKey: peerKeyParsed,
|
||||
ReplaceAllowedIPs: true,
|
||||
AllowedIPs: allowedIPs,
|
||||
AllowedIPs: prefixesToIPNets(allowedIPs),
|
||||
}
|
||||
|
||||
config = wgtypes.Config{
|
||||
@@ -184,12 +200,15 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
|
||||
}
|
||||
|
||||
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
||||
return fmt.Errorf("remove endpoint address: %w", err)
|
||||
c.allowedIPs.forget(peerKeyParsed)
|
||||
return fmt.Errorf("re-add peer without endpoint: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemovePeer removes a peer, then clears its activity and allowed IP records.
|
||||
// A failed device write leaves both records intact.
|
||||
func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
@@ -204,18 +223,17 @@ func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
|
||||
config := wgtypes.Config{
|
||||
Peers: []wgtypes.PeerConfig{peer},
|
||||
}
|
||||
ipcErr := c.device.IpcSet(toWgUserspaceString(config))
|
||||
|
||||
c.activityRecorder.Remove(peerKey)
|
||||
return ipcErr
|
||||
}
|
||||
|
||||
func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
||||
ipNet := net.IPNet{
|
||||
IP: allowedIP.Addr().AsSlice(),
|
||||
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
|
||||
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
|
||||
return ipcErr
|
||||
}
|
||||
|
||||
c.activityRecorder.Remove(peerKey)
|
||||
c.allowedIPs.forget(peerKeyParsed)
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op.
|
||||
func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -224,79 +242,89 @@ func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) e
|
||||
PublicKey: peerKeyParsed,
|
||||
UpdateOnly: true,
|
||||
ReplaceAllowedIPs: false,
|
||||
AllowedIPs: []net.IPNet{ipNet},
|
||||
AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
|
||||
}
|
||||
|
||||
config := wgtypes.Config{
|
||||
Peers: []wgtypes.PeerConfig{peer},
|
||||
}
|
||||
|
||||
return c.device.IpcSet(toWgUserspaceString(config))
|
||||
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP})
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs.
|
||||
// It returns ErrAllowedIPNotFound if the prefix is not assigned to the peer.
|
||||
func (c *WGUSPConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
|
||||
ipc, err := c.device.IpcGet()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse peer key: %w", err)
|
||||
}
|
||||
|
||||
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
hexKey := hex.EncodeToString(peerKeyParsed[:])
|
||||
|
||||
lines := strings.Split(ipc, "\n")
|
||||
idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
|
||||
if idx < 0 {
|
||||
return ErrAllowedIPNotFound
|
||||
}
|
||||
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
|
||||
|
||||
peer := wgtypes.PeerConfig{
|
||||
PublicKey: peerKeyParsed,
|
||||
UpdateOnly: true,
|
||||
ReplaceAllowedIPs: true,
|
||||
AllowedIPs: []net.IPNet{},
|
||||
AllowedIPs: prefixesToIPNets(newAllowedIPs),
|
||||
}
|
||||
|
||||
foundPeer := false
|
||||
removedAllowedIP := false
|
||||
ip := allowedIP.String()
|
||||
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
|
||||
// If we're within the details of the found peer and encounter another public key,
|
||||
// this means we're starting another peer's details. So, reset the flag.
|
||||
if strings.HasPrefix(line, "public_key=") && foundPeer {
|
||||
foundPeer = false
|
||||
}
|
||||
|
||||
// Identify the peer with the specific public key
|
||||
if line == fmt.Sprintf("public_key=%s", hexKey) {
|
||||
foundPeer = true
|
||||
}
|
||||
|
||||
// If we're within the details of the found peer and find the specific allowed IP, skip this line
|
||||
if foundPeer && line == "allowed_ip="+ip {
|
||||
removedAllowedIP = true
|
||||
continue
|
||||
}
|
||||
|
||||
// Append the line to the output string
|
||||
if foundPeer && strings.HasPrefix(line, "allowed_ip=") {
|
||||
allowedIPStr := strings.TrimPrefix(line, "allowed_ip=")
|
||||
_, ipNet, err := net.ParseCIDR(allowedIPStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
peer.AllowedIPs = append(peer.AllowedIPs, *ipNet)
|
||||
}
|
||||
}
|
||||
|
||||
if !removedAllowedIP {
|
||||
return ErrAllowedIPNotFound
|
||||
}
|
||||
config := wgtypes.Config{
|
||||
Peers: []wgtypes.PeerConfig{peer},
|
||||
}
|
||||
return c.device.IpcSet(toWgUserspaceString(config))
|
||||
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
|
||||
return fmt.Errorf("remove allowed IP %s: %w", allowedIP, err)
|
||||
}
|
||||
|
||||
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
|
||||
return nil
|
||||
}
|
||||
|
||||
// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device
|
||||
// only for a peer the store has not seen. Reading them back means dumping and parsing the
|
||||
// whole device configuration, and this runs on every relay and ICE transition.
|
||||
func (c *WGUSPConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) {
|
||||
if prefixes, ok := c.allowedIPs.get(peerKey); ok {
|
||||
return prefixes, nil
|
||||
}
|
||||
|
||||
ipcStr, err := c.device.IpcGet()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get IPC config: %w", err)
|
||||
}
|
||||
|
||||
stats, err := parseStatus(c.deviceName, ipcStr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse IPC config: %w", err)
|
||||
}
|
||||
|
||||
// parseStatus reports keys in their textual form, so the comparison needs it once.
|
||||
wanted := peerKey.String()
|
||||
for _, peer := range stats.Peers {
|
||||
if peer.PublicKey != wanted {
|
||||
continue
|
||||
}
|
||||
|
||||
prefixes := ipNetsToPrefixes(peer.AllowedIPs)
|
||||
c.allowedIPs.set(peerKey, prefixes)
|
||||
return prefixes, nil
|
||||
}
|
||||
|
||||
return nil, ErrPeerNotFound
|
||||
}
|
||||
|
||||
func (c *WGUSPConfigurer) FullStats() (*Stats, error) {
|
||||
|
||||
@@ -0,0 +1,318 @@
|
||||
package configurer
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
wgconn "golang.zx2c4.com/wireguard/conn"
|
||||
wgdevice "golang.zx2c4.com/wireguard/device"
|
||||
"golang.zx2c4.com/wireguard/tun/tuntest"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface/bind"
|
||||
)
|
||||
|
||||
// newTestUSPConfigurer builds a configurer over a real wireguard-go device backed by an
|
||||
// in-memory TUN. The device stays down, so no socket is opened and no privileges are needed.
|
||||
func newTestUSPConfigurer(t *testing.T) *WGUSPConfigurer {
|
||||
t.Helper()
|
||||
|
||||
tun := tuntest.NewChannelTUN()
|
||||
dev := wgdevice.NewDevice(tun.TUN(), wgconn.NewDefaultBind(), wgdevice.NewLogger(wgdevice.LogLevelSilent, ""))
|
||||
t.Cleanup(dev.Close)
|
||||
|
||||
c := NewUSPConfigurerNoUAPI(dev, "wgtest0", bind.NewActivityRecorder())
|
||||
|
||||
key, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(t, err, "generate device private key")
|
||||
require.NoError(t, c.ConfigureInterface(key.String(), 0), "configure test device")
|
||||
|
||||
return c
|
||||
}
|
||||
|
||||
// seedPeers adds count peers, each with a /32 overlay address, and returns their public keys.
|
||||
func seedPeers(t *testing.T, c *WGUSPConfigurer, count int) []string {
|
||||
t.Helper()
|
||||
|
||||
keys := make([]string, 0, count)
|
||||
for i := 0; i < count; i++ {
|
||||
priv, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(t, err, "generate peer private key")
|
||||
pub := priv.PublicKey().String()
|
||||
|
||||
addr := netip.PrefixFrom(netip.AddrFrom4([4]byte{100, 64, byte(i >> 8), byte(i)}), 32)
|
||||
require.NoError(t, c.UpdatePeer(pub, []netip.Prefix{addr}, 25*time.Second, nil, nil), "add peer")
|
||||
keys = append(keys, pub)
|
||||
}
|
||||
return keys
|
||||
}
|
||||
|
||||
func peerAllowedIPs(t *testing.T, c *WGUSPConfigurer, peerKey string) []string {
|
||||
t.Helper()
|
||||
|
||||
stats, err := c.FullStats()
|
||||
require.NoError(t, err, "read device stats")
|
||||
|
||||
for _, p := range stats.Peers {
|
||||
if p.PublicKey != peerKey {
|
||||
continue
|
||||
}
|
||||
got := make([]string, 0, len(p.AllowedIPs))
|
||||
for _, ipNet := range p.AllowedIPs {
|
||||
got = append(got, ipNet.String())
|
||||
}
|
||||
return got
|
||||
}
|
||||
t.Fatalf("peer %s not found on device", peerKey)
|
||||
return nil
|
||||
}
|
||||
|
||||
// TestRemoveEndpointAddressPreservesRoutedAllowedIPs covers the prefixes the route manager
|
||||
// attaches to a routing peer through AddAllowedIP. Those are not known to the peer.Conn that
|
||||
// triggers the endpoint removal, so dropping them here would silently blackhole every route
|
||||
// behind that peer on each relay or ICE disconnect.
|
||||
func TestRemoveEndpointAddressPreservesRoutedAllowedIPs(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
peerKey := seedPeers(t, c, 3)[1]
|
||||
|
||||
routed := []netip.Prefix{
|
||||
netip.MustParsePrefix("10.20.0.0/16"),
|
||||
netip.MustParsePrefix("192.168.7.0/24"),
|
||||
}
|
||||
for _, prefix := range routed {
|
||||
require.NoError(t, c.AddAllowedIP(peerKey, prefix), "add routed prefix")
|
||||
}
|
||||
|
||||
before := peerAllowedIPs(t, c, peerKey)
|
||||
require.Len(t, before, 3, "peer should hold its overlay address plus both routed prefixes")
|
||||
|
||||
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
|
||||
|
||||
assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey),
|
||||
"allowed IPs must survive the endpoint removal unchanged")
|
||||
}
|
||||
|
||||
// TestRemoveEndpointAddressDoesNotScaleWithPeerCount is the regression guard for the actual
|
||||
// defect: clearing one peer's endpoint used to dump and parse the whole device, so its cost
|
||||
// grew with the size of the network map. On a routing peer with thousands of peers that dump
|
||||
// runs on every relay and ICE transition, under the interface lock.
|
||||
func TestRemoveEndpointAddressDoesNotScaleWithPeerCount(t *testing.T) {
|
||||
measure := func(peerCount int) float64 {
|
||||
c := newTestUSPConfigurer(t)
|
||||
peerKey := seedPeers(t, c, peerCount)[peerCount/2]
|
||||
|
||||
return testing.AllocsPerRun(5, func() {
|
||||
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
|
||||
})
|
||||
}
|
||||
|
||||
small := measure(64)
|
||||
large := measure(1024)
|
||||
|
||||
assert.Less(t, large, small*2,
|
||||
"clearing one endpoint allocated %.0f objects with 1024 peers against %.0f with 64: the cost still scales with the peer count",
|
||||
large, small)
|
||||
}
|
||||
|
||||
// TestRemoveEndpointAddressFallsBackToDevice covers a peer the store never saw, which is what
|
||||
// an out-of-band reconfiguration of the device leaves behind. The device stays the source of
|
||||
// truth in that case, so the allowed IPs must still be preserved.
|
||||
func TestRemoveEndpointAddressFallsBackToDevice(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
peerKey := seedPeers(t, c, 3)[1]
|
||||
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix")
|
||||
|
||||
before := peerAllowedIPs(t, c, peerKey)
|
||||
c.allowedIPs.reset()
|
||||
|
||||
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
|
||||
|
||||
assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey),
|
||||
"allowed IPs recovered from the device must be preserved")
|
||||
|
||||
recovered, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||
assert.True(t, ok, "the fallback must seed the store so the next call skips the device dump")
|
||||
assert.Len(t, recovered, 2, "seeded prefixes")
|
||||
}
|
||||
|
||||
func TestRemoveAllowedIPKeepsTheOtherPrefixes(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
peerKey := seedPeers(t, c, 3)[0]
|
||||
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||
require.NoError(t, c.AddAllowedIP(peerKey, routed), "add routed prefix")
|
||||
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("192.168.7.0/24")), "add routed prefix")
|
||||
|
||||
require.NoError(t, c.RemoveAllowedIP(peerKey, routed), "remove routed prefix")
|
||||
|
||||
assert.ElementsMatch(t, []string{"100.64.0.0/32", "192.168.7.0/24"}, peerAllowedIPs(t, c, peerKey),
|
||||
"only the removed prefix should be gone")
|
||||
|
||||
assert.ErrorIs(t, c.RemoveAllowedIP(peerKey, routed), ErrAllowedIPNotFound,
|
||||
"removing a prefix that is no longer configured must be reported")
|
||||
}
|
||||
|
||||
// TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt covers the lazy connection window documented
|
||||
// in #6863: AddAllowedIP is update-only, a silent no-op when the peer is absent, so it must not
|
||||
// leave the store claiming prefixes the device never took. RemoveEndpointAddress re-adds a peer
|
||||
// without update-only, so a phantom entry would create a peer the device had dropped, and a
|
||||
// created peer would steal those allowed IPs from whichever peer legitimately holds them.
|
||||
func TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
seedPeers(t, c, 2)
|
||||
|
||||
priv, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(t, err, "generate peer private key")
|
||||
absent := priv.PublicKey().String()
|
||||
|
||||
require.NoError(t, c.AddAllowedIP(absent, netip.MustParsePrefix("10.20.0.0/16")),
|
||||
"update-only add on an absent peer is a silent no-op")
|
||||
|
||||
stats, err := c.FullStats()
|
||||
require.NoError(t, err, "read device stats")
|
||||
require.Len(t, stats.Peers, 2, "the absent peer must not have been created by AddAllowedIP")
|
||||
|
||||
assert.ErrorIs(t, c.RemoveEndpointAddress(absent), ErrPeerNotFound,
|
||||
"clearing the endpoint of a peer the device does not have must fail")
|
||||
|
||||
stats, err = c.FullStats()
|
||||
require.NoError(t, err, "read device stats")
|
||||
assert.Len(t, stats.Peers, 2, "no peer may be created while clearing an endpoint")
|
||||
}
|
||||
|
||||
// TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer covers WireGuard's rule that an
|
||||
// allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away from
|
||||
// whichever peer held it before. UpdatePeer relies on that rule rather than removing the prefix
|
||||
// from the previous holder itself, so a prefix handed over between peers must not come back.
|
||||
func TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
keys := seedPeers(t, c, 2)
|
||||
peerA, peerB := keys[0], keys[1]
|
||||
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||
|
||||
require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A")
|
||||
require.Contains(t, peerAllowedIPs(t, c, peerA), routed.String(), "A must hold the prefix")
|
||||
|
||||
// The route moves to B. The device takes it away from A on its own.
|
||||
require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B")
|
||||
require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix")
|
||||
require.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), "the device must have taken it from A")
|
||||
|
||||
require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint")
|
||||
|
||||
assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(),
|
||||
"clearing A's endpoint must not take the prefix back from B")
|
||||
assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(),
|
||||
"B must still hold the prefix")
|
||||
}
|
||||
|
||||
// TestPresharedKeyCreatedPeerTakesPartInPrefixHandover covers a peer created by a preshared
|
||||
// key write rather than by a peer update. Rosenpass applies a peer's first key without
|
||||
// updateOnly, which creates the peer on the device, so a store that ignored that operation
|
||||
// would treat the peer as unknown and would not account for a prefix later handed over to it.
|
||||
func TestPresharedKeyCreatedPeerTakesPartInPrefixHandover(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
peerA := seedPeers(t, c, 1)[0]
|
||||
routed := netip.MustParsePrefix("10.20.0.0/16")
|
||||
require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A")
|
||||
|
||||
priv, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(t, err, "generate peer private key")
|
||||
peerB := priv.PublicKey().String()
|
||||
|
||||
psk, err := wgtypes.GenerateKey()
|
||||
require.NoError(t, err, "generate preshared key")
|
||||
require.NoError(t, c.SetPresharedKey(peerB, psk, false), "a first key creates the peer")
|
||||
|
||||
require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B")
|
||||
require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix")
|
||||
|
||||
require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint")
|
||||
|
||||
assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(),
|
||||
"clearing A's endpoint must not take the prefix back from B")
|
||||
assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must still hold the prefix")
|
||||
}
|
||||
|
||||
// TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice is the end to end form of the
|
||||
// conversion: a v4-mapped prefix must not reach the device as a zero length allowed IP,
|
||||
// which would route every v4 address to that peer.
|
||||
func TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
|
||||
priv, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(t, err, "generate peer private key")
|
||||
peerKey := priv.PublicKey().String()
|
||||
|
||||
mapped := netip.MustParsePrefix("::ffff:10.1.2.3/112")
|
||||
require.NoError(t, c.UpdatePeer(peerKey, []netip.Prefix{mapped}, 25*time.Second, nil, nil), "add peer")
|
||||
|
||||
onDevice := peerAllowedIPs(t, c, peerKey)
|
||||
assert.NotContains(t, onDevice, "0.0.0.0/0", "the device must not be given a catch-all allowed IP")
|
||||
assert.Equal(t, []string{"10.1.0.0/16"}, onDevice, "the device holds the normalized prefix")
|
||||
|
||||
recorded, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||
require.True(t, ok, "the peer must be recorded")
|
||||
require.Len(t, recorded, 1, "one prefix recorded")
|
||||
assert.Equal(t, onDevice[0], recorded[0].String(), "device and store must agree")
|
||||
}
|
||||
|
||||
// TestUpdatePeerWithAnUnusableEndpointTouchesNothing pins the ordering: the endpoint is
|
||||
// parsed before the device is configured, so a failure cannot leave the device holding a
|
||||
// peer that the store never learned about, with the prefix handover skipped along with it.
|
||||
func TestUpdatePeerWithAnUnusableEndpointTouchesNothing(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
seedPeers(t, c, 2)
|
||||
|
||||
priv, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(t, err, "generate peer private key")
|
||||
peerKey := priv.PublicKey().String()
|
||||
|
||||
// A three byte address has no textual form netip can parse back.
|
||||
endpoint := &net.UDPAddr{IP: net.IP{1, 2, 3}, Port: 51820}
|
||||
require.Error(t, c.UpdatePeer(peerKey, []netip.Prefix{netip.MustParsePrefix("10.30.0.0/16")},
|
||||
25*time.Second, endpoint, nil), "an unusable endpoint must fail the update")
|
||||
|
||||
stats, err := c.FullStats()
|
||||
require.NoError(t, err, "read device stats")
|
||||
assert.Len(t, stats.Peers, 2, "the peer must not have reached the device")
|
||||
|
||||
_, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||
assert.False(t, ok, "the peer must not have been recorded either")
|
||||
}
|
||||
|
||||
// TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses covers a removal that never reached the
|
||||
// device. A single peer removal is one write, so a failure leaves the peer on the device
|
||||
// exactly as it was, and the record still describes it; dropping it would only force the
|
||||
// next caller to read the whole device back for an answer it already had.
|
||||
func TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses(t *testing.T) {
|
||||
c := newTestUSPConfigurer(t)
|
||||
peerKey := seedPeers(t, c, 1)[0]
|
||||
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix")
|
||||
|
||||
before, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||
require.True(t, ok, "the peer must be recorded before the removal")
|
||||
require.Len(t, before, 2, "overlay address plus routed prefix")
|
||||
|
||||
// A closed device refuses every write, which is the shape of any failed removal.
|
||||
c.device.Close()
|
||||
|
||||
require.Error(t, c.RemovePeer(peerKey), "the removal must report the failure")
|
||||
|
||||
after, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
|
||||
require.True(t, ok, "a peer still on the device must stay recorded")
|
||||
assert.Equal(t, before, after, "the record must describe the peer the device kept")
|
||||
}
|
||||
|
||||
// mustParseKey turns the textual key the configurer API takes into the form the store
|
||||
// keys on.
|
||||
func mustParseKey(t *testing.T, key string) wgtypes.Key {
|
||||
t.Helper()
|
||||
|
||||
parsed, err := wgtypes.ParseKey(key)
|
||||
require.NoError(t, err, "parse peer key")
|
||||
return parsed
|
||||
}
|
||||
@@ -51,7 +51,6 @@ func ValidateMTU(mtu uint16) error {
|
||||
|
||||
type wgProxyFactory interface {
|
||||
GetProxy() wgproxy.Proxy
|
||||
GetProxyPort() uint16
|
||||
Free() error
|
||||
}
|
||||
|
||||
@@ -81,12 +80,6 @@ func (w *WGIface) GetProxy() wgproxy.Proxy {
|
||||
return w.wgProxyFactory.GetProxy()
|
||||
}
|
||||
|
||||
// GetProxyPort returns the proxy port used by the WireGuard proxy.
|
||||
// Returns 0 if no proxy port is used (e.g., for userspace WireGuard).
|
||||
func (w *WGIface) GetProxyPort() uint16 {
|
||||
return w.wgProxyFactory.GetProxyPort()
|
||||
}
|
||||
|
||||
// GetBind returns the EndpointManager userspace bind mode.
|
||||
func (w *WGIface) GetBind() device.EndpointManager {
|
||||
w.mu.Lock()
|
||||
|
||||
@@ -52,7 +52,6 @@ func (f *fakeTunDevice) Close() error {
|
||||
type fakeProxyFactory struct{}
|
||||
|
||||
func (fakeProxyFactory) GetProxy() wgproxy.Proxy { return nil }
|
||||
func (fakeProxyFactory) GetProxyPort() uint16 { return 0 }
|
||||
func (fakeProxyFactory) Free() error { return nil }
|
||||
|
||||
// TestWGIface_CloseReleasesMutexBeforeTunClose guards against a deadlock
|
||||
|
||||
@@ -6,27 +6,14 @@ import (
|
||||
"fmt"
|
||||
"os/exec"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/netbirdio/netbird/client/internal/wincmd"
|
||||
)
|
||||
|
||||
func (w *WGIface) Destroy() error {
|
||||
netshCmd := GetSystem32Command("netsh")
|
||||
netshCmd := wincmd.System32("netsh")
|
||||
out, err := exec.Command(netshCmd, "interface", "set", "interface", w.Name(), "admin=disable").CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
|
||||
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
|
||||
func GetSystem32Command(command string) string {
|
||||
_, err := exec.LookPath(command)
|
||||
if err == nil {
|
||||
return command
|
||||
}
|
||||
|
||||
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
|
||||
|
||||
return "C:\\windows\\system32\\" + command + ".exe"
|
||||
}
|
||||
|
||||
+38
-42
@@ -40,14 +40,18 @@ func init() {
|
||||
peerPubKey = peerPrivateKey.PublicKey().String()
|
||||
}
|
||||
|
||||
// testIFaceBlackList mirrors the prefixes profilemanager.DefaultInterfaceBlacklist
|
||||
// carries for the overlay interface. These tests create their own utun device, and
|
||||
// stdnet's filter probes with wgctrl every interface it is not told to skip, which
|
||||
// on a userspace WireGuard platform reaches the UAPI socket of this same process.
|
||||
// Declared here rather than imported because profilemanager imports this package.
|
||||
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
|
||||
|
||||
func TestWGIface_UpdateAddr(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||
addr := "100.64.0.1/8"
|
||||
wgPort := 33100
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -127,10 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) {
|
||||
func Test_CreateInterface(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1)
|
||||
wgIP := "10.99.99.1/32"
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
Address: wgaddr.MustParseWGAddress(wgIP),
|
||||
@@ -170,10 +171,7 @@ func Test_Close(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
|
||||
wgIP := "10.99.99.2/32"
|
||||
wgPort := 33100
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -215,10 +213,7 @@ func TestRecreation(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
|
||||
wgIP := "10.99.99.2/32"
|
||||
wgPort := 33100
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -288,10 +283,7 @@ func Test_ConfigureInterface(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3)
|
||||
wgIP := "10.99.99.5/30"
|
||||
wgPort := 33100
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
Address: wgaddr.MustParseWGAddress(wgIP),
|
||||
@@ -343,10 +335,7 @@ func Test_ConfigureInterface(t *testing.T) {
|
||||
func Test_UpdatePeer(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||
wgIP := "10.99.99.9/30"
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -413,10 +402,7 @@ func Test_UpdatePeer(t *testing.T) {
|
||||
func Test_RemovePeer(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||
wgIP := "10.99.99.13/30"
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -477,10 +463,7 @@ func Test_ConnectPeers(t *testing.T) {
|
||||
peer2wgPort := 33200
|
||||
|
||||
keepAlive := 1 * time.Second
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
guid := fmt.Sprintf("{%s}", uuid.New().String())
|
||||
device.CustomWindowsGUIDString = strings.ToLower(guid)
|
||||
@@ -516,10 +499,7 @@ func Test_ConnectPeers(t *testing.T) {
|
||||
guid = fmt.Sprintf("{%s}", uuid.New().String())
|
||||
device.CustomWindowsGUIDString = strings.ToLower(guid)
|
||||
|
||||
newNet, err = stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet = stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
optsPeer2 := WGIFaceOpts{
|
||||
IFaceName: peer2ifaceName,
|
||||
@@ -568,11 +548,14 @@ func Test_ConnectPeers(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The peers use userspace WireGuard (stdnet transport). A tight busy-loop
|
||||
// here starves the wireguard-go goroutines that process the handshake, so
|
||||
// poll on a ticker instead and yield the CPU between checks. WireGuard also
|
||||
// only retries a lost handshake initiation every REKEY_TIMEOUT (5s), which
|
||||
// is why the overall wait can occasionally stretch to tens of seconds.
|
||||
// On Linux with the kernel module both peers are kernel devices, elsewhere
|
||||
// they run on wireguard-go. A tight busy-loop here would starve the
|
||||
// wireguard-go goroutines that process the handshake, so poll on a ticker
|
||||
// instead and yield the CPU between checks. WireGuard also only retries a
|
||||
// lost handshake initiation every REKEY_TIMEOUT (5s), which is why the
|
||||
// overall wait can occasionally stretch to tens of seconds. Each side sends
|
||||
// its first initiation when its peer is configured, and the first one leaves
|
||||
// before the other device knows the peer, so that one is always wasted.
|
||||
timeout := 30 * time.Second
|
||||
timeoutChannel := time.After(timeout)
|
||||
ticker := time.NewTicker(500 * time.Millisecond)
|
||||
@@ -590,13 +573,26 @@ func Test_ConnectPeers(t *testing.T) {
|
||||
|
||||
select {
|
||||
case <-timeoutChannel:
|
||||
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String())
|
||||
// The counters tell whether initiations were sent at all, whether they
|
||||
// arrived, and whether only one direction is working.
|
||||
t.Fatalf("waiting for peer handshake timeout after %s\n%s\n%s", timeout.String(),
|
||||
describePeer(peer1ifaceName, peer2Key.PublicKey().String()),
|
||||
describePeer(peer2ifaceName, peer1Key.PublicKey().String()))
|
||||
case <-ticker.C:
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func describePeer(ifaceName, peerPubKey string) string {
|
||||
peer, err := getPeer(ifaceName, peerPubKey)
|
||||
if err != nil {
|
||||
return fmt.Sprintf("%s: peer %s: %v", ifaceName, peerPubKey, err)
|
||||
}
|
||||
return fmt.Sprintf("%s: peer %s endpoint=%v tx=%d rx=%d last_handshake=%v",
|
||||
ifaceName, peerPubKey, peer.Endpoint, peer.TransmitBytes, peer.ReceiveBytes, peer.LastHandshakeTime)
|
||||
}
|
||||
|
||||
func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
|
||||
wg, err := wgctrl.New()
|
||||
if err != nil {
|
||||
|
||||
@@ -200,10 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() {
|
||||
}
|
||||
if len(networks) > 0 {
|
||||
if m.params.Net == nil {
|
||||
var err error
|
||||
if m.params.Net, err = stdnet.NewNet(context.Background(), nil); err != nil {
|
||||
m.params.Logger.Errorf("failed to get create network: %v", err)
|
||||
}
|
||||
m.params.Net = stdnet.NewNet(context.Background(), nil)
|
||||
}
|
||||
|
||||
ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true)
|
||||
|
||||
@@ -1,32 +0,0 @@
|
||||
package ebpf
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
)
|
||||
|
||||
var (
|
||||
portRangeStart = 3128
|
||||
portRangeEnd = portRangeStart + 100
|
||||
)
|
||||
|
||||
type portLookup struct {
|
||||
}
|
||||
|
||||
func (pl portLookup) searchFreePort() (int, error) {
|
||||
for i := portRangeStart; i <= portRangeEnd; i++ {
|
||||
if pl.tryToBind(i) == nil {
|
||||
return i, nil
|
||||
}
|
||||
}
|
||||
return 0, fmt.Errorf("failed to bind free port for eBPF proxy")
|
||||
}
|
||||
|
||||
func (pl portLookup) tryToBind(port int) error {
|
||||
l, err := net.ListenPacket("udp", fmt.Sprintf(":%d", port))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_ = l.Close()
|
||||
return nil
|
||||
}
|
||||
@@ -1,45 +0,0 @@
|
||||
package ebpf
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func Test_portLookup_searchFreePort(t *testing.T) {
|
||||
pl := portLookup{}
|
||||
_, err := pl.searchFreePort()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func Test_portLookup_on_allocated(t *testing.T) {
|
||||
pl := portLookup{}
|
||||
|
||||
portRangeStart = 4128
|
||||
portRangeEnd = portRangeStart + 100
|
||||
|
||||
allocatedPort, err := allocatePort(portRangeStart)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer allocatedPort.Close()
|
||||
|
||||
fp, err := pl.searchFreePort()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if fp != (portRangeStart + 1) {
|
||||
t.Errorf("invalid free port, expected: %d, got: %d", portRangeStart+1, fp)
|
||||
}
|
||||
}
|
||||
|
||||
func allocatePort(port int) (net.PacketConn, error) {
|
||||
c, err := net.ListenPacket("udp", fmt.Sprintf(":%d", port))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return c, err
|
||||
}
|
||||
@@ -1,243 +0,0 @@
|
||||
//go:build linux && !android
|
||||
|
||||
package ebpf
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"sync"
|
||||
|
||||
"github.com/hashicorp/go-multierror"
|
||||
"github.com/pion/transport/v3"
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||
"github.com/netbirdio/netbird/client/iface/bufsize"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/rawsocket"
|
||||
"github.com/netbirdio/netbird/client/internal/ebpf"
|
||||
ebpfMgr "github.com/netbirdio/netbird/client/internal/ebpf/manager"
|
||||
nbnet "github.com/netbirdio/netbird/client/net"
|
||||
)
|
||||
|
||||
const (
|
||||
loopbackAddr = "127.0.0.1"
|
||||
)
|
||||
|
||||
// WGEBPFProxy definition for proxy with EBPF support
|
||||
type WGEBPFProxy struct {
|
||||
localWGListenPort int
|
||||
proxyPort int
|
||||
mtu uint16
|
||||
|
||||
ebpfManager ebpfMgr.Manager
|
||||
relayedConnStore map[uint16]net.Conn
|
||||
relayedConnMutex sync.Mutex
|
||||
|
||||
lastUsedPort uint16
|
||||
rawConnIPv4 net.PacketConn
|
||||
rawConnIPv6 net.PacketConn
|
||||
conn transport.UDPConn
|
||||
|
||||
ctx context.Context
|
||||
ctxCancel context.CancelFunc
|
||||
}
|
||||
|
||||
// NewWGEBPFProxy create new WGEBPFProxy instance
|
||||
func NewWGEBPFProxy(wgPort int, mtu uint16) *WGEBPFProxy {
|
||||
log.Debugf("instantiate ebpf proxy")
|
||||
wgProxy := &WGEBPFProxy{
|
||||
localWGListenPort: wgPort,
|
||||
mtu: mtu,
|
||||
ebpfManager: ebpf.GetEbpfManagerInstance(),
|
||||
relayedConnStore: make(map[uint16]net.Conn),
|
||||
}
|
||||
return wgProxy
|
||||
}
|
||||
|
||||
// Listen load ebpf program and listen the proxy
|
||||
func (p *WGEBPFProxy) Listen() error {
|
||||
pl := portLookup{}
|
||||
proxyPort, err := pl.searchFreePort()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.proxyPort = proxyPort
|
||||
|
||||
// Prepare IPv4 raw socket (required)
|
||||
p.rawConnIPv4, err = rawsocket.PrepareSenderRawSocketIPv4()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Prepare IPv6 raw socket (optional)
|
||||
p.rawConnIPv6, err = rawsocket.PrepareSenderRawSocketIPv6()
|
||||
if err != nil {
|
||||
log.Warnf("failed to prepare IPv6 raw socket, continuing with IPv4 only: %v", err)
|
||||
}
|
||||
|
||||
err = p.ebpfManager.LoadWgProxy(proxyPort, p.localWGListenPort)
|
||||
if err != nil {
|
||||
if closeErr := p.rawConnIPv4.Close(); closeErr != nil {
|
||||
log.Warnf("failed to close IPv4 raw socket: %v", closeErr)
|
||||
}
|
||||
if p.rawConnIPv6 != nil {
|
||||
if closeErr := p.rawConnIPv6.Close(); closeErr != nil {
|
||||
log.Warnf("failed to close IPv6 raw socket: %v", closeErr)
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
addr := net.UDPAddr{
|
||||
Port: proxyPort,
|
||||
IP: net.ParseIP(loopbackAddr),
|
||||
}
|
||||
|
||||
p.ctx, p.ctxCancel = context.WithCancel(context.Background())
|
||||
|
||||
conn, err := nbnet.ListenUDP("udp", &addr)
|
||||
if err != nil {
|
||||
if cErr := p.Free(); cErr != nil {
|
||||
log.Errorf("Failed to close the wgproxy: %s", cErr)
|
||||
}
|
||||
return err
|
||||
}
|
||||
p.conn = conn
|
||||
|
||||
go p.proxyToRemote()
|
||||
log.Infof("local wg proxy listening on: %d", proxyPort)
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddRelayedConn add new relayed connection for the proxy
|
||||
func (p *WGEBPFProxy) AddRelayedConn(relayedConn net.Conn) (*net.UDPAddr, error) {
|
||||
wgEndpointPort, err := p.storeRelayedConn(relayedConn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
log.Infof("relayed conn added to wg proxy store: %s, endpoint port: :%d", relayedConn.RemoteAddr(), wgEndpointPort)
|
||||
|
||||
wgEndpoint := &net.UDPAddr{
|
||||
IP: net.ParseIP(loopbackAddr),
|
||||
Port: int(wgEndpointPort),
|
||||
}
|
||||
return wgEndpoint, nil
|
||||
}
|
||||
|
||||
// Free resources except the remoteConns will be keep open.
|
||||
func (p *WGEBPFProxy) Free() error {
|
||||
log.Debugf("free up ebpf wg proxy")
|
||||
if p.ctx != nil && p.ctx.Err() != nil {
|
||||
//nolint
|
||||
return nil
|
||||
}
|
||||
|
||||
p.ctxCancel()
|
||||
|
||||
var result *multierror.Error
|
||||
if p.conn != nil {
|
||||
if err := p.conn.Close(); err != nil {
|
||||
result = multierror.Append(result, err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := p.ebpfManager.FreeWGProxy(); err != nil {
|
||||
result = multierror.Append(result, err)
|
||||
}
|
||||
|
||||
if p.rawConnIPv4 != nil {
|
||||
if err := p.rawConnIPv4.Close(); err != nil {
|
||||
result = multierror.Append(result, err)
|
||||
}
|
||||
}
|
||||
|
||||
if p.rawConnIPv6 != nil {
|
||||
if err := p.rawConnIPv6.Close(); err != nil {
|
||||
result = multierror.Append(result, err)
|
||||
}
|
||||
}
|
||||
return nberrors.FormatErrorOrNil(result)
|
||||
}
|
||||
|
||||
// GetProxyPort returns the proxy listening port.
|
||||
func (p *WGEBPFProxy) GetProxyPort() uint16 {
|
||||
return uint16(p.proxyPort)
|
||||
}
|
||||
|
||||
// proxyToRemote read messages from local WireGuard interface and forward it to remote conn
|
||||
// From this go routine has only one instance.
|
||||
func (p *WGEBPFProxy) proxyToRemote() {
|
||||
buf := make([]byte, p.mtu+bufsize.WGBufferOverhead)
|
||||
for p.ctx.Err() == nil {
|
||||
if err := p.readAndForwardPacket(buf); err != nil {
|
||||
if p.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
log.Errorf("failed to proxy packet to remote conn: %s", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *WGEBPFProxy) readAndForwardPacket(buf []byte) error {
|
||||
n, addr, err := p.conn.ReadFromUDP(buf)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read UDP packet from WG: %w", err)
|
||||
}
|
||||
|
||||
p.relayedConnMutex.Lock()
|
||||
conn, ok := p.relayedConnStore[uint16(addr.Port)]
|
||||
p.relayedConnMutex.Unlock()
|
||||
if !ok {
|
||||
if p.ctx.Err() == nil {
|
||||
log.Debugf("relayed conn not found by port because conn already has been closed: %d", addr.Port)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if _, err := conn.Write(buf[:n]); err != nil {
|
||||
return fmt.Errorf("forward local WG packet (%d) to remote relayed conn: %w", addr.Port, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *WGEBPFProxy) storeRelayedConn(relayedConn net.Conn) (uint16, error) {
|
||||
p.relayedConnMutex.Lock()
|
||||
defer p.relayedConnMutex.Unlock()
|
||||
|
||||
np, err := p.nextFreePort()
|
||||
if err != nil {
|
||||
return np, err
|
||||
}
|
||||
p.relayedConnStore[np] = relayedConn
|
||||
return np, nil
|
||||
}
|
||||
|
||||
func (p *WGEBPFProxy) removeRelayedConn(relayedConnID uint16) {
|
||||
p.relayedConnMutex.Lock()
|
||||
defer p.relayedConnMutex.Unlock()
|
||||
|
||||
_, ok := p.relayedConnStore[relayedConnID]
|
||||
if ok {
|
||||
log.Debugf("remove relayed conn from store by port: %d", relayedConnID)
|
||||
}
|
||||
delete(p.relayedConnStore, relayedConnID)
|
||||
}
|
||||
|
||||
func (p *WGEBPFProxy) nextFreePort() (uint16, error) {
|
||||
if len(p.relayedConnStore) == 65535 {
|
||||
return 0, fmt.Errorf("reached maximum relayed connection numbers")
|
||||
}
|
||||
generatePort:
|
||||
if p.lastUsedPort == 65535 {
|
||||
p.lastUsedPort = 1
|
||||
} else {
|
||||
p.lastUsedPort++
|
||||
}
|
||||
|
||||
if _, ok := p.relayedConnStore[p.lastUsedPort]; ok {
|
||||
goto generatePort
|
||||
}
|
||||
return p.lastUsedPort, nil
|
||||
}
|
||||
@@ -1,56 +0,0 @@
|
||||
//go:build linux && !android
|
||||
|
||||
package ebpf
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestWGEBPFProxy_connStore(t *testing.T) {
|
||||
wgProxy := NewWGEBPFProxy(1, 1280)
|
||||
|
||||
p, _ := wgProxy.storeRelayedConn(nil)
|
||||
if p != 1 {
|
||||
t.Errorf("invalid initial port: %d", wgProxy.lastUsedPort)
|
||||
}
|
||||
|
||||
numOfConns := 10
|
||||
for i := 0; i < numOfConns; i++ {
|
||||
p, _ = wgProxy.storeRelayedConn(nil)
|
||||
}
|
||||
if p != uint16(numOfConns)+1 {
|
||||
t.Errorf("invalid last used port: %d, expected: %d", p, numOfConns+1)
|
||||
}
|
||||
if len(wgProxy.relayedConnStore) != numOfConns+1 {
|
||||
t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), numOfConns+1)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWGEBPFProxy_portCalculation_overflow(t *testing.T) {
|
||||
wgProxy := NewWGEBPFProxy(1, 1280)
|
||||
|
||||
_, _ = wgProxy.storeRelayedConn(nil)
|
||||
wgProxy.lastUsedPort = 65535
|
||||
p, _ := wgProxy.storeRelayedConn(nil)
|
||||
|
||||
if len(wgProxy.relayedConnStore) != 2 {
|
||||
t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), 2)
|
||||
}
|
||||
|
||||
if p != 2 {
|
||||
t.Errorf("invalid last used port: %d, expected: %d", p, 2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWGEBPFProxy_portCalculation_maxConn(t *testing.T) {
|
||||
wgProxy := NewWGEBPFProxy(1, 1280)
|
||||
|
||||
for i := 0; i < 65535; i++ {
|
||||
_, _ = wgProxy.storeRelayedConn(nil)
|
||||
}
|
||||
|
||||
_, err := wgProxy.storeRelayedConn(nil)
|
||||
if err == nil {
|
||||
t.Errorf("invalid relayed conn store calculation")
|
||||
}
|
||||
}
|
||||
@@ -8,11 +8,13 @@ import (
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
|
||||
udpProxy "github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
||||
)
|
||||
|
||||
const (
|
||||
envDisableKernelWGProxy = "NB_DISABLE_KERNEL_WG_PROXY"
|
||||
// envDisableEBPFWGProxy is a deprecated alias for envDisableKernelWGProxy.
|
||||
envDisableEBPFWGProxy = "NB_DISABLE_EBPF_WG_PROXY"
|
||||
)
|
||||
|
||||
@@ -20,7 +22,7 @@ type KernelFactory struct {
|
||||
wgPort int
|
||||
mtu uint16
|
||||
|
||||
ebpfProxy *ebpf.WGEBPFProxy
|
||||
loopbackProxy *loopback.Proxy
|
||||
}
|
||||
|
||||
func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
|
||||
@@ -29,55 +31,56 @@ func NewKernelFactory(wgPort int, mtu uint16) *KernelFactory {
|
||||
mtu: mtu,
|
||||
}
|
||||
|
||||
if isEBPFDisabled() {
|
||||
if isKernelProxyDisabled() {
|
||||
log.Infof("WireGuard Proxy Factory will produce UDP proxy")
|
||||
log.Infof("eBPF WireGuard proxy is disabled via %s environment variable", envDisableEBPFWGProxy)
|
||||
return f
|
||||
}
|
||||
|
||||
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, mtu)
|
||||
if err := ebpfProxy.Listen(); err != nil {
|
||||
loopbackProxy := loopback.NewProxy(wgPort, mtu)
|
||||
if err := loopbackProxy.Listen(); err != nil {
|
||||
log.Infof("WireGuard Proxy Factory will produce UDP proxy")
|
||||
log.Warnf("failed to initialize ebpf proxy, fallback to user space proxy: %s", err)
|
||||
log.Warnf("failed to initialize loopback proxy, fallback to user space proxy: %s", err)
|
||||
return f
|
||||
}
|
||||
log.Infof("WireGuard Proxy Factory will produce eBPF proxy")
|
||||
f.ebpfProxy = ebpfProxy
|
||||
log.Infof("WireGuard Proxy Factory will produce loopback proxy")
|
||||
f.loopbackProxy = loopbackProxy
|
||||
return f
|
||||
}
|
||||
|
||||
func (w *KernelFactory) GetProxy() Proxy {
|
||||
if w.ebpfProxy == nil {
|
||||
if w.loopbackProxy == nil {
|
||||
return udpProxy.NewWGUDPProxy(w.wgPort, w.mtu)
|
||||
}
|
||||
|
||||
return ebpf.NewProxyWrapper(w.ebpfProxy)
|
||||
}
|
||||
|
||||
// GetProxyPort returns the eBPF proxy port, or 0 if eBPF is not active.
|
||||
func (w *KernelFactory) GetProxyPort() uint16 {
|
||||
if w.ebpfProxy == nil {
|
||||
return 0
|
||||
}
|
||||
return w.ebpfProxy.GetProxyPort()
|
||||
return loopback.NewProxyWrapper(w.loopbackProxy)
|
||||
}
|
||||
|
||||
func (w *KernelFactory) Free() error {
|
||||
if w.ebpfProxy == nil {
|
||||
if w.loopbackProxy == nil {
|
||||
return nil
|
||||
}
|
||||
return w.ebpfProxy.Free()
|
||||
return w.loopbackProxy.Free()
|
||||
}
|
||||
|
||||
func isEBPFDisabled() bool {
|
||||
val := os.Getenv(envDisableEBPFWGProxy)
|
||||
func isKernelProxyDisabled() bool {
|
||||
env := envDisableKernelWGProxy
|
||||
val := os.Getenv(env)
|
||||
if val == "" {
|
||||
env = envDisableEBPFWGProxy
|
||||
val = os.Getenv(env)
|
||||
}
|
||||
if val == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
disabled, err := strconv.ParseBool(val)
|
||||
if err != nil {
|
||||
log.Warnf("failed to parse %s: %v", envDisableEBPFWGProxy, err)
|
||||
log.Warnf("failed to parse %s: %v", env, err)
|
||||
return false
|
||||
}
|
||||
|
||||
if disabled {
|
||||
log.Infof("kernel WireGuard proxy is disabled via %s", env)
|
||||
}
|
||||
return disabled
|
||||
}
|
||||
|
||||
@@ -24,11 +24,6 @@ func (w *USPFactory) GetProxy() Proxy {
|
||||
return proxyBind.NewProxyBind(w.bind, w.mtu)
|
||||
}
|
||||
|
||||
// GetProxyPort returns 0 as userspace WireGuard doesn't use a separate proxy port.
|
||||
func (w *USPFactory) GetProxyPort() uint16 {
|
||||
return 0
|
||||
}
|
||||
|
||||
func (w *USPFactory) Free() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
//go:build linux && !android
|
||||
|
||||
package loopback
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
)
|
||||
|
||||
// Peer endpoints live in the upper half of 127.0.0.0/8. Everything in that
|
||||
// range is delivered to the loopback device without any address or route being
|
||||
// configured, and staying out of 127.0.0.0/9 keeps well-known squatters such as
|
||||
// 127.0.0.53 (systemd-resolved) and 127.0.1.1 out of the way.
|
||||
const (
|
||||
addrRangeBase uint32 = 0x7f800000 // 127.128.0.0
|
||||
addrRangeSize uint32 = 1 << 23 // /9
|
||||
addrRangePrefix = "127.128.0.0/9"
|
||||
)
|
||||
|
||||
// allocator hands out one loopback address per relayed connection. The address
|
||||
// is the peer's identity: WireGuard sends to it, and the proxy recovers which
|
||||
// peer a packet belongs to from the destination address.
|
||||
type allocator struct {
|
||||
cursor uint32
|
||||
}
|
||||
|
||||
// next returns the first free address at or after the cursor, wrapping once.
|
||||
// inUse reports whether an address is already handed out.
|
||||
func (a *allocator) next(inUse func(netip.Addr) bool) (netip.Addr, error) {
|
||||
for i := uint32(0); i < addrRangeSize; i++ {
|
||||
a.cursor = (a.cursor + 1) % addrRangeSize
|
||||
addr := addrFromOffset(a.cursor)
|
||||
if !addr.IsValid() {
|
||||
continue
|
||||
}
|
||||
if inUse(addr) {
|
||||
continue
|
||||
}
|
||||
return addr, nil
|
||||
}
|
||||
return netip.Addr{}, fmt.Errorf("no free endpoint address in %s", addrRangePrefix)
|
||||
}
|
||||
|
||||
// addrFromOffset maps an offset in the range to an address, skipping the .0 and
|
||||
// .255 hosts. They are unremarkable on loopback, but tools and firewall rules
|
||||
// tend to treat them as network and broadcast addresses.
|
||||
func addrFromOffset(offset uint32) netip.Addr {
|
||||
last := offset & 0xff
|
||||
if last == 0 || last == 0xff {
|
||||
return netip.Addr{}
|
||||
}
|
||||
|
||||
v := addrRangeBase + offset
|
||||
return netip.AddrFrom4([4]byte{
|
||||
byte(v >> 24),
|
||||
byte(v >> 16),
|
||||
byte(v >> 8),
|
||||
byte(v),
|
||||
})
|
||||
}
|
||||
|
||||
// inRange reports whether addr is one this proxy could have handed out.
|
||||
func inRange(addr netip.Addr) bool {
|
||||
if !addr.Is4() {
|
||||
return false
|
||||
}
|
||||
b := addr.As4()
|
||||
v := uint32(b[0])<<24 | uint32(b[1])<<16 | uint32(b[2])<<8 | uint32(b[3])
|
||||
return v >= addrRangeBase && v < addrRangeBase+addrRangeSize && b[3] != 0 && b[3] != 0xff
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
//go:build linux && !android
|
||||
|
||||
package loopback
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAllocatorHandsOutDistinctAddresses(t *testing.T) {
|
||||
var a allocator
|
||||
taken := make(map[netip.Addr]bool)
|
||||
|
||||
for i := 0; i < 1000; i++ {
|
||||
addr, err := a.next(func(candidate netip.Addr) bool { return taken[candidate] })
|
||||
if err != nil {
|
||||
t.Fatalf("allocate %d: %v", i, err)
|
||||
}
|
||||
if taken[addr] {
|
||||
t.Fatalf("address %s handed out twice", addr)
|
||||
}
|
||||
if !inRange(addr) {
|
||||
t.Fatalf("address %s outside %s", addr, addrRangePrefix)
|
||||
}
|
||||
taken[addr] = true
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocatorSkipsNetworkAndBroadcastHosts(t *testing.T) {
|
||||
var a allocator
|
||||
taken := make(map[netip.Addr]bool)
|
||||
|
||||
// enough allocations to walk past a .255/.0 boundary
|
||||
for i := 0; i < 600; i++ {
|
||||
addr, err := a.next(func(candidate netip.Addr) bool { return taken[candidate] })
|
||||
if err != nil {
|
||||
t.Fatalf("allocate %d: %v", i, err)
|
||||
}
|
||||
last := addr.As4()[3]
|
||||
if last == 0 || last == 255 {
|
||||
t.Fatalf("address %s ends in .%d", addr, last)
|
||||
}
|
||||
taken[addr] = true
|
||||
}
|
||||
}
|
||||
|
||||
func TestAllocatorReusesReleasedAddresses(t *testing.T) {
|
||||
var a allocator
|
||||
taken := make(map[netip.Addr]bool)
|
||||
inUse := func(candidate netip.Addr) bool { return taken[candidate] }
|
||||
alloc := func() netip.Addr {
|
||||
t.Helper()
|
||||
addr, err := a.next(inUse)
|
||||
if err != nil {
|
||||
t.Fatalf("allocate: %v", err)
|
||||
}
|
||||
taken[addr] = true
|
||||
return addr
|
||||
}
|
||||
|
||||
first := alloc()
|
||||
second := alloc()
|
||||
delete(taken, first)
|
||||
|
||||
// The cursor only moves forward, so a released address comes back after a
|
||||
// wrap. Park the cursor near the end of the range instead of allocating
|
||||
// 2^23 addresses: the next call takes the last usable address, and the one
|
||||
// after that wraps past the skipped .255 and .0 hosts to the released one.
|
||||
a.cursor = addrRangeSize - 3
|
||||
last := alloc()
|
||||
if want := netip.MustParseAddr("127.255.255.254"); last != want {
|
||||
t.Fatalf("expected the last usable address %s before the wrap, got %s", want, last)
|
||||
}
|
||||
|
||||
if reused := alloc(); reused != first {
|
||||
t.Fatalf("expected the released address %s after the wrap, got %s", first, reused)
|
||||
}
|
||||
|
||||
// second is still held, so the allocator must step over it.
|
||||
if next := alloc(); next == second {
|
||||
t.Fatalf("allocator handed out %s while it was still in use", second)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInRange(t *testing.T) {
|
||||
tests := []struct {
|
||||
addr string
|
||||
want bool
|
||||
}{
|
||||
{"127.128.0.1", true},
|
||||
{"127.255.255.254", true},
|
||||
{"127.128.0.0", false}, // network host, never handed out
|
||||
{"127.128.5.255", false}, // broadcast host, never handed out
|
||||
{"127.127.255.255", false}, // below the range, where 127.0.0.53 and friends live
|
||||
{"127.0.0.1", false},
|
||||
{"127.0.0.53", false},
|
||||
{"127.0.1.1", false},
|
||||
{"128.0.0.1", false},
|
||||
{"10.0.0.1", false},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
addr := netip.MustParseAddr(tc.addr)
|
||||
if got := inRange(addr); got != tc.want {
|
||||
t.Errorf("inRange(%s) = %v, want %v", tc.addr, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestInRangeIgnoresIPv6(t *testing.T) {
|
||||
if inRange(netip.MustParseAddr("::1")) {
|
||||
t.Error("inRange(::1) = true, want false")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,291 @@
|
||||
//go:build linux && !android
|
||||
|
||||
package loopback
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"github.com/hashicorp/go-multierror"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||
"github.com/netbirdio/netbird/client/iface/bufsize"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/rawsocket"
|
||||
)
|
||||
|
||||
const (
|
||||
loopbackDevice = "lo"
|
||||
|
||||
portRangeStart = 3128
|
||||
portRangeEnd = portRangeStart + 100
|
||||
)
|
||||
|
||||
// Proxy forwards packets between relayed connections and a local kernel
|
||||
// WireGuard instance. Every relayed peer gets its own loopback address as its
|
||||
// WireGuard endpoint, so a single socket serves all of them: the destination
|
||||
// address of an incoming packet identifies the peer.
|
||||
type Proxy struct {
|
||||
localWGListenPort int
|
||||
mtu uint16
|
||||
proxyPort int
|
||||
|
||||
conn *net.UDPConn
|
||||
packetConn *ipv4.PacketConn
|
||||
loIndex int
|
||||
rawConnIPv4 net.PacketConn
|
||||
rawConnIPv6 net.PacketConn
|
||||
|
||||
relayedConnMutex sync.Mutex
|
||||
relayedConnStore map[netip.Addr]net.Conn
|
||||
addrs allocator
|
||||
|
||||
ctx context.Context
|
||||
ctxCancel context.CancelFunc
|
||||
}
|
||||
|
||||
// NewProxy creates a proxy for the WireGuard instance listening on wgPort.
|
||||
func NewProxy(wgPort int, mtu uint16) *Proxy {
|
||||
log.Debugf("instantiate loopback wg proxy")
|
||||
return &Proxy{
|
||||
localWGListenPort: wgPort,
|
||||
mtu: mtu,
|
||||
relayedConnStore: make(map[netip.Addr]net.Conn),
|
||||
}
|
||||
}
|
||||
|
||||
// Listen opens the shared socket and starts forwarding WireGuard packets to the
|
||||
// relayed connections.
|
||||
func (p *Proxy) Listen() error {
|
||||
rawConnIPv4, err := rawsocket.PrepareSenderRawSocketIPv4()
|
||||
if err != nil {
|
||||
return fmt.Errorf("prepare IPv4 raw socket: %w", err)
|
||||
}
|
||||
p.rawConnIPv4 = rawConnIPv4
|
||||
|
||||
p.rawConnIPv6, err = rawsocket.PrepareSenderRawSocketIPv6()
|
||||
if err != nil {
|
||||
log.Warnf("failed to prepare IPv6 raw socket, continuing with IPv4 only: %v", err)
|
||||
}
|
||||
|
||||
loopback, err := net.InterfaceByName(loopbackDevice)
|
||||
if err != nil {
|
||||
if freeErr := p.Free(); freeErr != nil {
|
||||
log.Errorf("failed to free the wgproxy: %s", freeErr)
|
||||
}
|
||||
return fmt.Errorf("look up %s: %w", loopbackDevice, err)
|
||||
}
|
||||
p.loIndex = loopback.Index
|
||||
|
||||
if err := p.listen(); err != nil {
|
||||
if freeErr := p.Free(); freeErr != nil {
|
||||
log.Errorf("failed to free the wgproxy: %s", freeErr)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
p.ctx, p.ctxCancel = context.WithCancel(context.Background())
|
||||
|
||||
go p.proxyToRemote()
|
||||
log.Infof("local wg proxy listening on %s:%d", addrRangePrefix, p.proxyPort)
|
||||
return nil
|
||||
}
|
||||
|
||||
// listen binds the shared socket on the first free port of the range. The bind
|
||||
// has to be a wildcard one to receive every peer address in the range, so it is
|
||||
// restricted to the loopback device: without that the port would be reachable
|
||||
// on every interface.
|
||||
func (p *Proxy) listen() error {
|
||||
var lastErr error
|
||||
for port := portRangeStart; port <= portRangeEnd; port++ {
|
||||
err := p.listenOn(port)
|
||||
if err == nil {
|
||||
p.proxyPort = port
|
||||
return nil
|
||||
}
|
||||
lastErr = err
|
||||
}
|
||||
return fmt.Errorf("bind proxy port in range %d-%d: %w", portRangeStart, portRangeEnd, lastErr)
|
||||
}
|
||||
|
||||
func (p *Proxy) listenOn(proxyPort int) error {
|
||||
lc := net.ListenConfig{
|
||||
Control: func(_, _ string, c syscall.RawConn) error {
|
||||
var sockErr error
|
||||
if err := c.Control(func(fd uintptr) {
|
||||
if err := unix.SetsockoptString(int(fd), unix.SOL_SOCKET, unix.SO_BINDTODEVICE, loopbackDevice); err != nil {
|
||||
sockErr = fmt.Errorf("bind to %s: %w", loopbackDevice, err)
|
||||
return
|
||||
}
|
||||
}); err != nil {
|
||||
return fmt.Errorf("control socket: %w", err)
|
||||
}
|
||||
return sockErr
|
||||
},
|
||||
}
|
||||
|
||||
conn, err := lc.ListenPacket(context.Background(), "udp4", fmt.Sprintf(":%d", proxyPort))
|
||||
if err != nil {
|
||||
return fmt.Errorf("listen on :%d: %w", proxyPort, err)
|
||||
}
|
||||
|
||||
udpConn, ok := conn.(*net.UDPConn)
|
||||
if !ok {
|
||||
if closeErr := conn.Close(); closeErr != nil {
|
||||
log.Errorf("failed to close proxy conn: %s", closeErr)
|
||||
}
|
||||
return fmt.Errorf("unexpected conn type %T", conn)
|
||||
}
|
||||
|
||||
packetConn := ipv4.NewPacketConn(udpConn)
|
||||
// the destination address carries the peer identity, the interface index is
|
||||
// checked on receive as a second line of defense behind SO_BINDTODEVICE
|
||||
if err := packetConn.SetControlMessage(ipv4.FlagDst|ipv4.FlagInterface, true); err != nil {
|
||||
if closeErr := udpConn.Close(); closeErr != nil {
|
||||
log.Errorf("failed to close proxy conn: %s", closeErr)
|
||||
}
|
||||
return fmt.Errorf("request destination address: %w", err)
|
||||
}
|
||||
|
||||
p.conn = udpConn
|
||||
p.packetConn = packetConn
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddRelayedConn assigns an endpoint address to the relayed connection and
|
||||
// returns the address WireGuard should send to, along with the key the
|
||||
// connection is stored under.
|
||||
func (p *Proxy) AddRelayedConn(relayedConn net.Conn) (*net.UDPAddr, netip.Addr, error) {
|
||||
addr, err := p.storeRelayedConn(relayedConn)
|
||||
if err != nil {
|
||||
return nil, netip.Addr{}, err
|
||||
}
|
||||
|
||||
log.Infof("relayed conn added to wg proxy store: %s, endpoint address: %s", relayedConn.RemoteAddr(), addr)
|
||||
|
||||
return &net.UDPAddr{
|
||||
IP: addr.AsSlice(),
|
||||
Port: p.proxyPort,
|
||||
}, addr, nil
|
||||
}
|
||||
|
||||
// Free releases the proxy resources. The relayed connections are left open.
|
||||
func (p *Proxy) Free() error {
|
||||
log.Debugf("free up loopback wg proxy")
|
||||
if p.ctx != nil && p.ctx.Err() != nil {
|
||||
//nolint
|
||||
return nil
|
||||
}
|
||||
|
||||
if p.ctxCancel != nil {
|
||||
p.ctxCancel()
|
||||
}
|
||||
|
||||
var result *multierror.Error
|
||||
if p.conn != nil {
|
||||
if err := p.conn.Close(); err != nil {
|
||||
result = multierror.Append(result, err)
|
||||
}
|
||||
}
|
||||
|
||||
if p.rawConnIPv4 != nil {
|
||||
if err := p.rawConnIPv4.Close(); err != nil {
|
||||
result = multierror.Append(result, err)
|
||||
}
|
||||
}
|
||||
|
||||
if p.rawConnIPv6 != nil {
|
||||
if err := p.rawConnIPv6.Close(); err != nil {
|
||||
result = multierror.Append(result, err)
|
||||
}
|
||||
}
|
||||
return nberrors.FormatErrorOrNil(result)
|
||||
}
|
||||
|
||||
// proxyToRemote reads packets from the local WireGuard instance and forwards
|
||||
// them to the relayed connection the destination address belongs to.
|
||||
func (p *Proxy) proxyToRemote() {
|
||||
buf := make([]byte, p.mtu+bufsize.WGBufferOverhead)
|
||||
for p.ctx.Err() == nil {
|
||||
if err := p.readAndForwardPacket(buf); err != nil {
|
||||
if p.ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
log.Errorf("failed to proxy packet to remote conn: %s", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *Proxy) readAndForwardPacket(buf []byte) error {
|
||||
n, cm, _, err := p.packetConn.ReadFrom(buf)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read UDP packet from WG: %w", err)
|
||||
}
|
||||
|
||||
if cm == nil {
|
||||
return fmt.Errorf("no control message on packet")
|
||||
}
|
||||
|
||||
if cm.IfIndex != p.loIndex {
|
||||
log.Tracef("dropping packet received on interface %d instead of %s", cm.IfIndex, loopbackDevice)
|
||||
return nil
|
||||
}
|
||||
|
||||
dst, ok := netip.AddrFromSlice(cm.Dst.To4())
|
||||
if !ok || !inRange(dst) {
|
||||
log.Tracef("dropping packet for unexpected destination %s", cm.Dst)
|
||||
return nil
|
||||
}
|
||||
|
||||
p.relayedConnMutex.Lock()
|
||||
conn, ok := p.relayedConnStore[dst]
|
||||
p.relayedConnMutex.Unlock()
|
||||
if !ok {
|
||||
if p.ctx.Err() == nil {
|
||||
log.Debugf("relayed conn not found by address because conn already has been closed: %s", dst)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if _, err := conn.Write(buf[:n]); err != nil {
|
||||
return fmt.Errorf("forward local WG packet (%s) to remote relayed conn: %w", dst, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *Proxy) storeRelayedConn(relayedConn net.Conn) (netip.Addr, error) {
|
||||
p.relayedConnMutex.Lock()
|
||||
defer p.relayedConnMutex.Unlock()
|
||||
|
||||
addr, err := p.addrs.next(func(a netip.Addr) bool {
|
||||
_, ok := p.relayedConnStore[a]
|
||||
return ok
|
||||
})
|
||||
if err != nil {
|
||||
return netip.Addr{}, err
|
||||
}
|
||||
|
||||
p.relayedConnStore[addr] = relayedConn
|
||||
return addr, nil
|
||||
}
|
||||
|
||||
// removeRelayedConn releases an endpoint address. It only removes the entry
|
||||
// while it still belongs to relayedConn, so a late release cannot take an
|
||||
// address away from the peer it was handed to next.
|
||||
func (p *Proxy) removeRelayedConn(addr netip.Addr, relayedConn net.Conn) {
|
||||
p.relayedConnMutex.Lock()
|
||||
defer p.relayedConnMutex.Unlock()
|
||||
|
||||
if stored, ok := p.relayedConnStore[addr]; !ok || stored != relayedConn {
|
||||
return
|
||||
}
|
||||
|
||||
log.Debugf("remove relayed conn from store by address: %s", addr)
|
||||
delete(p.relayedConnStore, addr)
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
//go:build linux && !android && privileged
|
||||
|
||||
package loopback
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
const testWGPort = 51862
|
||||
|
||||
// relayEnd stands in for a relayed connection: the proxy writes what it read
|
||||
// from WireGuard into it, and the test reads it back out here.
|
||||
func relayEnd(t *testing.T) (proxySide net.Conn, testSide *net.UDPConn) {
|
||||
t.Helper()
|
||||
|
||||
testSide, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
|
||||
if err != nil {
|
||||
t.Fatalf("relay listener: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
if err := testSide.Close(); err != nil {
|
||||
t.Logf("close relay listener: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
proxySide, err = net.Dial("udp", testSide.LocalAddr().String())
|
||||
if err != nil {
|
||||
t.Fatalf("relay conn: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
if err := proxySide.Close(); err != nil {
|
||||
t.Logf("close relay conn: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
return proxySide, testSide
|
||||
}
|
||||
|
||||
// TestProxyDemuxesByDestinationAddress is the core of the design: one socket
|
||||
// serves every peer, and the destination address decides which relayed
|
||||
// connection a WireGuard packet belongs to.
|
||||
func TestProxyDemuxesByDestinationAddress(t *testing.T) {
|
||||
proxy := NewProxy(testWGPort, 1280)
|
||||
if err := proxy.Listen(); err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := proxy.Free(); err != nil {
|
||||
t.Errorf("free proxy: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
const peers = 3
|
||||
endpoints := make([]*net.UDPAddr, 0, peers)
|
||||
readers := make([]*net.UDPConn, 0, peers)
|
||||
for i := 0; i < peers; i++ {
|
||||
proxySide, testSide := relayEnd(t)
|
||||
endpoint, _, err := proxy.AddRelayedConn(proxySide)
|
||||
if err != nil {
|
||||
t.Fatalf("add relayed conn %d: %v", i, err)
|
||||
}
|
||||
if endpoint.Port != proxy.proxyPort {
|
||||
t.Errorf("peer %d endpoint port = %d, want the shared proxy port %d", i, endpoint.Port, proxy.proxyPort)
|
||||
}
|
||||
endpoints = append(endpoints, endpoint)
|
||||
readers = append(readers, testSide)
|
||||
}
|
||||
|
||||
// every peer must have its own address, otherwise they are indistinguishable
|
||||
seen := make(map[string]bool, peers)
|
||||
for i, endpoint := range endpoints {
|
||||
if seen[endpoint.IP.String()] {
|
||||
t.Fatalf("peer %d reuses endpoint address %s", i, endpoint.IP)
|
||||
}
|
||||
seen[endpoint.IP.String()] = true
|
||||
}
|
||||
|
||||
wgSock, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")})
|
||||
if err != nil {
|
||||
t.Fatalf("wg socket: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := wgSock.Close(); err != nil {
|
||||
t.Logf("close wg socket: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
for i, endpoint := range endpoints {
|
||||
payload := []byte{byte(i), 'p', 'k', 't'}
|
||||
if _, err := wgSock.WriteTo(payload, endpoint); err != nil {
|
||||
t.Fatalf("write to peer %d endpoint %s: %v", i, endpoint, err)
|
||||
}
|
||||
|
||||
buf := make([]byte, 1500)
|
||||
if err := readers[i].SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
|
||||
t.Fatalf("set read deadline: %v", err)
|
||||
}
|
||||
n, _, err := readers[i].ReadFrom(buf)
|
||||
if err != nil {
|
||||
t.Fatalf("peer %d did not receive its packet: %v", i, err)
|
||||
}
|
||||
if string(buf[:n]) != string(payload) {
|
||||
t.Errorf("peer %d got %q, want %q", i, buf[:n], payload)
|
||||
}
|
||||
|
||||
// no other peer may see it
|
||||
for j, other := range readers {
|
||||
if j == i {
|
||||
continue
|
||||
}
|
||||
if err := other.SetReadDeadline(time.Now().Add(100 * time.Millisecond)); err != nil {
|
||||
t.Fatalf("set read deadline: %v", err)
|
||||
}
|
||||
if _, _, err := other.ReadFrom(buf); err == nil {
|
||||
t.Errorf("packet for peer %d also delivered to peer %d", i, j)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestProxyDropsPacketsOutsideTheRange guards the wildcard bind: anything that
|
||||
// is not addressed to a handed-out endpoint must not reach a relayed peer.
|
||||
func TestProxyDropsPacketsOutsideTheRange(t *testing.T) {
|
||||
proxy := NewProxy(testWGPort+1, 1280)
|
||||
if err := proxy.Listen(); err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := proxy.Free(); err != nil {
|
||||
t.Errorf("free proxy: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
proxySide, testSide := relayEnd(t)
|
||||
if _, _, err := proxy.AddRelayedConn(proxySide); err != nil {
|
||||
t.Fatalf("add relayed conn: %v", err)
|
||||
}
|
||||
|
||||
sender, err := net.Dial("udp", net.JoinHostPort("127.0.0.1", strconv.Itoa(proxy.proxyPort)))
|
||||
if err != nil {
|
||||
t.Fatalf("sender: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := sender.Close(); err != nil {
|
||||
t.Logf("close sender: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := sender.Write([]byte("stray")); err != nil {
|
||||
t.Fatalf("write stray packet: %v", err)
|
||||
}
|
||||
|
||||
buf := make([]byte, 1500)
|
||||
if err := testSide.SetReadDeadline(time.Now().Add(500 * time.Millisecond)); err != nil {
|
||||
t.Fatalf("set read deadline: %v", err)
|
||||
}
|
||||
if _, _, err := testSide.ReadFrom(buf); err == nil {
|
||||
t.Error("packet addressed to 127.0.0.1 was forwarded to a relayed peer")
|
||||
}
|
||||
}
|
||||
|
||||
// A wrapper that is closed before it starts forwarding still has to give its
|
||||
// endpoint address back, otherwise the range leaks an address per attempt.
|
||||
func TestClosingBeforeWorkReleasesTheAddress(t *testing.T) {
|
||||
proxy := NewProxy(testWGPort+2, 1280)
|
||||
if err := proxy.Listen(); err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := proxy.Free(); err != nil {
|
||||
t.Errorf("free proxy: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
proxySide, _ := relayEnd(t)
|
||||
wrapper := NewProxyWrapper(proxy)
|
||||
if err := wrapper.AddRelayedConn(context.Background(), nil, proxySide); err != nil {
|
||||
t.Fatalf("add relayed conn: %v", err)
|
||||
}
|
||||
|
||||
if got := len(proxy.relayedConnStore); got != 1 {
|
||||
t.Fatalf("store holds %d entries after adding one conn, want 1", got)
|
||||
}
|
||||
|
||||
if err := wrapper.CloseConn(); err != nil {
|
||||
t.Fatalf("close conn: %v", err)
|
||||
}
|
||||
|
||||
if got := len(proxy.relayedConnStore); got != 0 {
|
||||
t.Errorf("store holds %d entries after close, want 0", got)
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
//go:build linux && !android
|
||||
|
||||
package ebpf
|
||||
package loopback
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"sync"
|
||||
|
||||
"github.com/google/gopacket"
|
||||
@@ -95,13 +96,14 @@ func NewPacketHeaders(localWGListenPort int, endpoint *net.UDPAddr) (*PacketHead
|
||||
|
||||
// ProxyWrapper help to keep the remoteConn instance for net.Conn.Close function call
|
||||
type ProxyWrapper struct {
|
||||
wgeBPFProxy *WGEBPFProxy
|
||||
proxy *Proxy
|
||||
|
||||
remoteConn net.Conn
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
|
||||
wgRelayedEndpointAddr *net.UDPAddr
|
||||
peerAddr netip.Addr
|
||||
headers *PacketHeaders
|
||||
headerCurrentUsed *PacketHeaders
|
||||
rawConn net.PacketConn
|
||||
@@ -113,36 +115,44 @@ type ProxyWrapper struct {
|
||||
closeListener *listener.CloseListener
|
||||
}
|
||||
|
||||
func NewProxyWrapper(proxy *WGEBPFProxy) *ProxyWrapper {
|
||||
func NewProxyWrapper(proxy *Proxy) *ProxyWrapper {
|
||||
return &ProxyWrapper{
|
||||
wgeBPFProxy: proxy,
|
||||
proxy: proxy,
|
||||
pausedCond: sync.NewCond(&sync.Mutex{}),
|
||||
closeListener: listener.NewCloseListener(),
|
||||
}
|
||||
}
|
||||
|
||||
func (p *ProxyWrapper) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error {
|
||||
addr, err := p.wgeBPFProxy.AddRelayedConn(remoteConn)
|
||||
addr, peerAddr, err := p.proxy.AddRelayedConn(remoteConn)
|
||||
if err != nil {
|
||||
return fmt.Errorf("add relayed conn: %w", err)
|
||||
}
|
||||
|
||||
headers, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, addr)
|
||||
// the endpoint address is otherwise only released by the forwarding
|
||||
// goroutine, which never starts when the setup below fails
|
||||
release := func() { p.proxy.removeRelayedConn(peerAddr, remoteConn) }
|
||||
|
||||
headers, err := NewPacketHeaders(p.proxy.localWGListenPort, addr)
|
||||
if err != nil {
|
||||
release()
|
||||
return fmt.Errorf("create packet sender: %w", err)
|
||||
}
|
||||
|
||||
// Check if required raw connection is available
|
||||
if !headers.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil {
|
||||
if !headers.isIPv4 && p.proxy.rawConnIPv6 == nil {
|
||||
release()
|
||||
return errIPv6ConnNotAvailable
|
||||
}
|
||||
if headers.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil {
|
||||
if headers.isIPv4 && p.proxy.rawConnIPv4 == nil {
|
||||
release()
|
||||
return errIPv4ConnNotAvailable
|
||||
}
|
||||
|
||||
p.remoteConn = remoteConn
|
||||
p.ctx, p.cancel = context.WithCancel(ctx)
|
||||
p.wgRelayedEndpointAddr = addr
|
||||
p.peerAddr = peerAddr
|
||||
p.headers = headers
|
||||
p.rawConn = p.selectRawConn(headers)
|
||||
return nil
|
||||
@@ -193,18 +203,18 @@ func (p *ProxyWrapper) RedirectAs(endpoint *net.UDPAddr) {
|
||||
return
|
||||
}
|
||||
|
||||
header, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, endpoint)
|
||||
header, err := NewPacketHeaders(p.proxy.localWGListenPort, endpoint)
|
||||
if err != nil {
|
||||
log.Errorf("failed to create packet headers: %s", err)
|
||||
return
|
||||
}
|
||||
|
||||
// Check if required raw connection is available
|
||||
if !header.isIPv4 && p.wgeBPFProxy.rawConnIPv6 == nil {
|
||||
if !header.isIPv4 && p.proxy.rawConnIPv6 == nil {
|
||||
log.Error(errIPv6ConnNotAvailable)
|
||||
return
|
||||
}
|
||||
if header.isIPv4 && p.wgeBPFProxy.rawConnIPv4 == nil {
|
||||
if header.isIPv4 && p.proxy.rawConnIPv4 == nil {
|
||||
log.Error(errIPv4ConnNotAvailable)
|
||||
return
|
||||
}
|
||||
@@ -240,6 +250,10 @@ func (p *ProxyWrapper) CloseConn() error {
|
||||
|
||||
p.closeListener.SetCloseListener(nil)
|
||||
|
||||
// releases the endpoint address for a wrapper that was never started, and
|
||||
// is a no-op once the forwarding goroutine has released it
|
||||
p.proxy.removeRelayedConn(p.peerAddr, p.remoteConn)
|
||||
|
||||
p.pausedCond.L.Lock()
|
||||
p.paused = false
|
||||
p.pausedCond.Signal()
|
||||
@@ -252,9 +266,9 @@ func (p *ProxyWrapper) CloseConn() error {
|
||||
}
|
||||
|
||||
func (p *ProxyWrapper) proxyToLocal(ctx context.Context) {
|
||||
defer p.wgeBPFProxy.removeRelayedConn(uint16(p.wgRelayedEndpointAddr.Port))
|
||||
defer p.proxy.removeRelayedConn(p.peerAddr, p.remoteConn)
|
||||
|
||||
buf := make([]byte, p.wgeBPFProxy.mtu+bufsize.WGBufferOverhead)
|
||||
buf := make([]byte, p.proxy.mtu+bufsize.WGBufferOverhead)
|
||||
for {
|
||||
n, err := p.readFromRemote(ctx, buf)
|
||||
if err != nil {
|
||||
@@ -286,7 +300,7 @@ func (p *ProxyWrapper) readFromRemote(ctx context.Context, buf []byte) (int, err
|
||||
}
|
||||
p.closeListener.Notify()
|
||||
if !errors.Is(err, io.EOF) {
|
||||
log.Errorf("failed to read from relayed conn (endpoint: :%d): %s", p.wgRelayedEndpointAddr.Port, err)
|
||||
log.Errorf("failed to read from relayed conn (endpoint: %s): %s", p.wgRelayedEndpointAddr, err)
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
@@ -314,7 +328,7 @@ func (p *ProxyWrapper) sendPkg(data []byte, header *PacketHeaders) error {
|
||||
|
||||
func (p *ProxyWrapper) selectRawConn(header *PacketHeaders) net.PacketConn {
|
||||
if header.isIPv4 {
|
||||
return p.wgeBPFProxy.rawConnIPv4
|
||||
return p.proxy.rawConnIPv4
|
||||
}
|
||||
return p.wgeBPFProxy.rawConnIPv6
|
||||
return p.proxy.rawConnIPv6
|
||||
}
|
||||
@@ -9,25 +9,25 @@ import (
|
||||
"github.com/netbirdio/netbird/client/iface/bind"
|
||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||
bindproxy "github.com/netbirdio/netbird/client/iface/wgproxy/bind"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
||||
)
|
||||
|
||||
func seedProxies() ([]proxyInstance, error) {
|
||||
pl := make([]proxyInstance, 0)
|
||||
|
||||
ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280)
|
||||
if err := ebpfProxy.Listen(); err != nil {
|
||||
return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err)
|
||||
loopbackProxy := loopback.NewProxy(51831, 1280)
|
||||
if err := loopbackProxy.Listen(); err != nil {
|
||||
return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err)
|
||||
}
|
||||
|
||||
pEbpf := proxyInstance{
|
||||
name: "ebpf kernel proxy",
|
||||
proxy: ebpf.NewProxyWrapper(ebpfProxy),
|
||||
pLoopback := proxyInstance{
|
||||
name: "loopback kernel proxy",
|
||||
proxy: loopback.NewProxyWrapper(loopbackProxy),
|
||||
wgPort: 51831,
|
||||
closeFn: ebpfProxy.Free,
|
||||
closeFn: loopbackProxy.Free,
|
||||
}
|
||||
pl = append(pl, pEbpf)
|
||||
pl = append(pl, pLoopback)
|
||||
|
||||
pUDP := proxyInstance{
|
||||
name: "udp kernel proxy",
|
||||
@@ -42,18 +42,18 @@ func seedProxies() ([]proxyInstance, error) {
|
||||
func seedProxyForProxyCloseByRemoteConn() ([]proxyInstance, error) {
|
||||
pl := make([]proxyInstance, 0)
|
||||
|
||||
ebpfProxy := ebpf.NewWGEBPFProxy(51831, 1280)
|
||||
if err := ebpfProxy.Listen(); err != nil {
|
||||
return nil, fmt.Errorf("failed to initialize ebpf proxy: %s", err)
|
||||
loopbackProxy := loopback.NewProxy(51831, 1280)
|
||||
if err := loopbackProxy.Listen(); err != nil {
|
||||
return nil, fmt.Errorf("failed to initialize loopback proxy: %s", err)
|
||||
}
|
||||
|
||||
pEbpf := proxyInstance{
|
||||
name: "ebpf kernel proxy",
|
||||
proxy: ebpf.NewProxyWrapper(ebpfProxy),
|
||||
pLoopback := proxyInstance{
|
||||
name: "loopback kernel proxy",
|
||||
proxy: loopback.NewProxyWrapper(loopbackProxy),
|
||||
wgPort: 51831,
|
||||
closeFn: ebpfProxy.Free,
|
||||
closeFn: loopbackProxy.Free,
|
||||
}
|
||||
pl = append(pl, pEbpf)
|
||||
pl = append(pl, pLoopback)
|
||||
|
||||
pUDP := proxyInstance{
|
||||
name: "udp kernel proxy",
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/ebpf"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/loopback"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy/udp"
|
||||
)
|
||||
|
||||
@@ -198,20 +198,20 @@ func testRedirectAs(t *testing.T, proxy Proxy, wgPort int, nbAddr, p2pEndpoint *
|
||||
}
|
||||
}
|
||||
|
||||
// TestRedirectAs_eBPF_IPv4 tests RedirectAs with eBPF proxy using IPv4 addresses
|
||||
func TestRedirectAs_eBPF_IPv4(t *testing.T) {
|
||||
// TestRedirectAs_Loopback_IPv4 tests RedirectAs with the loopback proxy using IPv4 addresses
|
||||
func TestRedirectAs_Loopback_IPv4(t *testing.T) {
|
||||
wgPort := 51850
|
||||
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
|
||||
if err := ebpfProxy.Listen(); err != nil {
|
||||
t.Fatalf("failed to initialize ebpf proxy: %v", err)
|
||||
loopbackProxy := loopback.NewProxy(wgPort, 1280)
|
||||
if err := loopbackProxy.Listen(); err != nil {
|
||||
t.Fatalf("failed to initialize loopback proxy: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := ebpfProxy.Free(); err != nil {
|
||||
t.Errorf("failed to free ebpf proxy: %v", err)
|
||||
if err := loopbackProxy.Free(); err != nil {
|
||||
t.Errorf("failed to free loopback proxy: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
proxy := ebpf.NewProxyWrapper(ebpfProxy)
|
||||
proxy := loopback.NewProxyWrapper(loopbackProxy)
|
||||
|
||||
// NetBird UDP address of the remote peer
|
||||
nbAddr := &net.UDPAddr{
|
||||
@@ -227,20 +227,20 @@ func TestRedirectAs_eBPF_IPv4(t *testing.T) {
|
||||
testRedirectAs(t, proxy, wgPort, nbAddr, p2pEndpoint)
|
||||
}
|
||||
|
||||
// TestRedirectAs_eBPF_IPv6 tests RedirectAs with eBPF proxy using IPv6 addresses
|
||||
func TestRedirectAs_eBPF_IPv6(t *testing.T) {
|
||||
// TestRedirectAs_Loopback_IPv6 tests RedirectAs with the loopback proxy using IPv6 addresses
|
||||
func TestRedirectAs_Loopback_IPv6(t *testing.T) {
|
||||
wgPort := 51851
|
||||
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
|
||||
if err := ebpfProxy.Listen(); err != nil {
|
||||
t.Fatalf("failed to initialize ebpf proxy: %v", err)
|
||||
loopbackProxy := loopback.NewProxy(wgPort, 1280)
|
||||
if err := loopbackProxy.Listen(); err != nil {
|
||||
t.Fatalf("failed to initialize loopback proxy: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := ebpfProxy.Free(); err != nil {
|
||||
t.Errorf("failed to free ebpf proxy: %v", err)
|
||||
if err := loopbackProxy.Free(); err != nil {
|
||||
t.Errorf("failed to free loopback proxy: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
proxy := ebpf.NewProxyWrapper(ebpfProxy)
|
||||
proxy := loopback.NewProxyWrapper(loopbackProxy)
|
||||
|
||||
// NetBird UDP address of the remote peer
|
||||
nbAddr := &net.UDPAddr{
|
||||
@@ -259,17 +259,17 @@ func TestRedirectAs_eBPF_IPv6(t *testing.T) {
|
||||
// TestRedirectAs_Multiple_Switches tests switching between multiple endpoints
|
||||
func TestRedirectAs_Multiple_Switches(t *testing.T) {
|
||||
wgPort := 51856
|
||||
ebpfProxy := ebpf.NewWGEBPFProxy(wgPort, 1280)
|
||||
if err := ebpfProxy.Listen(); err != nil {
|
||||
t.Fatalf("failed to initialize ebpf proxy: %v", err)
|
||||
loopbackProxy := loopback.NewProxy(wgPort, 1280)
|
||||
if err := loopbackProxy.Listen(); err != nil {
|
||||
t.Fatalf("failed to initialize loopback proxy: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := ebpfProxy.Free(); err != nil {
|
||||
t.Errorf("failed to free ebpf proxy: %v", err)
|
||||
if err := loopbackProxy.Free(); err != nil {
|
||||
t.Errorf("failed to free loopback proxy: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
proxy := ebpf.NewProxyWrapper(ebpfProxy)
|
||||
proxy := loopback.NewProxyWrapper(loopbackProxy)
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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 := ®istryConfigurator{
|
||||
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.
|
||||
|
||||
@@ -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 := ®istryConfigurator{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 := ®istryConfigurator{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")
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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.
@@ -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.
@@ -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
|
||||
@@ -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";
|
||||
@@ -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__); \
|
||||
})
|
||||
```
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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.
|
||||
//
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
+55
@@ -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())
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
})
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user