mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-30 02:29:08 +02:00
Merge branch 'main' into loopback-wg-proxy
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
//
|
||||
|
||||
@@ -1040,7 +1040,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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
// Package wincmd locates the Windows utilities the client shells out to.
|
||||
package wincmd
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
// defaultSystem32Dir is where the system directory is on every supported
|
||||
// install, used only when the API that reports it fails.
|
||||
const defaultSystem32Dir = `C:\Windows\System32`
|
||||
|
||||
// System32 returns the full path of a Windows utility under the system
|
||||
// directory.
|
||||
//
|
||||
// PATH is deliberately not consulted. The daemon runs as LocalSystem with an
|
||||
// environment of its own, so whoever can place an entry in that PATH chooses
|
||||
// which binary runs with those privileges. The system directory is read from
|
||||
// the API rather than from %SystemRoot% for the same reason.
|
||||
func System32(command string) string {
|
||||
sysDir, err := windows.GetSystemDirectory()
|
||||
if err != nil {
|
||||
log.Warnf("Failed to locate the Windows system directory, falling back to %s: %v", defaultSystem32Dir, err)
|
||||
sysDir = defaultSystem32Dir
|
||||
}
|
||||
|
||||
return filepath.Join(sysDir, command+".exe")
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package wincmd
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSystem32IgnoresPATH(t *testing.T) {
|
||||
// A directory holding something that would win a PATH lookup, in front of
|
||||
// everything else: the daemon runs as LocalSystem, so a PATH entry must not
|
||||
// be able to decide what it executes.
|
||||
planted := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(planted, "netsh.exe"), []byte("not really netsh"), 0o600))
|
||||
t.Setenv("PATH", planted+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
|
||||
got := System32("netsh")
|
||||
|
||||
assert.True(t, filepath.IsAbs(got), "the path must be absolute, got %q", got)
|
||||
assert.NotContains(t, got, planted, "a PATH entry must not be consulted")
|
||||
assert.True(t, strings.EqualFold(filepath.Base(got), "netsh.exe"), "unexpected file name in %q", got)
|
||||
|
||||
// The system directory is what Windows reports it to be, not %SystemRoot%,
|
||||
// which the same caller could have set alongside PATH.
|
||||
t.Setenv("SystemRoot", planted)
|
||||
assert.Equal(t, got, System32("netsh"), "%SystemRoot% must not move the lookup")
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
import * as DropdownMenuPrimitive from "@radix-ui/react-dropdown-menu";
|
||||
import { cva } from "class-variance-authority";
|
||||
import { Check, ChevronRight, Circle } from "lucide-react";
|
||||
import { Check, ChevronRight } from "lucide-react";
|
||||
import * as React from "react";
|
||||
import { cn } from "@/lib/cn";
|
||||
|
||||
@@ -159,19 +159,23 @@ const DropdownMenuRadioItem = React.forwardRef<
|
||||
<DropdownMenuPrimitive.RadioItem
|
||||
ref={ref}
|
||||
className={cn(
|
||||
"relative flex cursor-default select-none items-center rounded-sm py-1.5 pl-8 pr-2 text-sm outline-none",
|
||||
"text-nb-gray-200 transition-colors hover:bg-nb-gray-900 hover:text-nb-gray-50 focus-visible:bg-nb-gray-900 focus-visible:text-nb-gray-50",
|
||||
"my-0.5 flex cursor-default select-none items-center gap-2 rounded-md px-2 py-2 outline-none",
|
||||
"text-xs font-semibold text-nb-gray-200 transition-colors",
|
||||
"data-[highlighted]:bg-nb-gray-850 data-[highlighted]:text-nb-gray-50",
|
||||
"data-[disabled]:pointer-events-none data-[disabled]:opacity-50",
|
||||
className,
|
||||
)}
|
||||
{...props}
|
||||
>
|
||||
<span className={"absolute left-2 flex h-3.5 w-3.5 items-center justify-center"}>
|
||||
{children}
|
||||
<span
|
||||
aria-hidden={"true"}
|
||||
className={"ml-auto flex w-4 shrink-0 items-center justify-center"}
|
||||
>
|
||||
<DropdownMenuPrimitive.ItemIndicator>
|
||||
<Circle className={"h-2 w-2 fill-current"} />
|
||||
<Check size={14} className={"text-netbird"} />
|
||||
</DropdownMenuPrimitive.ItemIndicator>
|
||||
</span>
|
||||
{children}
|
||||
</DropdownMenuPrimitive.RadioItem>
|
||||
));
|
||||
DropdownMenuRadioItem.displayName = DropdownMenuPrimitive.RadioItem.displayName;
|
||||
|
||||
@@ -89,7 +89,11 @@ export function LanguagePicker() {
|
||||
tabIndex={0}
|
||||
disabled={busy || languages.length === 0}
|
||||
onKeyDown={handleTriggerKeyDown}
|
||||
aria-label={t("settings.general.language.label")}
|
||||
aria-label={
|
||||
current
|
||||
? `${t("settings.general.language.label")}: ${labelFor(current)}`
|
||||
: t("settings.general.language.label")
|
||||
}
|
||||
aria-haspopup={"listbox"}
|
||||
aria-expanded={open}
|
||||
className={cn(
|
||||
|
||||
@@ -1,18 +1,10 @@
|
||||
import { useState } from "react";
|
||||
import { useTranslation } from "react-i18next";
|
||||
import { ChevronDown, MonitorIcon, MoonIcon, SunMediumIcon, type LucideIcon } from "lucide-react";
|
||||
import {
|
||||
DropdownMenu,
|
||||
DropdownMenuContent,
|
||||
DropdownMenuRadioGroup,
|
||||
DropdownMenuRadioItem,
|
||||
DropdownMenuTrigger,
|
||||
} from "@/components/DropdownMenu";
|
||||
import { MonitorIcon, MoonIcon, SunMediumIcon, type LucideIcon } from "lucide-react";
|
||||
import { Select } from "@/components/inputs/Select";
|
||||
import { HelpText } from "@/components/typography/HelpText";
|
||||
import { Label } from "@/components/typography/Label";
|
||||
import { useTheme, type ThemePreference } from "@/contexts/ThemeContext";
|
||||
import { useFocusVisible } from "@/hooks/useFocusVisible";
|
||||
import { cn } from "@/lib/cn";
|
||||
import { errorDialog, formatErrorMessage } from "@/lib/errors";
|
||||
|
||||
const OPTIONS: { value: ThemePreference; icon: LucideIcon; labelKey: string }[] = [
|
||||
@@ -25,16 +17,12 @@ export function ThemePicker() {
|
||||
const { t } = useTranslation();
|
||||
const { theme, setTheme } = useTheme();
|
||||
const [busy, setBusy] = useState(false);
|
||||
const isFocusVisible = useFocusVisible();
|
||||
|
||||
const current = OPTIONS.find((o) => o.value === theme) ?? OPTIONS[0];
|
||||
const CurrentIcon = current.icon;
|
||||
|
||||
const select = async (value: string) => {
|
||||
const select = async (value: ThemePreference) => {
|
||||
if (busy || value === theme) return;
|
||||
setBusy(true);
|
||||
try {
|
||||
await setTheme(value as ThemePreference);
|
||||
await setTheme(value);
|
||||
} catch (e) {
|
||||
await errorDialog({
|
||||
Title: t("settings.error.saveTitle"),
|
||||
@@ -52,57 +40,17 @@ export function ThemePicker() {
|
||||
<HelpText margin={false}>{t("settings.general.theme.help")}</HelpText>
|
||||
</div>
|
||||
<div className={"shrink-0"}>
|
||||
<DropdownMenu>
|
||||
<DropdownMenuTrigger asChild>
|
||||
<button
|
||||
type={"button"}
|
||||
tabIndex={0}
|
||||
disabled={busy}
|
||||
aria-label={t("settings.general.theme.label")}
|
||||
className={cn(
|
||||
"inline-flex h-[40px] min-w-[160px] items-center gap-2 px-3",
|
||||
"rounded-md border bg-white dark:bg-nb-gray-900",
|
||||
"border-neutral-200 dark:border-nb-gray-700",
|
||||
"cursor-default text-xs font-semibold text-nb-gray-100 outline-none",
|
||||
"hover:border-nb-gray-700 data-[state=open]:border-nb-gray-700 dark:hover:border-nb-gray-600 dark:data-[state=open]:border-nb-gray-600",
|
||||
isFocusVisible &&
|
||||
"focus-visible:ring-2 focus-visible:ring-nb-gray-50/60 focus-visible:ring-offset-2 focus-visible:ring-offset-nb-gray-940",
|
||||
"disabled:opacity-50",
|
||||
)}
|
||||
>
|
||||
<CurrentIcon
|
||||
size={16}
|
||||
aria-hidden={"true"}
|
||||
className={"shrink-0 text-nb-gray-200"}
|
||||
/>
|
||||
<span className={"flex-1 truncate text-left"}>
|
||||
{t(current.labelKey)}
|
||||
</span>
|
||||
<ChevronDown
|
||||
size={12}
|
||||
aria-hidden={"true"}
|
||||
className={"shrink-0 text-nb-gray-400"}
|
||||
/>
|
||||
</button>
|
||||
</DropdownMenuTrigger>
|
||||
<DropdownMenuContent
|
||||
align={"end"}
|
||||
className={"w-[var(--radix-dropdown-menu-trigger-width)]"}
|
||||
>
|
||||
<DropdownMenuRadioGroup value={theme} onValueChange={(v) => void select(v)}>
|
||||
{OPTIONS.map(({ value, icon: Icon, labelKey }) => (
|
||||
<DropdownMenuRadioItem key={value} value={value}>
|
||||
<Icon
|
||||
size={14}
|
||||
aria-hidden={"true"}
|
||||
className={"mr-2 shrink-0 text-nb-gray-300"}
|
||||
/>
|
||||
{t(labelKey)}
|
||||
</DropdownMenuRadioItem>
|
||||
))}
|
||||
</DropdownMenuRadioGroup>
|
||||
</DropdownMenuContent>
|
||||
</DropdownMenu>
|
||||
<Select
|
||||
value={theme}
|
||||
options={OPTIONS.map(({ value, icon, labelKey }) => ({
|
||||
value,
|
||||
icon,
|
||||
label: t(labelKey),
|
||||
}))}
|
||||
onChange={(v) => void select(v)}
|
||||
ariaLabel={t("settings.general.theme.label")}
|
||||
disabled={busy}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
|
||||
@@ -147,7 +147,8 @@ export const Button = forwardRef<HTMLButtonElement, ButtonProps>(function Button
|
||||
ref={ref}
|
||||
type={type}
|
||||
tabIndex={0}
|
||||
disabled={disabled || loading}
|
||||
disabled={disabled}
|
||||
aria-disabled={loading || undefined}
|
||||
aria-busy={loading || undefined}
|
||||
className={cn(
|
||||
buttonVariants({
|
||||
@@ -156,10 +157,15 @@ export const Button = forwardRef<HTMLButtonElement, ButtonProps>(function Button
|
||||
border: border ? 1 : 0,
|
||||
size,
|
||||
}),
|
||||
loading && "pointer-events-none",
|
||||
className,
|
||||
)}
|
||||
onClick={(e) => {
|
||||
if (stopPropagation) e.stopPropagation();
|
||||
if (loading) {
|
||||
e.preventDefault();
|
||||
return;
|
||||
}
|
||||
if (copy !== undefined) {
|
||||
void navigator.clipboard
|
||||
.writeText(copy)
|
||||
|
||||
@@ -14,6 +14,7 @@ type ConfirmModalProps = {
|
||||
cancelLabel?: string;
|
||||
danger?: boolean;
|
||||
busy?: boolean;
|
||||
cancellable?: boolean;
|
||||
onConfirm: () => void;
|
||||
onCancel: () => void;
|
||||
};
|
||||
@@ -26,11 +27,13 @@ export const ConfirmModal = ({
|
||||
cancelLabel,
|
||||
danger = false,
|
||||
busy = false,
|
||||
cancellable,
|
||||
onConfirm,
|
||||
onCancel,
|
||||
}: ConfirmModalProps) => {
|
||||
const { t } = useTranslation();
|
||||
const resolvedCancel = cancelLabel ?? t("common.cancel");
|
||||
const canCancel = cancellable ?? !busy;
|
||||
|
||||
const srTitle = typeof title === "string" ? title : undefined;
|
||||
const srDescription = typeof description === "string" ? description : undefined;
|
||||
@@ -39,7 +42,7 @@ export const ConfirmModal = ({
|
||||
<Dialog.Root
|
||||
open={open}
|
||||
onOpenChange={(next) => {
|
||||
if (!next && !busy) onCancel();
|
||||
if (!next && canCancel) onCancel();
|
||||
}}
|
||||
>
|
||||
<Dialog.Content
|
||||
@@ -62,7 +65,7 @@ export const ConfirmModal = ({
|
||||
<Button
|
||||
variant={"secondary"}
|
||||
size={"sm"}
|
||||
disabled={busy}
|
||||
disabled={!canCancel}
|
||||
onClick={onCancel}
|
||||
>
|
||||
{resolvedCancel}
|
||||
@@ -71,7 +74,7 @@ export const ConfirmModal = ({
|
||||
autoFocus
|
||||
variant={danger ? "danger" : "primary"}
|
||||
size={"sm"}
|
||||
disabled={busy}
|
||||
loading={busy}
|
||||
onClick={onConfirm}
|
||||
>
|
||||
{confirmLabel}
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
import type { LucideIcon } from "lucide-react";
|
||||
import { ChevronDown } from "lucide-react";
|
||||
import {
|
||||
DropdownMenu,
|
||||
DropdownMenuContent,
|
||||
DropdownMenuRadioGroup,
|
||||
DropdownMenuRadioItem,
|
||||
DropdownMenuTrigger,
|
||||
} from "@/components/DropdownMenu";
|
||||
import { useFocusVisible } from "@/hooks/useFocusVisible";
|
||||
import { cn } from "@/lib/cn";
|
||||
|
||||
export type SelectOption<T extends string> = {
|
||||
value: T;
|
||||
label: string;
|
||||
icon?: LucideIcon;
|
||||
};
|
||||
|
||||
type SelectProps<T extends string> = {
|
||||
value: T;
|
||||
options: SelectOption<T>[];
|
||||
onChange: (value: T) => void;
|
||||
ariaLabel: string;
|
||||
disabled?: boolean;
|
||||
className?: string;
|
||||
};
|
||||
|
||||
export function Select<T extends string>({
|
||||
value,
|
||||
options,
|
||||
onChange,
|
||||
ariaLabel,
|
||||
disabled,
|
||||
className,
|
||||
}: SelectProps<T>) {
|
||||
const isFocusVisible = useFocusVisible();
|
||||
const current = options.find((o) => o.value === value) ?? options[0];
|
||||
const CurrentIcon = current?.icon;
|
||||
|
||||
return (
|
||||
<DropdownMenu>
|
||||
<DropdownMenuTrigger asChild>
|
||||
<button
|
||||
type={"button"}
|
||||
tabIndex={0}
|
||||
disabled={disabled}
|
||||
aria-label={current ? `${ariaLabel}: ${current.label}` : ariaLabel}
|
||||
className={cn(
|
||||
"inline-flex h-[40px] min-w-[160px] items-center gap-2 px-3",
|
||||
"rounded-md border bg-white dark:bg-nb-gray-900",
|
||||
"border-neutral-200 dark:border-nb-gray-700",
|
||||
"cursor-default text-xs font-semibold text-nb-gray-100 outline-none",
|
||||
"hover:border-nb-gray-700 data-[state=open]:border-nb-gray-700 dark:hover:border-nb-gray-600 dark:data-[state=open]:border-nb-gray-600",
|
||||
isFocusVisible &&
|
||||
"focus-visible:ring-2 focus-visible:ring-nb-gray-50/60 focus-visible:ring-offset-2 focus-visible:ring-offset-nb-gray-940",
|
||||
"disabled:opacity-50",
|
||||
className,
|
||||
)}
|
||||
>
|
||||
{CurrentIcon && (
|
||||
<CurrentIcon
|
||||
size={16}
|
||||
aria-hidden={"true"}
|
||||
className={"shrink-0 text-nb-gray-200"}
|
||||
/>
|
||||
)}
|
||||
<span className={"flex-1 truncate text-left"}>{current?.label ?? "—"}</span>
|
||||
<ChevronDown
|
||||
size={12}
|
||||
aria-hidden={"true"}
|
||||
className={"shrink-0 text-nb-gray-400"}
|
||||
/>
|
||||
</button>
|
||||
</DropdownMenuTrigger>
|
||||
<DropdownMenuContent
|
||||
align={"start"}
|
||||
sideOffset={6}
|
||||
className={
|
||||
"w-[var(--radix-dropdown-menu-trigger-width)] border-nb-gray-850 bg-nb-gray-920"
|
||||
}
|
||||
>
|
||||
<DropdownMenuRadioGroup value={value} onValueChange={(v) => onChange(v as T)}>
|
||||
{options.map(({ value: optionValue, label, icon: Icon }) => (
|
||||
<DropdownMenuRadioItem key={optionValue} value={optionValue}>
|
||||
{Icon && (
|
||||
<Icon
|
||||
size={14}
|
||||
aria-hidden={"true"}
|
||||
className={"shrink-0 text-nb-gray-300"}
|
||||
/>
|
||||
)}
|
||||
<span className={"min-w-0 flex-1 truncate"}>{label}</span>
|
||||
</DropdownMenuRadioItem>
|
||||
))}
|
||||
</DropdownMenuRadioGroup>
|
||||
</DropdownMenuContent>
|
||||
</DropdownMenu>
|
||||
);
|
||||
}
|
||||
@@ -8,6 +8,28 @@ import {
|
||||
useState,
|
||||
} from "react";
|
||||
import { ConfirmModal } from "@/components/dialog/ConfirmModal";
|
||||
import i18next from "@/lib/i18n";
|
||||
|
||||
// Nothing on the daemon path carries a deadline, so a hung call would leave the
|
||||
// modal spinning with no way out. Cancel comes back once the wait stops looking
|
||||
// normal, and the wait is abandoned entirely at the deadline.
|
||||
const CANCELLABLE_AFTER_MS = 15_000;
|
||||
const TIMEOUT_MS = 30_000;
|
||||
|
||||
const withTimeout = async (action: () => Promise<unknown>) => {
|
||||
let timer: ReturnType<typeof setTimeout> | undefined;
|
||||
const expiry = new Promise<never>((_, reject) => {
|
||||
timer = setTimeout(
|
||||
() => reject(new Error(i18next.t("error.daemon_unreachable"))),
|
||||
TIMEOUT_MS,
|
||||
);
|
||||
});
|
||||
try {
|
||||
await Promise.race([action(), expiry]);
|
||||
} finally {
|
||||
clearTimeout(timer);
|
||||
}
|
||||
};
|
||||
|
||||
export type ConfirmOptions = {
|
||||
title: ReactNode;
|
||||
@@ -15,6 +37,7 @@ export type ConfirmOptions = {
|
||||
confirmLabel: string;
|
||||
cancelLabel?: string;
|
||||
danger?: boolean;
|
||||
onConfirm?: () => Promise<unknown>;
|
||||
};
|
||||
|
||||
type DialogContextValue = {
|
||||
@@ -23,23 +46,50 @@ type DialogContextValue = {
|
||||
|
||||
const DialogContext = createContext<DialogContextValue | null>(null);
|
||||
|
||||
type Settler = { resolve: (result: boolean) => void; reject: (reason: unknown) => void };
|
||||
|
||||
export function DialogProvider({ children }: Readonly<{ children: ReactNode }>) {
|
||||
const [open, setOpen] = useState(false);
|
||||
const [busy, setBusy] = useState(false);
|
||||
const [stalled, setStalled] = useState(false);
|
||||
const [options, setOptions] = useState<ConfirmOptions | null>(null);
|
||||
const resolverRef = useRef<((result: boolean) => void) | null>(null);
|
||||
const resolverRef = useRef<Settler | null>(null);
|
||||
|
||||
const confirm = useCallback((opts: ConfirmOptions) => {
|
||||
setOptions(opts);
|
||||
setOpen(true);
|
||||
return new Promise<boolean>((resolve) => {
|
||||
resolverRef.current = resolve;
|
||||
return new Promise<boolean>((resolve, reject) => {
|
||||
resolverRef.current = { resolve, reject };
|
||||
});
|
||||
}, []);
|
||||
|
||||
const settle = (result: boolean) => {
|
||||
resolverRef.current?.(result);
|
||||
const take = (expected?: Settler | null) => {
|
||||
const settler = resolverRef.current;
|
||||
if (expected && settler !== expected) return null;
|
||||
resolverRef.current = null;
|
||||
setBusy(false);
|
||||
setStalled(false);
|
||||
setOpen(false);
|
||||
return settler;
|
||||
};
|
||||
|
||||
const handleConfirm = async () => {
|
||||
const action = options?.onConfirm;
|
||||
if (!action) {
|
||||
take()?.resolve(true);
|
||||
return;
|
||||
}
|
||||
const dispatched = resolverRef.current;
|
||||
setBusy(true);
|
||||
const stallTimer = setTimeout(() => setStalled(true), CANCELLABLE_AFTER_MS);
|
||||
try {
|
||||
await withTimeout(action);
|
||||
take(dispatched)?.resolve(true);
|
||||
} catch (e) {
|
||||
take(dispatched)?.reject(e);
|
||||
} finally {
|
||||
clearTimeout(stallTimer);
|
||||
}
|
||||
};
|
||||
|
||||
const value = useMemo<DialogContextValue>(() => ({ confirm }), [confirm]);
|
||||
@@ -54,8 +104,10 @@ export function DialogProvider({ children }: Readonly<{ children: ReactNode }>)
|
||||
confirmLabel={options?.confirmLabel ?? ""}
|
||||
cancelLabel={options?.cancelLabel}
|
||||
danger={options?.danger}
|
||||
onConfirm={() => settle(true)}
|
||||
onCancel={() => settle(false)}
|
||||
busy={busy}
|
||||
cancellable={!busy || stalled}
|
||||
onConfirm={() => void handleConfirm()}
|
||||
onCancel={() => take()?.resolve(false)}
|
||||
/>
|
||||
</DialogContext.Provider>
|
||||
);
|
||||
|
||||
@@ -78,7 +78,7 @@ export function ProfilesTab() {
|
||||
return items;
|
||||
}, [profiles, activeProfileId]);
|
||||
|
||||
const guarded = async (title: string, fn: () => Promise<void>) => {
|
||||
const guarded = async (title: string, fn: () => Promise<unknown>) => {
|
||||
if (busy) return;
|
||||
setBusy(true);
|
||||
try {
|
||||
@@ -115,14 +115,15 @@ export function ProfilesTab() {
|
||||
|
||||
const handleDelete = async (id: string, name: string) => {
|
||||
if (id === DEFAULT_PROFILE_ID) return;
|
||||
const ok = await confirm({
|
||||
title: t("profile.delete.title", { name }),
|
||||
description: t("profile.delete.message", { name }),
|
||||
confirmLabel: t("common.delete"),
|
||||
danger: true,
|
||||
});
|
||||
if (!ok) return;
|
||||
void guarded(i18next.t("profile.error.deleteTitle"), () => removeProfile(id));
|
||||
await guarded(i18next.t("profile.error.deleteTitle"), () =>
|
||||
confirm({
|
||||
title: t("profile.delete.title", { name }),
|
||||
description: t("profile.delete.message", { name }),
|
||||
confirmLabel: t("common.delete"),
|
||||
danger: true,
|
||||
onConfirm: () => removeProfile(id),
|
||||
}),
|
||||
);
|
||||
};
|
||||
|
||||
const handleCreate = async (name: string, managementUrl: string) => {
|
||||
|
||||
@@ -1,5 +1,21 @@
|
||||
import { useCallback, useEffect, useRef, useState } from "react";
|
||||
import { createRoot } from "react-dom/client";
|
||||
import netbirdLogo from "@/assets/logos/netbird.svg";
|
||||
|
||||
const scratch = new Uint32Array(1);
|
||||
|
||||
function random() {
|
||||
crypto.getRandomValues(scratch);
|
||||
return scratch[0] / 2 ** 32;
|
||||
}
|
||||
|
||||
type Mask = {
|
||||
cols: number;
|
||||
rows: number;
|
||||
cells: Uint8Array;
|
||||
seeds: Uint8Array;
|
||||
glow: Float32Array;
|
||||
};
|
||||
|
||||
export function useAccentTrigger() {
|
||||
const clicksRef = useRef(0);
|
||||
@@ -50,24 +66,45 @@ function Accent({ onDone }: Readonly<{ onDone: () => void }>) {
|
||||
const ctx = canvas.getContext("2d");
|
||||
if (!ctx) return;
|
||||
|
||||
const chars = "DRIBTENMAET".split("").reverse().join("");
|
||||
|
||||
let disposed = false;
|
||||
let mask: Mask | null = null;
|
||||
|
||||
const dpr = window.devicePixelRatio || 1;
|
||||
let columns = 0;
|
||||
let drops: number[] = [];
|
||||
let latestBuild = 0;
|
||||
let rebuild: ReturnType<typeof setTimeout> | undefined;
|
||||
|
||||
const resize = () => {
|
||||
canvas.width = window.innerWidth * dpr;
|
||||
canvas.height = window.innerHeight * dpr;
|
||||
canvas.style.width = `${window.innerWidth}px`;
|
||||
canvas.style.height = `${window.innerHeight}px`;
|
||||
ctx.setTransform(dpr, 0, 0, dpr, 0, 0);
|
||||
|
||||
const next = Math.floor(window.innerWidth / 15);
|
||||
if (next !== columns) {
|
||||
columns = next;
|
||||
drops = Array.from({ length: columns }, () => random() * -60);
|
||||
mask = null;
|
||||
}
|
||||
|
||||
globalThis.clearTimeout(rebuild);
|
||||
rebuild = globalThis.setTimeout(() => {
|
||||
const build = ++latestBuild;
|
||||
void buildMask().then((m) => {
|
||||
if (!disposed && build === latestBuild) mask = m;
|
||||
});
|
||||
}, 100);
|
||||
};
|
||||
resize();
|
||||
window.addEventListener("resize", resize);
|
||||
|
||||
const chars = "TEAMNETBIRD";
|
||||
const fontSize = 16;
|
||||
const columns = Math.floor(window.innerWidth / fontSize);
|
||||
const drops = Array.from({ length: columns }, () => Math.random() * -50);
|
||||
|
||||
let raf = 0;
|
||||
let last = 0;
|
||||
let frame = 0;
|
||||
const draw = (t: number) => {
|
||||
if (t - last > 50) {
|
||||
last = t;
|
||||
@@ -77,18 +114,26 @@ function Accent({ onDone }: Readonly<{ onDone: () => void }>) {
|
||||
ctx.fillRect(0, 0, window.innerWidth, window.innerHeight);
|
||||
|
||||
ctx.globalCompositeOperation = "source-over";
|
||||
ctx.font = `${fontSize}px ui-monospace, monospace`;
|
||||
ctx.fillStyle = "#f68330";
|
||||
ctx.font = "15px ui-monospace, monospace";
|
||||
ctx.textBaseline = "top";
|
||||
|
||||
ctx.shadowBlur = 0;
|
||||
ctx.fillStyle = "rgba(246, 131, 48, 0.5)";
|
||||
for (let i = 0; i < drops.length; i++) {
|
||||
const ch = chars[Math.floor(Math.random() * chars.length)];
|
||||
const y = drops[i] * fontSize;
|
||||
ctx.fillText(ch, i * fontSize, y);
|
||||
if (y > window.innerHeight && Math.random() > 0.975) {
|
||||
drops[i] = 0;
|
||||
const ch = chars[Math.floor(random() * chars.length)];
|
||||
const y = drops[i] * 15;
|
||||
ctx.fillText(ch, i * 15, y);
|
||||
|
||||
igniteTrail(mask, i, Math.floor(drops[i]));
|
||||
|
||||
if (y > window.innerHeight && random() > 0.86) {
|
||||
drops[i] = random() * -12;
|
||||
}
|
||||
drops[i]++;
|
||||
}
|
||||
|
||||
drawGlow(ctx, mask, frame, chars);
|
||||
frame++;
|
||||
}
|
||||
raf = requestAnimationFrame(draw);
|
||||
};
|
||||
@@ -100,8 +145,10 @@ function Accent({ onDone }: Readonly<{ onDone: () => void }>) {
|
||||
}, 9000);
|
||||
|
||||
return () => {
|
||||
disposed = true;
|
||||
cancelAnimationFrame(raf);
|
||||
globalThis.clearTimeout(timeout);
|
||||
globalThis.clearTimeout(rebuild);
|
||||
window.removeEventListener("resize", resize);
|
||||
};
|
||||
}, [onDone]);
|
||||
@@ -114,3 +161,103 @@ function Accent({ onDone }: Readonly<{ onDone: () => void }>) {
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
function igniteTrail(mask: Mask | null, col: number, row: number) {
|
||||
if (!mask || col < 0 || col >= mask.cols) return;
|
||||
for (let k = 0; k < 6; k++) {
|
||||
const r = row - k;
|
||||
if (r < 0 || r >= mask.rows) continue;
|
||||
const idx = r * mask.cols + col;
|
||||
if (mask.cells[idx] === 0) continue;
|
||||
const heat = 1 - k / 6;
|
||||
if (heat > mask.glow[idx]) mask.glow[idx] = heat;
|
||||
}
|
||||
}
|
||||
|
||||
function drawGlow(ctx: CanvasRenderingContext2D, mask: Mask | null, frame: number, chars: string) {
|
||||
if (!mask) return;
|
||||
|
||||
for (let idx = 0; idx < mask.cells.length; idx++) {
|
||||
const heat = fade(mask, idx);
|
||||
if (heat === 0) continue;
|
||||
|
||||
const seed = mask.seeds[idx];
|
||||
const core = mask.cells[idx] === 2;
|
||||
|
||||
ctx.shadowColor = core ? "#f05252" : "#f68330";
|
||||
ctx.shadowBlur = 10 * heat;
|
||||
ctx.fillStyle = core ? `rgba(255, 226, 210, ${heat})` : `rgba(255, 255, 255, ${heat})`;
|
||||
ctx.fillText(
|
||||
chars[(seed + Math.floor(frame / (3 + (seed % 5)))) % chars.length],
|
||||
(idx % mask.cols) * 15,
|
||||
Math.floor(idx / mask.cols) * 15,
|
||||
);
|
||||
}
|
||||
ctx.shadowBlur = 0;
|
||||
}
|
||||
|
||||
function fade(mask: Mask, idx: number) {
|
||||
if (mask.cells[idx] === 0) return 0;
|
||||
|
||||
const heat = mask.glow[idx];
|
||||
if (heat <= 0.02) {
|
||||
mask.glow[idx] = 0;
|
||||
return 0;
|
||||
}
|
||||
mask.glow[idx] = heat * 0.94;
|
||||
return heat;
|
||||
}
|
||||
|
||||
function loadLogo() {
|
||||
return new Promise<HTMLImageElement>((resolve, reject) => {
|
||||
const img = new Image();
|
||||
img.onload = () => resolve(img);
|
||||
img.onerror = reject;
|
||||
img.src = netbirdLogo;
|
||||
});
|
||||
}
|
||||
|
||||
async function buildMask(): Promise<Mask | null> {
|
||||
const cols = Math.floor(window.innerWidth / 15);
|
||||
const rows = Math.ceil(window.innerHeight / 15);
|
||||
if (cols <= 0 || rows <= 0) return null;
|
||||
|
||||
let img: HTMLImageElement;
|
||||
try {
|
||||
img = await loadLogo();
|
||||
} catch {
|
||||
return null;
|
||||
}
|
||||
|
||||
const off = document.createElement("canvas");
|
||||
off.width = cols;
|
||||
off.height = rows;
|
||||
const offCtx = off.getContext("2d", { willReadFrequently: true });
|
||||
if (!offCtx) return null;
|
||||
|
||||
const aspect = (img.naturalWidth || 31) / (img.naturalHeight || 23);
|
||||
let w = cols * 0.8;
|
||||
let h = w / aspect;
|
||||
if (h > rows * 0.8) {
|
||||
h = rows * 0.8;
|
||||
w = h * aspect;
|
||||
}
|
||||
|
||||
offCtx.imageSmoothingEnabled = false;
|
||||
offCtx.drawImage(img, (cols - w) / 2, (rows - h) / 2, w, h);
|
||||
|
||||
const { data } = offCtx.getImageData(0, 0, cols, rows);
|
||||
const cells = new Uint8Array(cols * rows);
|
||||
const seeds = new Uint8Array(cols * rows);
|
||||
const glow = new Float32Array(cols * rows);
|
||||
for (let i = 0; i < cells.length; i++) {
|
||||
seeds[i] = Math.floor(random() * 251);
|
||||
const alpha = data[i * 4 + 3];
|
||||
if (alpha < 64) continue;
|
||||
const r = data[i * 4];
|
||||
const g = data[i * 4 + 1];
|
||||
const b = data[i * 4 + 2];
|
||||
cells[i] = r > 180 && g < 130 && b < 130 && g <= b + 24 ? 2 : 1;
|
||||
}
|
||||
return { cols, rows, cells, seeds, glow };
|
||||
}
|
||||
|
||||
@@ -1,6 +1,15 @@
|
||||
import { useId, type ReactNode } from "react";
|
||||
import { Trans, useTranslation } from "react-i18next";
|
||||
import { ChevronDown, CircleCheckBig, FolderOpen, Info, Loader2 } from "lucide-react";
|
||||
import {
|
||||
CircleCheckBig,
|
||||
FolderOpen,
|
||||
Info,
|
||||
Loader2,
|
||||
Shield,
|
||||
ShieldCheck,
|
||||
ShieldOff,
|
||||
type LucideIcon,
|
||||
} from "lucide-react";
|
||||
import { Browser } from "@wailsio/runtime";
|
||||
import { Debug as DebugSvc } from "@bindings/services";
|
||||
import type { DebugBundleResult } from "@bindings/services/models.js";
|
||||
@@ -8,20 +17,13 @@ import { Button } from "@/components/buttons/Button";
|
||||
import { DialogActions } from "@/components/dialog/DialogActions";
|
||||
import { DialogDescription } from "@/components/dialog/DialogDescription";
|
||||
import { DialogHeading } from "@/components/dialog/DialogHeading";
|
||||
import {
|
||||
DropdownMenu,
|
||||
DropdownMenuContent,
|
||||
DropdownMenuRadioGroup,
|
||||
DropdownMenuRadioItem,
|
||||
DropdownMenuTrigger,
|
||||
} from "@/components/DropdownMenu";
|
||||
import FancyToggleSwitch from "@/components/switches/FancyToggleSwitch";
|
||||
import HelpText from "@/components/typography/HelpText.tsx";
|
||||
import { Input } from "@/components/inputs/Input";
|
||||
import { Label } from "@/components/typography/Label";
|
||||
import { Select } from "@/components/inputs/Select";
|
||||
import { SquareIcon } from "@/components/SquareIcon";
|
||||
import { Tooltip } from "@/components/Tooltip";
|
||||
import { cn } from "@/lib/cn";
|
||||
import { formatRemaining } from "@/lib/formatters";
|
||||
import type { AnonymizeLevel, DebugStage } from "@/contexts/DebugBundleContext";
|
||||
import { useDebugBundleContext } from "@/contexts/DebugBundleContext";
|
||||
@@ -29,6 +31,12 @@ import { SectionGroup, SettingsBottomBar } from "@/modules/settings/SettingsSect
|
||||
|
||||
const SUPPORT_DOCS_URL = "https://docs.netbird.io/help/report-bug-issues";
|
||||
|
||||
const ANONYMIZE_LEVELS: { value: AnonymizeLevel; icon: LucideIcon }[] = [
|
||||
{ value: "none", icon: ShieldOff },
|
||||
{ value: "default", icon: Shield },
|
||||
{ value: "strict", icon: ShieldCheck },
|
||||
];
|
||||
|
||||
export function SettingsTroubleshooting() {
|
||||
const { t } = useTranslation();
|
||||
const durationId = useId();
|
||||
@@ -89,44 +97,16 @@ export function SettingsTroubleshooting() {
|
||||
</HelpText>
|
||||
</div>
|
||||
<div className={"shrink-0"}>
|
||||
<DropdownMenu>
|
||||
<DropdownMenuTrigger asChild>
|
||||
<button
|
||||
type={"button"}
|
||||
aria-label={t("settings.troubleshooting.anonymize.label")}
|
||||
className={cn(
|
||||
"inline-flex h-[40px] min-w-[160px] items-center justify-between gap-2 px-3",
|
||||
"rounded-md border bg-white dark:bg-nb-gray-900",
|
||||
"border-neutral-200 dark:border-nb-gray-700",
|
||||
"cursor-default text-xs font-semibold text-nb-gray-100 outline-none",
|
||||
"hover:border-nb-gray-700 data-[state=open]:border-nb-gray-700 dark:hover:border-nb-gray-600 dark:data-[state=open]:border-nb-gray-600",
|
||||
)}
|
||||
>
|
||||
{t(`settings.troubleshooting.anonymize.${anonymizeLevel}`)}
|
||||
<ChevronDown
|
||||
size={16}
|
||||
aria-hidden={"true"}
|
||||
className={"shrink-0 text-nb-gray-200"}
|
||||
/>
|
||||
</button>
|
||||
</DropdownMenuTrigger>
|
||||
<DropdownMenuContent align={"end"} className={"min-w-[160px]"}>
|
||||
<DropdownMenuRadioGroup
|
||||
value={anonymizeLevel}
|
||||
onValueChange={(v) => setAnonymizeLevel(v as AnonymizeLevel)}
|
||||
>
|
||||
<DropdownMenuRadioItem value={"none"}>
|
||||
{t("settings.troubleshooting.anonymize.none")}
|
||||
</DropdownMenuRadioItem>
|
||||
<DropdownMenuRadioItem value={"default"}>
|
||||
{t("settings.troubleshooting.anonymize.default")}
|
||||
</DropdownMenuRadioItem>
|
||||
<DropdownMenuRadioItem value={"strict"}>
|
||||
{t("settings.troubleshooting.anonymize.strict")}
|
||||
</DropdownMenuRadioItem>
|
||||
</DropdownMenuRadioGroup>
|
||||
</DropdownMenuContent>
|
||||
</DropdownMenu>
|
||||
<Select
|
||||
value={anonymizeLevel}
|
||||
options={ANONYMIZE_LEVELS.map(({ value, icon }) => ({
|
||||
value,
|
||||
icon,
|
||||
label: t(`settings.troubleshooting.anonymize.${value}`),
|
||||
}))}
|
||||
onChange={setAnonymizeLevel}
|
||||
ariaLabel={t("settings.troubleshooting.anonymize.label")}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
<FancyToggleSwitch
|
||||
|
||||
Reference in New Issue
Block a user