Compare commits

...

2 Commits

Author SHA1 Message Date
riccardom
afb0525db3 Remove obvious comments 2026-08-06 13:39:48 +02:00
riccardom
ad0033f851 [client] Add catch-all NRPT rule when NetBird is the primary DNS resolver 2026-08-05 19:05:18 +02:00
2 changed files with 162 additions and 0 deletions

View File

@@ -6,8 +6,10 @@ import (
"fmt"
"io"
"net/netip"
"os"
"os/exec"
"slices"
"strconv"
"strings"
"syscall"
"time"
@@ -36,6 +38,15 @@ const (
gpoDnsPolicyRoot = `SOFTWARE\Policies\Microsoft\Windows NT\DNSClient\DnsPolicyConfig`
gpoDnsPolicyConfigMatchPath = gpoDnsPolicyRoot + `\NetBird-Match`
dnsPolicyConfigCatchAllPath = `SYSTEM\CurrentControlSet\Services\Dnscache\Parameters\DnsPolicyConfig\NetBird-CatchAll`
gpoDnsPolicyConfigCatchAllPath = gpoDnsPolicyRoot + `\NetBird-CatchAll`
nrptCatchAllNamespace = "."
// envDisableCatchAllNRPT turns off the catch-all NRPT rule, restoring the
// previous behavior where the OS is free to query other adapters' resolvers.
envDisableCatchAllNRPT = "NB_DISABLE_DNS_CATCHALL_NRPT"
dnsPolicyConfigVersionKey = "Version"
dnsPolicyConfigVersionValue = 2
dnsPolicyConfigNameKey = "Name"
@@ -318,6 +329,12 @@ func (r *registryConfigurator) applyDNSConfig(config HostDNSConfig, stateManager
r.updateState(stateManager)
if config.RouteAll {
if err := r.addDNSCatchAllPolicy(config.ServerIP); err != nil {
return fmt.Errorf("add dns catch-all policy: %w", err)
}
}
if err := r.updateSearchDomains(searchDomains); err != nil {
return fmt.Errorf("update search domains: %w", err)
}
@@ -388,6 +405,29 @@ func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr
return ruleIndex, nil
}
func (r *registryConfigurator) addDNSCatchAllPolicy(ip netip.Addr) error {
if parseBoolEnv(envDisableCatchAllNRPT) {
log.Infof("%s is set, not forcing all DNS queries through %s", envDisableCatchAllNRPT, ip)
return nil
}
if err := r.configureDNSPolicy(dnsPolicyConfigCatchAllPath, []string{nrptCatchAllNamespace}, ip); err != nil {
return fmt.Errorf("configure catch-all DNS policy: %w", err)
}
if r.gpo {
if err := r.configureDNSPolicy(gpoDnsPolicyConfigCatchAllPath, []string{nrptCatchAllNamespace}, ip); err != nil {
return fmt.Errorf("configure gpo catch-all DNS policy: %w", err)
}
if err := refreshGroupPolicy(); err != nil {
log.Warnf("failed to refresh group policy: %v", err)
}
}
log.Infof("added catch-all NRPT rule: all DNS queries now resolve exclusively through %s", ip)
return nil
}
func (r *registryConfigurator) configureDNSPolicy(policyPath string, domains []string, ip netip.Addr) error {
if err := removeRegistryKeyFromDNSPolicyConfig(policyPath); err != nil {
return fmt.Errorf("remove existing dns policy: %w", err)
@@ -530,6 +570,14 @@ func (r *registryConfigurator) removeDNSMatchPolicies() error {
merr = multierror.Append(merr, fmt.Errorf("remove GPO base entry: %w", err))
}
if err := removeRegistryKeyFromDNSPolicyConfig(dnsPolicyConfigCatchAllPath); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove local catch-all entry: %w", err))
}
if err := removeRegistryKeyFromDNSPolicyConfig(gpoDnsPolicyConfigCatchAllPath); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove GPO catch-all entry: %w", err))
}
for i := 0; i < r.nrptEntryCount; i++ {
localPath := fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i)
gpoPath := fmt.Sprintf("%s-%d", gpoDnsPolicyConfigMatchPath, i)
@@ -594,6 +642,20 @@ func refreshGroupPolicy() error {
return nil
}
func parseBoolEnv(key string) bool {
val := os.Getenv(key)
if val == "" {
return false
}
parsed, err := strconv.ParseBool(val)
if err != nil {
log.Warnf("failed to parse %s=%q: %v", key, val, err)
return false
}
return parsed
}
func closer(closer io.Closer) {
if err := closer.Close(); err != nil {
log.Errorf("failed to close: %s", err)

View File

@@ -94,6 +94,106 @@ func TestNRPTEntriesCleanupOnConfigChange(t *testing.T) {
assert.False(t, exists, "NRPT rule 2 should NOT exist after reducing to 75 domains")
}
func TestNRPTCatchAllRule(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
defer cleanupRegistryKeys(t)
cleanupRegistryKeys(t)
testIP := netip.MustParseAddr("100.64.0.1")
testGUID := "{12345678-1234-1234-1234-123456789ABC}"
interfacePath := `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\` + testGUID
testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE)
require.NoError(t, err, "Should create test interface registry key")
testKey.Close()
defer func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath)
}()
cfg := &registryConfigurator{guid: testGUID}
matchOnly := HostDNSConfig{
ServerIP: testIP,
Domains: []DomainConfig{{Domain: "example.com", MatchOnly: true}},
}
primary := HostDNSConfig{
ServerIP: testIP,
RouteAll: true,
Domains: []DomainConfig{{Domain: "example.com", MatchOnly: true}},
}
require.NoError(t, cfg.applyDNSConfig(matchOnly, nil))
exists, err := registryKeyExists(dnsPolicyConfigCatchAllPath)
require.NoError(t, err)
assert.False(t, exists, "catch-all rule should not exist for a match-only config")
require.NoError(t, cfg.applyDNSConfig(primary, nil))
exists, err = registryKeyExists(dnsPolicyConfigCatchAllPath)
require.NoError(t, err)
require.True(t, exists, "catch-all rule should exist when RouteAll is set")
k, err := registry.OpenKey(registry.LOCAL_MACHINE, dnsPolicyConfigCatchAllPath, registry.QUERY_VALUE)
require.NoError(t, err)
names, _, err := k.GetStringsValue(dnsPolicyConfigNameKey)
require.NoError(t, err)
assert.Equal(t, []string{nrptCatchAllNamespace}, names, "catch-all rule should match the root namespace")
servers, _, err := k.GetStringValue(dnsPolicyConfigGenericDNSServersKey)
require.NoError(t, err)
assert.Equal(t, testIP.String(), servers, "catch-all rule should list only our resolver")
opts, _, err := k.GetIntegerValue(dnsPolicyConfigConfigOptionsKey)
require.NoError(t, err)
assert.EqualValues(t, dnsPolicyConfigConfigOptionsValue, opts)
k.Close()
require.NoError(t, cfg.applyDNSConfig(matchOnly, nil))
exists, err = registryKeyExists(dnsPolicyConfigCatchAllPath)
require.NoError(t, err)
assert.False(t, exists, "catch-all rule should be removed when RouteAll is cleared")
require.NoError(t, cfg.applyDNSConfig(primary, nil))
require.NoError(t, cfg.restoreHostDNS())
exists, err = registryKeyExists(dnsPolicyConfigCatchAllPath)
require.NoError(t, err)
assert.False(t, exists, "catch-all rule should be removed on restore")
}
func TestNRPTCatchAllRuleDisabledByEnv(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
defer cleanupRegistryKeys(t)
cleanupRegistryKeys(t)
t.Setenv(envDisableCatchAllNRPT, "true")
testGUID := "{12345678-1234-1234-1234-123456789ABC}"
interfacePath := `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\` + testGUID
testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE)
require.NoError(t, err, "Should create test interface registry key")
testKey.Close()
defer func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, interfacePath)
}()
cfg := &registryConfigurator{guid: testGUID}
config := HostDNSConfig{
ServerIP: netip.MustParseAddr("100.64.0.1"),
RouteAll: true,
}
require.NoError(t, cfg.applyDNSConfig(config, nil))
exists, err := registryKeyExists(dnsPolicyConfigCatchAllPath)
require.NoError(t, err)
assert.False(t, exists, "catch-all rule should not be installed when disabled by env")
}
func registryKeyExists(path string) (bool, error) {
k, err := registry.OpenKey(registry.LOCAL_MACHINE, path, registry.QUERY_VALUE)
if err != nil {