From d2e62e358a07333462fa1e60587ef2964af85b1f Mon Sep 17 00:00:00 2001 From: Riccardo Manfrin <3090891+riccardomanfrin@users.noreply.github.com> Date: Tue, 8 Sep 2026 16:32:16 +0200 Subject: [PATCH 1/9] [client] Compare MDM-managed URLs as endpoints, not as strings (#7472) A policy that enforces a management URL refuses any SetConfig or Login whose URL differs from it. The comparison normalized only the default port, so three ways of writing the very endpoint the policy names were reported as conflicts: policy https://mgmt.example.com vs https://mgmt.example.com/ refused https://MGMT.example.com refused https://mgmt.example.com:0443 refused For an MDM-managed deployment whose stored or command-line URL is spelled differently from the policy's value, that means every settings update is refused with an MDMManagedFieldsViolation naming a field the caller did not change. `netbird up --management-url https://MGMT.example.com` reproduces it. The rules now live in util.SameServiceURL, and ConflictURL delegates: scheme and host compared case-insensitively, the effective port normalized numerically, a trailing slash ignored, and a path otherwise still part of the identity so /other remains a divergence. Unparseable input falls back to string equality. util rather than either caller, because comparing two service URLs is neither device management nor profile storage, and more than one place does it: an MDM-enforced management URL against a requested one here, a stored profile URL against a command-line one in profilemanager and the SSH gate. Every copy of these rules that drifts turns an equivalent URL into a refused request, which is how this one arose. CanonicalURL is left alone: besides comparison it is the canonical value handed to mdm.Restrictions and to the Android and iOS Preferences getters, and normalizing what those return is a separate decision. --- client/mdm/conflicts.go | 13 +++++-- client/mdm/conflicts_test.go | 40 ++++++++++++++++++++ util/serviceurl.go | 69 ++++++++++++++++++++++++++++++++++ util/serviceurl_test.go | 73 ++++++++++++++++++++++++++++++++++++ 4 files changed, 191 insertions(+), 4 deletions(-) create mode 100644 client/mdm/conflicts_test.go create mode 100644 util/serviceurl.go create mode 100644 util/serviceurl_test.go diff --git a/client/mdm/conflicts.go b/client/mdm/conflicts.go index a04cfb05c..160212afb 100644 --- a/client/mdm/conflicts.go +++ b/client/mdm/conflicts.go @@ -1,6 +1,10 @@ package mdm -import "net/url" +import ( + "net/url" + + "github.com/netbirdio/netbird/util" +) // PreSharedKeyRedactedSentinel is the redaction mask returned in place of a // real pre-shared key; an incoming value equal to it is a round-trip echo, @@ -44,8 +48,9 @@ func ConflictStringPtr(key string, p *string) ConflictCheck { } } -// ConflictURL builds a ConflictCheck for a URL-typed MDM key; both sides are -// normalized via CanonicalURL before comparison. +// ConflictURL builds a ConflictCheck for a URL-typed MDM key. The two sides are +// compared as the endpoints they address, not as strings: see +// util.SameServiceURL. func ConflictURL(key, got string) ConflictCheck { return ConflictCheck{ Key: key, @@ -54,7 +59,7 @@ func ConflictURL(key, got string) ConflictCheck { return true } want, ok := pol.GetString(key) - return ok && CanonicalURL(want) == CanonicalURL(got) + return ok && util.SameServiceURLStrings(want, got) }, } } diff --git a/client/mdm/conflicts_test.go b/client/mdm/conflicts_test.go new file mode 100644 index 000000000..d145ec103 --- /dev/null +++ b/client/mdm/conflicts_test.go @@ -0,0 +1,40 @@ +package mdm + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// The same spellings, through the conflict check that decides whether a request +// is refused. An enforced URL restated in another spelling addresses the very +// server the policy names, so it must not be reported as a conflict. +func TestConflictURLComparesEndpoints(t *testing.T) { + policy := NewPolicy(map[string]any{KeyManagementURL: "https://mgmt.example.com"}) + require.True(t, policy.HasKey(KeyManagementURL)) + + for _, restated := range []string{ + "https://mgmt.example.com", + "https://mgmt.example.com:443", + "https://mgmt.example.com/", + "https://MGMT.example.com", + "https://mgmt.example.com:0443", + } { + conflicts := ResolveConflicts(policy, []ConflictCheck{ConflictURL(KeyManagementURL, restated)}) + assert.Empty(t, conflicts, "%q is the enforced endpoint written differently", restated) + } + + for _, diverging := range []string{ + "https://other.example.com", + "http://mgmt.example.com", + "https://mgmt.example.com:8443", + "https://mgmt.example.com/other", + } { + conflicts := ResolveConflicts(policy, []ConflictCheck{ConflictURL(KeyManagementURL, diverging)}) + assert.Equal(t, []string{KeyManagementURL}, conflicts, "%q addresses another endpoint", diverging) + } + + // An unset field is not a request to change anything. + assert.Empty(t, ResolveConflicts(policy, []ConflictCheck{ConflictURL(KeyManagementURL, "")})) +} diff --git a/util/serviceurl.go b/util/serviceurl.go new file mode 100644 index 000000000..ffd287df1 --- /dev/null +++ b/util/serviceurl.go @@ -0,0 +1,69 @@ +package util + +import ( + "net/url" + "strconv" + "strings" +) + +// SameServiceURL reports whether two service URLs address the same endpoint. +// One endpoint can be written several ways, and every spelling below reaches +// the same server, so none of them is a divergence from another: +// +// an implicit default port https://mgmt.example.com :443 +// a zero-padded port https://mgmt.example.com:0443 +// a different host case https://MGMT.example.com +// a trailing slash https://mgmt.example.com/ +// +// A path is otherwise part of the identity: https://mgmt.example.com and +// https://mgmt.example.com/other are two endpoints. +// +// It lives here rather than next to any one caller because several of them +// compare the same kind of URL — an MDM-enforced management URL against a +// requested one, a stored profile URL against a command-line one — and every +// copy of these rules that drifts turns an equivalent URL into a refused +// request. +func SameServiceURL(a, b *url.URL) bool { + if a == nil || b == nil { + return a == b + } + + return strings.EqualFold(a.Hostname(), b.Hostname()) && + strings.EqualFold(a.Scheme, b.Scheme) && + ServiceURLPort(a) == ServiceURLPort(b) && + strings.TrimSuffix(a.Path, "/") == strings.TrimSuffix(b.Path, "/") +} + +// SameServiceURLStrings is SameServiceURL for unparsed input. Input that does +// not parse falls back to string equality, which is the strictest thing left +// to do with it. +func SameServiceURLStrings(a, b string) bool { + ua, errA := url.ParseRequestURI(a) + ub, errB := url.ParseRequestURI(b) + if errA != nil || errB != nil { + return a == b + } + + return SameServiceURL(ua, ub) +} + +// ServiceURLPort is the port a URL addresses: the one it carries, normalized +// numerically so ":0443" and ":443" are one port, or the scheme's default. +func ServiceURLPort(u *url.URL) string { + port := u.Port() + if port == "" { + switch strings.ToLower(u.Scheme) { + case "https": + return "443" + case "http": + return "80" + default: + return "" + } + } + + if n, err := strconv.Atoi(port); err == nil { + return strconv.Itoa(n) + } + return port +} diff --git a/util/serviceurl_test.go b/util/serviceurl_test.go new file mode 100644 index 000000000..af32a7c29 --- /dev/null +++ b/util/serviceurl_test.go @@ -0,0 +1,73 @@ +package util + +import ( + "net/url" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSameServiceURLSpellings(t *testing.T) { + tests := []struct { + a, b string + want bool + }{ + // One endpoint, written several ways. + {a: "https://mgmt.example.com", b: "https://mgmt.example.com:443", want: true}, + {a: "https://mgmt.example.com", b: "https://mgmt.example.com/", want: true}, + {a: "https://mgmt.example.com/", b: "https://mgmt.example.com:443/", want: true}, + {a: "https://MGMT.example.com", b: "https://mgmt.example.com", want: true}, + {a: "https://mgmt.example.com:0443", b: "https://mgmt.example.com:443", want: true}, + {a: "http://mgmt.example.com", b: "http://mgmt.example.com:80", want: true}, + {a: "HTTPS://mgmt.example.com", b: "https://mgmt.example.com", want: true}, + + // Different endpoints. + {a: "https://mgmt.example.com", b: "http://mgmt.example.com", want: false}, + {a: "https://mgmt.example.com", b: "https://mgmt.example.com:8443", want: false}, + {a: "https://mgmt.example.com", b: "https://other.example.com", want: false}, + {a: "https://mgmt.example.com", b: "https://mgmt.example.com/other", want: false}, + + // Unparseable input falls back to string equality. + {a: "mgmt.example.com", b: "mgmt.example.com", want: true}, + {a: "mgmt.example.com", b: "https://mgmt.example.com", want: false}, + } + + for _, tt := range tests { + t.Run(tt.a+" vs "+tt.b, func(t *testing.T) { + assert.Equal(t, tt.want, SameServiceURLStrings(tt.a, tt.b)) + assert.Equal(t, tt.want, SameServiceURLStrings(tt.b, tt.a), "the comparison must be symmetric") + }) + } +} + +// The parsed form is the primitive the string form delegates to, so it must +// answer the same for a spelling that only the parser can tell apart. +func TestSameServiceURLParsed(t *testing.T) { + parse := func(raw string) *url.URL { + t.Helper() + u, err := url.ParseRequestURI(raw) + require.NoError(t, err) + return u + } + + assert.True(t, SameServiceURL(parse("https://mgmt.example.com:0443/"), parse("https://MGMT.example.com"))) + assert.False(t, SameServiceURL(parse("https://mgmt.example.com"), parse("https://mgmt.example.com:8443"))) + + assert.True(t, SameServiceURL(nil, nil), "two absent URLs are the same absence") + assert.False(t, SameServiceURL(nil, parse("https://mgmt.example.com"))) +} + +func TestServiceURLPort(t *testing.T) { + parse := func(raw string) *url.URL { + t.Helper() + u, err := url.ParseRequestURI(raw) + require.NoError(t, err) + return u + } + + assert.Equal(t, "443", ServiceURLPort(parse("https://mgmt.example.com"))) + assert.Equal(t, "80", ServiceURLPort(parse("http://mgmt.example.com"))) + assert.Equal(t, "443", ServiceURLPort(parse("https://mgmt.example.com:0443"))) + assert.Equal(t, "8443", ServiceURLPort(parse("https://mgmt.example.com:8443"))) +} From d101f6cc46724129f954cb01ab22d1cba42ba30a Mon Sep 17 00:00:00 2001 From: Viktor Liu <17948409+lixmal@users.noreply.github.com> Date: Wed, 9 Sep 2026 18:32:11 +0900 Subject: [PATCH 2/9] [client] Redirect DNS port 53 with UDP and TCP DNAT instead of the eBPF forwarder (#7439) --- client/internal/dns/service_listener.go | 223 +++++++++++------- client/internal/dns/service_listener_test.go | 133 +++++++++++ client/internal/ebpf/ebpf/bpf_bpfeb.go | 36 ++- client/internal/ebpf/ebpf/bpf_bpfeb.o | Bin 14408 -> 8712 bytes client/internal/ebpf/ebpf/bpf_bpfel.go | 36 ++- client/internal/ebpf/ebpf/bpf_bpfel.o | Bin 14408 -> 8712 bytes client/internal/ebpf/ebpf/dns_fwd_linux.go | 52 ---- client/internal/ebpf/ebpf/manager_linux.go | 7 +- .../internal/ebpf/ebpf/manager_linux_test.go | 17 +- client/internal/ebpf/ebpf/src/bpf_map_def.h | 16 ++ client/internal/ebpf/ebpf/src/dns_fwd.c | 67 ------ client/internal/ebpf/ebpf/src/prog.c | 6 - client/internal/ebpf/ebpf/src/readme.md | 18 +- client/internal/ebpf/manager/manager.go | 6 +- 14 files changed, 363 insertions(+), 254 deletions(-) delete mode 100644 client/internal/ebpf/ebpf/dns_fwd_linux.go create mode 100644 client/internal/ebpf/ebpf/src/bpf_map_def.h delete mode 100644 client/internal/ebpf/ebpf/src/dns_fwd.c diff --git a/client/internal/dns/service_listener.go b/client/internal/dns/service_listener.go index 3dc29c4dc..d65a727b1 100644 --- a/client/internal/dns/service_listener.go +++ b/client/internal/dns/service_listener.go @@ -6,6 +6,7 @@ import ( "net" "net/netip" "runtime" + "slices" "strconv" "sync" "time" @@ -17,17 +18,20 @@ import ( nberrors "github.com/netbirdio/netbird/client/errors" firewall "github.com/netbirdio/netbird/client/firewall/manager" - "github.com/netbirdio/netbird/client/internal/ebpf" - ebpfMgr "github.com/netbirdio/netbird/client/internal/ebpf/manager" ) const ( customPort = 5053 + // randomPortAttempts bounds the search for a port free on both protocols. + randomPortAttempts = 5 ) var ( defaultIP = netip.MustParseAddr("127.0.0.1") customIP = netip.MustParseAddr("127.0.0.153") + + // dnatProtocols are the protocols the port 53 redirect covers. + dnatProtocols = []firewall.Protocol{firewall.ProtocolUDP, firewall.ProtocolTCP} ) type serviceViaListener struct { @@ -40,9 +44,20 @@ type serviceViaListener struct { listenPort uint16 listenerIsRunning bool listenerFlagLock sync.Mutex - ebpfService ebpfMgr.Manager firewall Firewall - tcpDNATConfigured bool + // dnatRules holds the port 53 redirects that are installed and not yet + // removed, so a removal that fails can be retried. + dnatRules []dnatRule +} + +// dnatRule is a port 53 redirect as it was installed. The target is kept with +// the rule because the listener can come back on a different address or port, +// and a retried removal has to name the address and port the rule was added +// with, not the ones in use now. +type dnatRule struct { + protocol firewall.Protocol + ip netip.Addr + port uint16 } func newServiceViaListener(wgIface WGIface, customAddr *netip.AddrPort, fw Firewall) *serviceViaListener { @@ -112,34 +127,93 @@ func (s *serviceViaListener) Listen() error { } }() - // When eBPF redirects UDP port 53 to our listen port, TCP still needs - // a DNAT rule because eBPF only handles UDP. - if s.ebpfService != nil && s.firewall != nil && s.listenPort != DefaultPort { - if err := s.firewall.AddOutputDNAT(s.listenIP, firewall.ProtocolTCP, DefaultPort, s.listenPort); err != nil { - log.Warnf("failed to add DNS TCP DNAT rule, TCP DNS on port 53 will not work: %v", err) - } else { - s.tcpDNATConfigured = true - log.Infof("added DNS TCP DNAT rule: %s:%d -> %s:%d", s.listenIP, DefaultPort, s.listenIP, s.listenPort) - } + if s.listenPort != DefaultPort { + s.setupDNAT() } return nil } +// setupDNAT redirects port 53 to the port the DNS server actually listens on. +// Both protocols must be redirected or none: RuntimePort reports port 53 only +// while the full redirect is in place, so a half-configured redirect would +// advertise a resolver that answers over one protocol. +func (s *serviceViaListener) setupDNAT() { + if s.firewall == nil { + log.Errorf("no firewall manager available to redirect DNS port %d to %d, "+ + "clients pointed at %s will not reach the resolver", DefaultPort, s.listenPort, s.listenIP) + return + } + + // Clear whatever an earlier removal left behind first. Those rules can point + // at an address or port this listener no longer uses, and they are matched + // before anything added now, so adding a redirect on top of one would keep + // sending port 53 traffic to the previous listener while reporting the + // redirect as complete. The rules stay recorded for a later attempt. + if err := s.removeDNAT(); err != nil { + log.Errorf("failed to remove stale DNS DNAT rules, leaving port %d redirected to the previous listener: %v", + DefaultPort, err) + return + } + + for _, proto := range dnatProtocols { + if err := s.firewall.AddOutputDNAT(s.listenIP, proto, DefaultPort, s.listenPort); err != nil { + log.Errorf("failed to add DNS %s DNAT rule, DNS on port %d will not work: %v", + proto, DefaultPort, err) + if err := s.removeDNAT(); err != nil { + log.Warnf("failed to roll back DNS DNAT rules, retrying on stop: %v", err) + } + return + } + s.dnatRules = append(s.dnatRules, dnatRule{protocol: proto, ip: s.listenIP, port: s.listenPort}) + } + + log.Infof("added DNS DNAT rules: %s:%d -> %s:%d (UDP + TCP)", s.listenIP, DefaultPort, s.listenIP, s.listenPort) +} + +// removeDNAT removes every installed port 53 redirect. A rule whose removal +// fails stays recorded so a later setup or Stop retries it, rather than leaving +// port 53 pointing at a resolver that is no longer listening. +func (s *serviceViaListener) removeDNAT() error { + if s.firewall == nil { + return nil + } + + var merr *multierror.Error + var remaining []dnatRule + for _, rule := range s.dnatRules { + if err := s.firewall.RemoveOutputDNAT(rule.ip, rule.protocol, DefaultPort, rule.port); err != nil { + merr = multierror.Append(merr, fmt.Errorf("remove DNS %s DNAT rule for %s:%d: %w", + rule.protocol, rule.ip, rule.port, err)) + remaining = append(remaining, rule) + } + } + s.dnatRules = remaining + + return nberrors.FormatErrorOrNil(merr) +} + func (s *serviceViaListener) Stop() error { s.listenerFlagLock.Lock() defer s.listenerFlagLock.Unlock() + var merr *multierror.Error + + // Redirects are removed even when the listener is already stopped, so that + // a removal which failed earlier is retried instead of leaving port 53 + // pointing at a resolver that no longer listens. + if err := s.removeDNAT(); err != nil { + merr = multierror.Append(merr, err) + } + if !s.listenerIsRunning { - return nil + return nberrors.FormatErrorOrNil(merr) } s.listenerIsRunning = false ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - var merr *multierror.Error - if err := s.server.ShutdownContext(ctx); err != nil { merr = multierror.Append(merr, fmt.Errorf("stop DNS UDP server: %w", err)) } @@ -148,19 +222,6 @@ func (s *serviceViaListener) Stop() error { merr = multierror.Append(merr, fmt.Errorf("stop DNS TCP server: %w", err)) } - if s.tcpDNATConfigured && s.firewall != nil { - if err := s.firewall.RemoveOutputDNAT(s.listenIP, firewall.ProtocolTCP, DefaultPort, s.listenPort); err != nil { - merr = multierror.Append(merr, fmt.Errorf("remove DNS TCP DNAT rule: %w", err)) - } - s.tcpDNATConfigured = false - } - - if s.ebpfService != nil { - if err := s.ebpfService.FreeDNSFwd(); err != nil { - merr = multierror.Append(merr, fmt.Errorf("stop traffic forwarder: %w", err)) - } - } - return nberrors.FormatErrorOrNil(merr) } @@ -177,11 +238,23 @@ func (s *serviceViaListener) RuntimePort() int { s.listenerFlagLock.Lock() defer s.listenerFlagLock.Unlock() - if s.ebpfService != nil { + if s.redirectInstalled() { return DefaultPort - } else { - return int(s.listenPort) } + return int(s.listenPort) +} + +// redirectInstalled reports whether every protocol is redirected from port 53 +// to the address and port the listener currently serves. Rules left over from +// an earlier listener do not count. +func (s *serviceViaListener) redirectInstalled() bool { + for _, proto := range dnatProtocols { + current := dnatRule{protocol: proto, ip: s.listenIP, port: s.listenPort} + if !slices.Contains(s.dnatRules, current) { + return false + } + } + return true } func (s *serviceViaListener) RuntimeIP() netip.Addr { @@ -190,30 +263,29 @@ func (s *serviceViaListener) RuntimeIP() netip.Addr { // evalListenAddress figures out the listen address for the DNS server. // IPv4-only: all peers have a v4 overlay address, and DNS config points to v4. -// First checks port 53 on WG interface or lo, then tries eBPF on a random port, -// then falls back to port 5053. +// Prefers port 53 on the overlay interface or lo, so no redirect is needed at +// all; when it is taken it falls back to port 5053 and then to a random free +// port, both of which need the port 53 redirect set up by setupDNAT. func (s *serviceViaListener) evalListenAddress() (netip.Addr, uint16, error) { if s.customAddr != nil { return s.customAddr.Addr(), s.customAddr.Port(), nil } - ip, ok := s.testFreePort(DefaultPort) - if ok { + if ip, ok := s.testFreePort(DefaultPort); ok { return ip, DefaultPort, nil } - ebpfSrv, port, ok := s.tryToUseeBPF() - if ok { - s.ebpfService = ebpfSrv - return s.wgInterface.Address().IP, port, nil - } - - ip, ok = s.testFreePort(customPort) - if ok { + if ip, ok := s.testFreePort(customPort); ok { return ip, customPort, nil } - return netip.Addr{}, 0, fmt.Errorf("failed to find a free port for DNS server") + ip := s.wgInterface.Address().IP + port, err := s.randomFreePort(ip) + if err != nil { + return netip.Addr{}, 0, fmt.Errorf("find a free port for DNS server: %w", err) + } + + return ip, port, nil } func (s *serviceViaListener) testFreePort(port int) (netip.Addr, bool) { @@ -260,48 +332,25 @@ func (s *serviceViaListener) tryToBind(ip netip.Addr, port int) bool { return true } -// tryToUseeBPF decides whether to apply eBPF program to capture DNS traffic on port 53. -// This is needed because on some operating systems if we start a DNS server not on a default port 53, -// the domain name resolution won't work. So, in case we are running on Linux and picked a free -// port we should fall back to the eBPF solution that will capture traffic on port 53 and forward -// it to a local DNS server running on the chosen port. -func (s *serviceViaListener) tryToUseeBPF() (ebpfMgr.Manager, uint16, bool) { - if runtime.GOOS != "linux" { - return nil, 0, false +// randomFreePort returns a port that is free on ip for both UDP and TCP, since +// the DNS server binds both. The probe listeners are closed again, so the port +// is only likely, not guaranteed, to still be free when the server binds it. +func (s *serviceViaListener) randomFreePort(ip netip.Addr) (uint16, error) { + for range randomPortAttempts { + probeListener, err := net.ListenUDP("udp4", &net.UDPAddr{}) + if err != nil { + return 0, fmt.Errorf("bind random port: %w", err) + } + + port := uint16(probeListener.LocalAddr().(*net.UDPAddr).Port) + if err := probeListener.Close(); err != nil { + return 0, fmt.Errorf("free up probed port: %w", err) + } + + if s.tryToBind(ip, int(port)) { + return port, nil + } } - port, err := s.generateFreePort() //nolint:staticcheck,unused - if err != nil { - log.Warnf("failed to generate a free port for eBPF DNS forwarder server: %s", err) - return nil, 0, false - } - - ebpfSrv := ebpf.GetEbpfManagerInstance() - err = ebpfSrv.LoadDNSFwd(s.wgInterface.Address().IP, int(port)) - if err != nil { - log.Warnf("failed to load DNS forwarder eBPF program, error: %s", err) - return nil, 0, false - } - - return ebpfSrv, port, true -} - -func (s *serviceViaListener) generateFreePort() (uint16, error) { - ok := s.tryToBind(s.wgInterface.Address().IP, customPort) - if ok { - return customPort, nil - } - - probeListener, err := net.ListenUDP("udp4", &net.UDPAddr{}) - if err != nil { - log.Debugf("failed to bind random port for DNS: %s", err) - return 0, err - } - - port := uint16(probeListener.LocalAddr().(*net.UDPAddr).Port) - if err = probeListener.Close(); err != nil { - log.Debugf("failed to free up DNS port: %s", err) - return 0, err - } - return port, nil + return 0, fmt.Errorf("no port free for UDP and TCP on %s after %d attempts", ip, randomPortAttempts) } diff --git a/client/internal/dns/service_listener_test.go b/client/internal/dns/service_listener_test.go index 90ef71d19..b158a79fd 100644 --- a/client/internal/dns/service_listener_test.go +++ b/client/internal/dns/service_listener_test.go @@ -1,6 +1,7 @@ package dns import ( + "errors" "fmt" "net" "net/netip" @@ -10,6 +11,8 @@ import ( "github.com/miekg/dns" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + firewall "github.com/netbirdio/netbird/client/firewall/manager" ) func TestServiceViaListener_TCPAndUDP(t *testing.T) { @@ -84,3 +87,133 @@ func TestServiceViaListener_TCPAndUDP(t *testing.T) { require.NotEmpty(t, tcpResp.Answer) assert.Contains(t, tcpResp.Answer[0].String(), "192.0.2.1", "TCP response should contain expected IP") } + +type dnatCall struct { + rule dnatRule + added bool +} + +// fakeFirewall records DNAT calls and fails the ones named in addErrs/removeErrs. +type fakeFirewall struct { + calls []dnatCall + addErrs map[firewall.Protocol]error + removeErrs map[firewall.Protocol]error +} + +func (f *fakeFirewall) AddOutputDNAT(ip netip.Addr, protocol firewall.Protocol, _, translatedPort uint16) error { + if err := f.addErrs[protocol]; err != nil { + return err + } + f.calls = append(f.calls, dnatCall{rule: dnatRule{protocol: protocol, ip: ip, port: translatedPort}, added: true}) + return nil +} + +func (f *fakeFirewall) RemoveOutputDNAT(ip netip.Addr, protocol firewall.Protocol, _, translatedPort uint16) error { + if err := f.removeErrs[protocol]; err != nil { + return err + } + f.calls = append(f.calls, dnatCall{rule: dnatRule{protocol: protocol, ip: ip, port: translatedPort}}) + return nil +} + +func newDNATTestService(fw Firewall) *serviceViaListener { + return &serviceViaListener{ + listenIP: netip.MustParseAddr("100.64.0.1"), + listenPort: customPort, + firewall: fw, + } +} + +func TestSetupDNAT_BothProtocols(t *testing.T) { + svc := newDNATTestService(&fakeFirewall{}) + + svc.setupDNAT() + + assert.Len(t, svc.dnatRules, len(dnatProtocols)) + assert.Equal(t, DefaultPort, svc.RuntimePort(), "port 53 is advertised once both redirects are installed") +} + +func TestSetupDNAT_RollsBackPartialRedirect(t *testing.T) { + fw := &fakeFirewall{addErrs: map[firewall.Protocol]error{firewall.ProtocolTCP: errors.New("nftables busy")}} + svc := newDNATTestService(fw) + + svc.setupDNAT() + + assert.Empty(t, svc.dnatRules, "the UDP redirect installed before the failure must be rolled back") + assert.Equal(t, int(svc.listenPort), svc.RuntimePort(), "an incomplete redirect must not advertise port 53") + udp := dnatRule{protocol: firewall.ProtocolUDP, ip: svc.listenIP, port: svc.listenPort} + assert.Contains(t, fw.calls, dnatCall{rule: udp}, "UDP removal should have been attempted") +} + +// A rollback that fails must keep the rule recorded, so port 53 is not left +// redirected to a resolver that no longer listens. +func TestStop_RetriesFailedDNATRemoval(t *testing.T) { + fw := &fakeFirewall{ + addErrs: map[firewall.Protocol]error{firewall.ProtocolTCP: errors.New("nftables busy")}, + removeErrs: map[firewall.Protocol]error{firewall.ProtocolUDP: errors.New("nftables busy")}, + } + svc := newDNATTestService(fw) + + svc.setupDNAT() + udp := dnatRule{protocol: firewall.ProtocolUDP, ip: svc.listenIP, port: svc.listenPort} + require.Equal(t, []dnatRule{udp}, svc.dnatRules, "a failed rollback keeps the rule for a later retry") + + require.Error(t, svc.Stop(), "the failing removal should be reported") + require.Equal(t, []dnatRule{udp}, svc.dnatRules) + + delete(fw.removeErrs, firewall.ProtocolUDP) + require.NoError(t, svc.Stop(), "a later stop retries the removal") + assert.Empty(t, svc.dnatRules) +} + +// A stale rule that cannot be removed is matched before anything added now, so +// no new redirect may be installed on top of it and port 53 must not be +// advertised as reaching this listener. +func TestSetupDNAT_AbortsWhileStaleRuleRemains(t *testing.T) { + fw := &fakeFirewall{removeErrs: map[firewall.Protocol]error{firewall.ProtocolUDP: errors.New("nftables busy")}} + svc := newDNATTestService(fw) + stalePort := svc.listenPort + + svc.setupDNAT() + require.Error(t, svc.Stop()) + staleUDP := dnatRule{protocol: firewall.ProtocolUDP, ip: svc.listenIP, port: stalePort} + require.Equal(t, []dnatRule{staleUDP}, svc.dnatRules) + + svc.listenPort = stalePort + 1 + fw.calls = nil + + svc.setupDNAT() + + assert.Equal(t, []dnatRule{staleUDP}, svc.dnatRules, "the stale rule stays recorded for a later attempt") + for _, call := range fw.calls { + assert.False(t, call.added, "no redirect may be installed while a stale one is still in place") + } + assert.Equal(t, int(svc.listenPort), svc.RuntimePort(), "port 53 must not be advertised") +} + +// A rule left behind by a failed removal must be removed with the address and +// port it was installed with, even when the listener has since moved to another +// port, and it must not count towards the redirect the new listener advertises. +func TestSetupDNAT_ClearsStaleRuleAfterPortChange(t *testing.T) { + fw := &fakeFirewall{removeErrs: map[firewall.Protocol]error{firewall.ProtocolUDP: errors.New("nftables busy")}} + svc := newDNATTestService(fw) + stalePort := svc.listenPort + + svc.setupDNAT() + require.Error(t, svc.Stop()) + staleUDP := dnatRule{protocol: firewall.ProtocolUDP, ip: svc.listenIP, port: stalePort} + require.Equal(t, []dnatRule{staleUDP}, svc.dnatRules) + + delete(fw.removeErrs, firewall.ProtocolUDP) + svc.listenPort = stalePort + 1 + fw.calls = nil + + svc.setupDNAT() + + assert.Contains(t, fw.calls, dnatCall{rule: staleUDP}, "the stale rule must be removed with its original port") + assert.Len(t, svc.dnatRules, len(dnatProtocols)) + assert.Equal(t, DefaultPort, svc.RuntimePort(), "the new listener is fully redirected") + for _, rule := range svc.dnatRules { + assert.Equal(t, svc.listenPort, rule.port, "only rules for the current listener remain") + } +} diff --git a/client/internal/ebpf/ebpf/bpf_bpfeb.go b/client/internal/ebpf/ebpf/bpf_bpfeb.go index 04b19883b..4b6230217 100644 --- a/client/internal/ebpf/ebpf/bpf_bpfeb.go +++ b/client/internal/ebpf/ebpf/bpf_bpfeb.go @@ -1,5 +1,5 @@ // Code generated by bpf2go; DO NOT EDIT. -//go:build arm64be || armbe || mips || mips64 || mips64p32 || ppc64 || s390 || s390x || sparc || sparc64 +//go:build mips || mips64 || ppc64 || s390x package ebpf @@ -47,9 +47,10 @@ func loadBpfObjects(obj interface{}, opts *ebpf.CollectionOptions) error { type bpfSpecs struct { bpfProgramSpecs bpfMapSpecs + bpfVariableSpecs } -// bpfSpecs contains programs before they are loaded into the kernel. +// bpfProgramSpecs contains programs before they are loaded into the kernel. // // It can be passed ebpf.CollectionSpec.Assign. type bpfProgramSpecs struct { @@ -61,17 +62,28 @@ type bpfProgramSpecs struct { // It can be passed ebpf.CollectionSpec.Assign. type bpfMapSpecs struct { NbFeatures *ebpf.MapSpec `ebpf:"nb_features"` - NbMapDnsIp *ebpf.MapSpec `ebpf:"nb_map_dns_ip"` - NbMapDnsPort *ebpf.MapSpec `ebpf:"nb_map_dns_port"` NbWgProxySettingsMap *ebpf.MapSpec `ebpf:"nb_wg_proxy_settings_map"` } +// bpfVariableSpecs contains global variables before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type bpfVariableSpecs struct { + FlagFeatureWgProxy *ebpf.VariableSpec `ebpf:"flag_feature_wg_proxy"` + MapKeyFeatures *ebpf.VariableSpec `ebpf:"map_key_features"` + MapKeyProxyPort *ebpf.VariableSpec `ebpf:"map_key_proxy_port"` + MapKeyWgPort *ebpf.VariableSpec `ebpf:"map_key_wg_port"` + ProxyPort *ebpf.VariableSpec `ebpf:"proxy_port"` + WgPort *ebpf.VariableSpec `ebpf:"wg_port"` +} + // bpfObjects contains all objects after they have been loaded into the kernel. // // It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. type bpfObjects struct { bpfPrograms bpfMaps + bpfVariables } func (o *bpfObjects) Close() error { @@ -86,20 +98,28 @@ func (o *bpfObjects) Close() error { // It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. type bpfMaps struct { NbFeatures *ebpf.Map `ebpf:"nb_features"` - NbMapDnsIp *ebpf.Map `ebpf:"nb_map_dns_ip"` - NbMapDnsPort *ebpf.Map `ebpf:"nb_map_dns_port"` NbWgProxySettingsMap *ebpf.Map `ebpf:"nb_wg_proxy_settings_map"` } func (m *bpfMaps) Close() error { return _BpfClose( m.NbFeatures, - m.NbMapDnsIp, - m.NbMapDnsPort, m.NbWgProxySettingsMap, ) } +// bpfVariables contains all global variables after they have been loaded into the kernel. +// +// It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. +type bpfVariables struct { + FlagFeatureWgProxy *ebpf.Variable `ebpf:"flag_feature_wg_proxy"` + MapKeyFeatures *ebpf.Variable `ebpf:"map_key_features"` + MapKeyProxyPort *ebpf.Variable `ebpf:"map_key_proxy_port"` + MapKeyWgPort *ebpf.Variable `ebpf:"map_key_wg_port"` + ProxyPort *ebpf.Variable `ebpf:"proxy_port"` + WgPort *ebpf.Variable `ebpf:"wg_port"` +} + // bpfPrograms contains all programs after they have been loaded into the kernel. // // It can be passed to loadBpfObjects or ebpf.CollectionSpec.LoadAndAssign. diff --git a/client/internal/ebpf/ebpf/bpf_bpfeb.o b/client/internal/ebpf/ebpf/bpf_bpfeb.o index 7433ad740ac150d19705f49d188d055d3c4c8cc2..b435d49647544d14fc150e72031fe0251e485a42 100644 GIT binary patch literal 8712 zcmds6Z;Vw(6`%X|57^drt3cP5wzCztZz(LhfRuzl9taf-DNh8UhI;qyzPr26ef!?y z-M6qvwHs3nsSV*nO;%&NhGEpi17nPP57XZ#2EbjX3pH5 zH?R@yHz%1p=Xd7JIcLtCnLGEr*Z1syDVK{>RwDHe&>kb}0rAXog9`HOQBiM?p|eM? z&PYRi-NNYNh$U7kWkytFTqr-)XXPGLFZ6YDgwuE3>g`!dulNyNHlvjKmNlm?k6PmL zaodLWE00_L${y84YN`Iq;{Rc-;M%=%pisrA2w~jdW7wS=2I;{u()2s#E^S}9fy0= z1~iYm^|0|!5!Qbl`$Iqf@AIQyJ;Qn5CqJGe|G&tOJmryBYuF-R%y=(93_r{d%A!Qx z$;ABP+OvHKmE|7Zr;Gd5pWi`7W6AiykP%zE%S0wA?6}BC!3LD_oLh!OANHj-4XDQa3IPu(8rwLz+gvTl0NL$=(x~`Pjph`I2HQxjYT;#G4h4w!V;YYPfk7) zY>|oH2LDl1l;}Hnkk+6VM8RXAe;_n%y(_jED#$&EF1siY#W@!J1Oqu1Dx;s;7Ecvj z6`D4FEwl#h;x+^xCqut+#74hC3q8?wv>!ox!}*0)1HUh8juUAzwDV5_!z1fP>q}8g_CQ>!nN@iPXa6{T0+i-FsZZzw&rv}HAMmbTc&Ln!f z9&L=r)1^jSt=Fe!8gWuhrt|%^(b$P+$|$u?_fvK%Ic0Sja|7Fq!pgYL_Zh7N`n2H~ zb4CPq)|xIW`ax~eH_THE3xbWwR=t+yttE7RMO7M;<)*4kR@I55nO5pG)vBkct+<-h zRHdxOo27}kK0dBmt*U~*t@;=uIX0P$O{qAZ*{Wu0X=S38l=awTsR`C-GQ5R7cxGMc zB%A>=otXx^R4zADm9kM~^W>;a8#xNkY*eMzQZr@aU|OGPjwPy`q%CC}Yp_RAVz%7K zmVylJGxBMxIWyMMW|~^BtzdzKWi{)yXx(yff&3-eu7~t%ujwp&z`kguv3`f~U-#L_ zR`VN&Hx0cwv}Nj5=RjGvIyMPx$aF0^8JE+R%Lcnri{8jm8cinpSj{1u zES>_Su~T3q889p4yRh9DbFyAP#i7@CWt$1`bQvqvDw#)|)GUgp6R@z2nn{|*mGMfg zoXo1`Y}I~vR;evF0kWCsWk)B zaCoUL>O{O<=~(`r%y_abPt+?g#R+q+(4V)Qfqb=In;3u(TzK1SqqjJMC&$twtTgkC z+tQPXY~--DDY~y%Jb2*Hf%x$5q6;=Pu61$J{d&BLV?l7X{V06z>6fu2?Uh@t za)MLkR3lM5Eg1-9jy+3Gmu3j86LU`8-B!yxXB}G9 z&dsPJ!sEO3NmTcaMGnfcvo=mHp=j>z}CIbH0WVmOL;ggk1lu^ z{+qk0M%Z6g>ifjkY`?+qH*)^ZD)k=krZ=UZ7+0L=mWL^STky9L3f!r3J%WFs)Gg4t zUctXt>Sr1lmGR2mRq7h<@VOz2JFvFD;?-rV#fEIYySza> zOuyuPmS-++8q)9bW>&Dv8_ds)CuW{3Se)ez{5SqGkKn({vmEr8ytyR(#LSy39;UpT zpX5!JSCThbUP<0WcVxWFo8=Z~dDCO#$-KE9p>N63fgs=88S$)eeEKdW>H{JCNDC89 z#P46n=UiO*7YN~-Ty}f~yPX-xh=^+|e6x3?hoAe7Z`R)N^NW15Cc;;fn1J)&@#xnA zHW|Lh+sX={ZQyseajm^ifzR0bg?eB1o;eBH@*XBtUTH~nmOa}X3(6k!}~$wo#yhdvkaR4I{gjMjF({Zn<``FF!lL( zH!$~*>*sv^e&&mhX9LW5`2F(%J|Eyq0lpgG8v(u*;5+R&S{~ru0Otd|HNblVd^Et5 z0iF%;nE=lR_9-0}eT1~?z!tpP3qd+VPY2{7}?@1G0sxd1N& z_(FiM1o&EjZwB~wJ2u_)37U1^8xw zZ?|K!O}+V>^TfwGkRJ+gA;84|j|8~k@KdXLS3bS^+4kRO?*C0T*i5SOLma7Hhh7@w zuW}FYDQP~(%;(pAe%o{7v(G#KQ;#T_c4f{l3Jxv zHJF~7Hm!UYn)zTBYX7cX!|_%&8{-!(`Xrm7c#~xgn|3(fD3Hy|k zFox>J2kX#z$w=GZ)P*#SKmLp0IDf+NH)W77PnPfJb>OA!-=r`~INr&#?+&Vu$&=Fg zM@u}!9P^8>?{DZ>D*jw|mUF0#zx%(?&6n}th`NOR$F86crE?AszxehC(XoX6-Z?v{ zvxt`VdADNm?f<-fcfX1Ioayo(tFH-i_jcFm`^y;l>)~Y3t|(YwzI!+(oPY1X0mtFu A(f|Me literal 14408 zcmds7Z)_aLb)UUEN@SfTR8rSW@sBL2kma5gDW(CzzQ{647qRzFyd`-M zd3U-yN~cz>*}^Uuphfw?MTo*kh}JELAVv8hfdwRi_`yMlwkX&(HCQA~RiFXrG#}Wa zKxm+d`}@tj*_#`ciHf@YkO6n+{ocHJ^WK}8x3jl=`TT`XrBZ=PN}zrS+F>LuAdYuy zl#}}b74-FLI=S_Q38{%!4Gi8a7~+1V4v;EJzi*j3^!rWr8~JLt^@yptJrJnCqKTiq2-5<~T5Bm;L{{O)5k5TWv`t93S zeuTZz|Ab1hesNUoXWci6Xs`vkUo!#H75Zb%f znO%psWj#M1DwuY+obRuV4}?EX`S&p{A93}`zTR{H6%U(r1iL?gWb^W`=vjL2y|UgV zSQq`Ohx+_^v3|c9vQ@7ey-v_iuV1^4?D%$l`s-+~{k>PbZWr`$>i&YNkSac)GP=HQ z__q(6dVi?v!K3}EOzP?b>SuL**sS;IkVthu+n%=)G2|X#1M&)Y!hy2RE=rlWs1(-0 zFx&Tv$4^~8I;jrQ-+RbM4{14@kFJ`2Nj^HM9?|VTIUoIk$eDapH~o@)bX8^9?#)Yp zc{xJ|dErfB!P(C`cW;1fcZ(2V%{= z$?Iouxy81zeqY)iQ%~hDK-s^3L7UQN2J3o5DPBE)POKAvu`3XR0~p8W<$gQ_$qFz&w>7u(Bq)L z0{T8&d0Z?oNfp5VDipKeJP@o~yMa00UrQTubny2=Bc8#xMUJk}mhZRa5D^mtVcD<* z{|;F>73fK0U&sWzs6U7cCHN6;q(i94h=LD+zE5b3l{z3}GgOef7hT36h~n9mdI%la z7b;Vani_YNIwCa37!f*zO*U@B;L#b{jmK*0G1SnLI)(bfs2{g>p$>s_GUu3EcWtC! z625*8?5!_T7P*hIKL%~q>d=4(+r8_#7ARW55DJC$ok!b23e65ognW8kz=GL`MOiUz`8 z*E}NtwKL6XWi~sQw4YYT<{R~6ty3qW6UUFvmMinEqtli7W97=^?0hLccDdYa9J_p@ z88@yMYeO?oR*#$W^-35;m&uAItCdExSZPLMqhXfJmY7ki)|>b24TU{_$Vxm4v;A#- zu2bZSTaE0XBXe5)hAP))N_AD9nN`=~dZS#esAjc+s~OG46;&>&sd{lbs!mO*W^-26 z>eXg-60&`XRonQ_QE~sRjxGEd`TPFsLt0XV^xY9O{ER1u|-j$RjTR9#S$S79h_}6>+_RM ztQj~NhM3YLp|FhJa1gbPRe37RauJw#oe78C2FM+Kg2RVr&V)~$f9CwfCoe?jo_Q|H zKYQ^~c<2zuKs%kGfGjpySy;89(_go85W5z`tZp?v+V0v4(W^sc8_J2niC)|DF2hcn zZKdbDy|?tPx;`CYHd{AjTm7UxWyO;%%3bCxN3BbT&RvjNs|{mY1gmjSH<<|hZ^x!> z?VTW~z2c1#r}aXyZX_$kGXCdF#dQaZCA@hvOosLrPkU|j+w?LzGB$j2_{8zCu>l(& zu0MAnKvonEm<(n1R#=IzN2Nwn)~QZcjcnT*iNDsQw<#hqYy#~lKw`BDlBH>cTocad z!K|#)ZgxV4DJH;kC2Z+tQQxAtVo)>}gN4Lij~htKQ{_r2ZmD`Jx;h`v$B}_?XX|9M zrI5?Ho8j9uXHT97gjluQ>yCh(Kz5ibg@fpC13MKLk&ojB0Tk=pTL}Y?Oa3u1;AmXQ@BdhW$Yx`K)jwsq*s)I{g_zbkMvU9+35?Y?oPhPzE=UD!x!FLY!g+5*Q`#R3RP4b6v5kp6d zV^9b2EBNElI*03zapC=Ii?2rGxNJLKn^ZyY50zTNyP;Y%I5?}+YkY5568r*Qzc}8C zgQv_#$A_G0TunfsRI7Y&9#@litD$rRQQU&BKYsH+^BC=@b{Fu#NQ6q91Po2h>!07 zBc(pW{yXjd1@?#kY|mgl%;I$@-R)q=r?Upb9@T<9+yh$vno4Jkl z3VO+4Yk&HdV6i{FE?DeOTYHINfBKHJi~am5NBe{Rv_C@`$No&{VA!9r_I_KbH;8l6 zPR#yde^(dQKjtHoH|-i@zB3aJrhGv#-gh@JPniY5&ndMHTQZA+r6NvXf$>*uP$jPDNOt8>dQD_!;>hwOv>pCN+Z8o#7e3F9!2XpFr!1G|_iKerSbPMZ^ZD6cz+-A#u?y(gV}#Z@D1b%_{}_{<%!u2KQw+rsjnf9YT3bTUlIJ5 z@QZlW!5;x$b1>Usr`~6;O5WM@@Rq?gk0LZ$KFK@Xf^FXE1x74v9u43p_3J##yfZA= z=A8+Hle|+9%sc~KLCwL4hs{e#-dQp28nYe#YJWa2d1uYR%sU$n9s=HUFyfKqsqdP+ z6KqR6^9c3tIJh5p*TEy$$Gr~D0*n7T4?(}pQ%T+#HuAXZ!n~#D7xDQH<#P_k{4sBJ zwCByb&8(SypPF$n_I+yEV7u>A>w>Xwb>7+#{5{EAn-1of1nq)`Y|PqfZW|2V6P?o^>%yk^)_EWFQqAXr_E0lgo!HnoWn<) zZ3+I(;UiwQ1aGzZjJqwt-#Yvy%Z2_wZS$FrZPDKqN-c7eZ{HnLRkXP|`7_}LO* z&sn~H{z^meTW!AHzn1^c4u4yisPymiqCe&Owb+*Dy0#_whc-W*vs~1#^Fi!jlAE~R z#sgr#$A1v`o8a&B__X66!SDWW_>rTB{aIK)ms@L1C8xP29RZ2T{gi$XiU=!~nsv4F?^A9VP=z$}d0afgqgbRqv?@a6M{ za69dw?m$|okM9f?kAuc}Ho*nZ4>*|d?sM=8Xyn;M{!S;>YMt={Z@eG}d1t&}0@xWZ zm;rXi3zmVM@pN9{^HuEE&o4)xe%`p4_VfFX$fp)SKjh#w(DLm9ls>n1$Zvre;*$P6 zr>cGjGk#8g=If63C17X#^a`*uetHAg?SBW@mEQplamn}@($4sq0G zqYl0W+KH#4f6jcVZD418=<|#dPko+o=PL*KK3p=MKEF8n?fE5XXS`i~yPY_d2X_3? z``z(Z=Rqevd=GKsncDRF(;ru!R&*JZZro6=* zkNO06fz;_moJv`yz84tHSyy)R&53hRomWMYyyyf9- z5ASy3RJVuwJe>9LDG!f(c*4Uo9&UMf(ZkCgUiI)T4{vyQ%fs6q-tELX-#hDz=YMy- z>HOigpYqzrJv`yz84tHSyy)R&53hRomWMYyyyf9-5ASwj{_f%Wo9Xj#*2AYfJnrEM z56^hG<>5sSFMD_m*vTiEbq_ONyX|*8ywizwzINohdcF40!^0lVc{uOkf`@A!Uhwde zhgUqj=HYb@Z+iHShj%)$J{CLv@9Xv2Lk|yoIOpNKhYKFAd3eFYOCDbF@S2C$J-q4R zJ09NY#QK=-_`ko`YY#mW2 zbA~PB`z@^IcU*9-Pst$s+Qx?DgYVGn?6tW_|F1}Sy8V5M?e;>HEM8Va zjT>{imakamo7P<)>U2yINp<1u+2heKnQBDWW6>$Xk?5pwPp}?~9y4%M>JeR!M90Y< zjmB6Fn`$f?6rDEB@}@UE$oJD>5kWIYp^a}|2an&#lJTv5V6k*tzxxbM?p6;Oui&!wd-Kq}DLraZ+sf_# zGMLsLuYXTYI{cX5LxLh?_833u@A;L?1G|R2;bQT}KWo@Y+yo}_!2Ipi#a4uM&$vM! zy5e#{JN^eP3QBuUYN@gkH~qT{=1S*sw%K@*H;T*J&o-y@@qG`adqr36n=`(`cZ9vh zA6PJ5x><89{3fo8Gk$QO`rZW{w*_fzS2Rr**@qUo_6DB@*PkB|?4K}1dJ8puY93jsQjMRQ=QEnks1|6h^OV=s z&+C!<78+Gt2yl$H^3)}VXU|m{RYF+w4)6>1i&i75HrrQlO z{k-Y@eHGRid0X96nptXmTKqp*-?pfp@0X2Fn|e{}+kNd zahQ+SK0P6qYcTzQnLpiU`e(3sC)n++d$iz8Zl&q_nsq;j@kVgRYWeEc_)N|ez{$ED ziYd4GU_4kiZFSr9eBVzShyUC9=rb3nVNM@s7#`nKwRxZTX3uMZ>l>x0US zKN|_@Q_MX3=`X)qOG(2NuE3G}idY=3j zAb`F1o&z6%mt`8Z3Xi25wh7PrN`q(5)3sOl0Q^4T)8G$@eg=Gx@IBza2z~|nEE6Fq1U1sf(}uTxtFOrzPwN@s8J}~a$Gy^*Z##N`jNMc8lVnbaQWv2k*PfO`r-3?` zAz!5>N?jE`0G^aJSA+kF@U-=T;p>>M0FCsxwTzak=PmY2%x7F_`&ZztpEJVK#_xp> zz}vW`!INOCnYZyrW*v14J&*lu^dRe39cvK(Rx-Cj^{kuy%XFR_$R7gF{gI9cAAmnD zd!Uy2537-M~GvRx|PYXW|{TkFhr!~9RG)my#$`R zE6l;D`MgoET{~s4ty@kZAp%OZqovSTFWW`W*3PT4&Sx;jNnZd)esxL$U=py}tU1FU zgAV3vosA;rdgd3`*xBOfZ*_Fuz+BkTb2}V)=-AyPa>oDD4zI$4<#JI~qg+0$l&eV? zM@dpB4aTY~iN+GuwIiILYR;uN(^IiQa0G-VhN26eW&{r7G(l zjilO3eW6vaEk>DScs2Uf_C6$z8v?MG>B~{3)fog6rEDsE*Bq=KR zo0R(z$^N0Je^`a#=w3BiiVK6KC?E6><*E?%McsQig1@z)cnHow7>|xXoXh8{v5MKL zy#C8syLNIG{?=v{N{JfHYX{@-H66kFP zL1zI85|LRs=tOVR!34P}*&n2Xx8DwG;RAM63zh8$wEw|o6WL;FXZNo3E9pJEckSwM zHDP1>0ka3{+U%GjFd+*{(YY`mC$<_)rslnoRO*#<^pVZ9DHdk|l(AhvYcgO)$#r4B zHfD95ajIio*Oh%bz$1BND#__TVN^0G9El*pF{(y!92N!&rF=A|s$=1q(P%UZ4U86! z_9tTsXNNh$-Uhk2oIXOFJZ8{QunFXZk$ljJ`8KcvQ7##UX*j&p5w$ALsk9>ZOx9#_ zEYFq;FvSgXuUJ22Bps<@xir`TADDPEYN9tkgEz<0GmJLtjPrw+gShbZs61eI86`tF zjvbH4b@@IjMozW*fxvpxaY5fJG=R0X9~fekld>()ICc0$PnZplWEYet3t8*oXjAn_ zHhcWoiDTi(Ls=VauGneg6to8eMO+JlYsZhukKTSYk(9l%*~&+_RnAu;#oLmBQ2N?4 z>_~15!B47%DAo(mV9e6tD!io^T)VuPAmW{;teHh!tT|)%!X<0GTw$*sKmO{Q;g^nl z<#2f7&FtYYRAI5uAC=-rb@j#g(WsUgaAk$qbF80j1NqJ2n+e-){J8NMm&9kK7tr{d zNa7Fd3>sgzuA1i!a7KsE4b-*2X94(bwB`(+W6=j_@1ivw!n3U4%jjEAD3x(=3-Bqy zAE0kOt<<<+J-=6}3BmtDzq+W@6$dv1PYLFGR^6EKIe_dhz+CJ$x&UnpHrDi>QgaSo z4;(;`0w|%AbyLu zg7+EG|1R+U@0q&bIB^p&meA|n2VOO4>V)IOO~4Z}{yd&9n-Pag4rW|fH`G6X@um;) z4&vZvfv-832SvQk{ZQWdRFze>JU|Tm+4rbllbTH#HBiPo>J;AnaY+c#9c_8DqZW_J1Vcjv_ zx^8YE2>*qa)oz`6rB!NS< zZ2x7^NBA|_p6<~f>&LkG`wK0%vz5LNI1xuq`)U2N{RJ(5T7P4po}J0}eEM%awRZoY z<$UuqWc4kk8~k?BA_C%-Xc-6dJHsglQ-8?8w0*_FdryR`M;j;#b z`rPkP!R+|WCGTMRTX8VswGE9#F(`!`D1K<>46*-}dl5 z4?plQe?Q#)+VA<+|CBe*=Mq;x>tN=;*TWSLk9+u%hbKLJ-NQFMJnP}R9-i~CJ_zuW z`=h@t&iNt^Je>A$#=}_;_jJx2zt1fGPS#b8it-_j+`$vCb@8j*6TC|5*D?M2waB+U zJHPgN=f3LECO+E~@B)81z|o>2Q}3cI&=7+$k`vQ}l`c+`()qg$3*UQ9iiZ-!VHV*Mx=odoN|o zp-UM7J#x|{#PNUPyYP7qUh`U ze%`{|#q2LNpSmvbchi?~{9`eE$NnduO!EEtJLHSSe@^0mN#-&BT*{CCPRJIsf0GkH zS~7af)wS00~1IJ@xeujrYM+x7%bwpD$oLQ+6T5M zU>Ydm{{J)QTnCgL7JC6IE zcR0Sn6Kj9_NzquYK{t{zST5YePPIE=JIaGp?`AX zzg%^@HN3prbu=b^e@E7D>=7&w(mS}=gRW25UN|51_xgTw<5mBJ_PQVLx848meYU4z zE8TF{InaH0OZM}{Eqm^17tUAxv&VeF+JAa}*}b7Xp6mbZ{N*ZcT`q3fZeGJt{AJjG z{IM~GLEDVEysP=1^U-;~8Y&YBzg zg)`gJ&}9|jK2?AG%*CVXemmb!jz^C$sUA;6ORiAL>pBRsRPV6Ke9dz~5c(gZT z%l&o1(qLTo!{>|zZLQK{hr^0;@65>w=wa1dVXa+!}Ck0)BjA&8|*1@ zXN4z|*F(7c#B)SF^G~bHJoV-Nb7=1y|AuEFte-jFf@5W~#-+(Nkr}3u9m_y^V$c!bgLuSs1`i966i`5wO6OkiuBOAxKK7&6Gi_E@uI=O-C8{uO5q_&)u*}^;5 zV>Fh&4?xy&9u}E>bc-B8R=;&YWp-hrQjz%le9XNZ3WIg*a6Aq&QIwsUPI*}d@! z!VdNBF0grf1vE~(&uML#ejf5)VB=t)G;$CA5HkJ4^q$BO*7CTkV|oHU1IAuz2sm?N7p83(o}G`+__BDMQ8d5Y-2y>a8J)1v-9_wuN zk?7gxK2Oega)&2(ii{&-j(W15F}CwfuivM{hWc+qkKFeAMbUGw?f2wEp4{!pCO%y% zjhjl6%g5EERx1=os;0A+)N7{mXtgqQG+(U7!`JeiLncBhl;f+p@yR4Q6%Ce$;_sGSC5$@SPNI!?MqhBN!Lp^oU+t44LAwdP#|JL!&lg^Zyr$>~h%-fE(E$cdeL)uo=jq4`ZjfZTJ^T1ErH^;Vq5`@-wyc@(x zyHRu+=}u9_|9z|Y)4^tmUVSx6yAC$b`rg>5;pJF&PuJ63Cyw{@bf|y0|NM;rTWPSv zMJRW+qGEC_&R1))Pi?rWX4Xd|{DF{c{t!;t&eYZ`cPC-P zP@ID&@;K79oc$9gMMv>O0u=&#C8;7T4;PC0q;4wp_{wB5nZyp0jip1ix;z$R^z@Xqto%hyj@bqq8!&%#qg!<8wQMK@ENoiPTnbbxR(mL+a+emk} z8f$9JpN(`hvZdGo+B)*=D7{=OMY2JwnX?!A;{N#g{terc8`;|79zfIQ`}@zoc;UtP z;+cMRHd{QTeu~dKdCJD=;zL#f9~b@rRRU;Y<%J6 z{zQmiJXbFhlXFMEbNDYn=b(F^frNKJTzoHK-7~nD_|C&;;MSTkd=4d-aW&P^hVTvK zE%=i~{wA&;;(BPsm_FpRXZd#g84&&u`L1bWZVI2qgLL!j#z1B5Uk10%x_2Bl{??c!k6XaY9v=j+c>FYY)#Fa^n#avxzLc~7e?-0; z>(MIw-^g25-M1l@llOrmue=+~cX%z&fU_PqgZqR}U@z_YFZe?kP1-adY}X6F@^J56 zf$eYO+L6Wj3jYnb`3!!e5&l>3t_$#w$F1Nc;h(VnuVVgw{Q=Bh_*^4igvPN)VaNW+ zJHQ!XJD$2b11!Iaax;bbNclXt`3Cl@$DadF3+wn7gmwIj!aDvXk2(I^!aDvY^uzvb zeT=_VSjQiE%<=aM{|Nbxx3C@_H-r0xx8nDVJ#XW829I07mptbA7!YoU{i928f$&S< z);l<#!Xx15UCdwjcfmWCv3|ne0e|K%v3|nVzt|s)Ex!QQGCc{o7ni}hH3`20Zp}k4 z{3f`06!Q{Z0C&9V;t1>in2dGnm-05eXrFo=^YEB)1WnlfuA{tV&czwlCpUwqrTi`M zj_(^Y<8d>1R`>_B=e#`L4W9Ry~E<~(NHnfI7+$B$EtJIh`<b zxBq*w=U(&E8-QKjdM!)gd&qb4mrhBam;BF>@6<0gkagz&VD(L6wXTn?{z0SlHGjeC zx25&9ondK5S*Mr5_&H7;A7dKR2=bi%p&UlzrZlym!XUcJIPo|9I;E&VQ{ z{r{rXH>ESy_Df>FUD{FmYhtgSP_gil^Sz-D7ZaK98w0qQ z$a`^>2}tBNT(h{C$o!^r8yAzZ6@+bnED&aYS&!Kt?}bdv*x#tf>~F?n_NQ@$YTjdb z?_y%c_9AhJvhEi$zX45o<$Qmd_n6-=cyC}*`xTEP=$m9cS-uZfpU12}>M_eT9%*~C zUOD?;@|g8+d(82!dK`fxSy=6Vz+=|eeaiCvxadPB@&Q~61SB&3xk5l9bG_~mkjPxG zHVR10=-)1nxn5b1*;mZPk{`#^0yW(GF9N)6XoZtH%+x|Redj`+iOjL9JZNlu9{@w2} z{nzC&=ktum5xCFecJKv{4}u3gJ^-!@v;PdPX^*+yZ+Oi0ob#CeecR*x;KR6>*dBd) zoPebKjKh@A1bo5char!8%=YR5-wb#r;JJVo0$vJuCEz;&uLa!1K_Th<^gJl@KIWJA z2HY3$rGQ5Rt_OTG;F*Bu0$vDsDd3fW?*zOSaFd)Doo`#fnSgr(?hE)*z@q`z1HKvX zOu%yiF9f_4@JhgU0$vNaN#a8~|9~}qXt_QQDE9^SeAf5(i^dzZuLtEf1D*+ZF5rcL zmjYf1_)frU0XK1ANIJf@fHMK>bGF*|1?Bur!*6di;CjF}1D*+(?||AK5-z+C}n1MUxaAmDPqQvpv0JR9(Qz>5Jd2Yfr=)qwFTv0=Wg0Y?FM1)L4I zKj49Y%K=XXJRR_C!1Dnw2D}{b?SNMU)(1#k@7AV`{)htZ3OE}u|96T>+ZzbD9Pm`Y z(*YlO{E5f*Wey!Y+Iei_=jjb!cGB-Xot0!<{xgQs=?h=#|;Oh*nh;KP5MzP_}?6@ntoLBKQL*bd?53>E#UhtM$PZI(CVDhPWZKr z1?dOh&hhc96S4hYk;2IO_bHa^r>+fNQe`Lp!=siPG@aG!6E>HxSniwFT_5VKPvJ>( z?)2&7@z1-g+vYuSFR?p*TC^va_ry;*JSKU!&Aa2{R3D3bn02|VCq5>&l6Tv@JMQkL z9(kABP(S70=kc5)x4I22Jf4>9Wz4~91eiwtuX?@qaG5-SLUj0W0okhRbKbzG0 z+CP7*rinZ*xk~uEAck$J#tAa4RC@yw>j%R$%Tm9bS~?-D{~FXZ4eMv+r=#^>LhfBZ zC*$9&emJj?qfP3cmj1)>a;|=|zrFejk^7|nnElVdb%U_IIvTv!`js=TPOEgt_2P4h zxHbL#hub_H6KlRL3pcA3s6-0;=l?$5tN*DEYo31FX8-)(X>CsT&rWpyz`2k4v)dG| j^(5zyWA;-x)?b1mToc+~+UIKNb1v)g|E0Bm-IxCZwP{$s diff --git a/client/internal/ebpf/ebpf/dns_fwd_linux.go b/client/internal/ebpf/ebpf/dns_fwd_linux.go deleted file mode 100644 index 1e7774573..000000000 --- a/client/internal/ebpf/ebpf/dns_fwd_linux.go +++ /dev/null @@ -1,52 +0,0 @@ -package ebpf - -import ( - "encoding/binary" - "fmt" - "net/netip" - - log "github.com/sirupsen/logrus" -) - -const ( - mapKeyDNSIP uint32 = 0 - mapKeyDNSPort uint32 = 1 -) - -func (tf *GeneralManager) LoadDNSFwd(ip netip.Addr, dnsPort int) error { - log.Debugf("load eBPF DNS forwarder, watching addr: %s:53, redirect to port: %d", ip, dnsPort) - tf.lock.Lock() - defer tf.lock.Unlock() - - err := tf.loadXdp() - if err != nil { - return err - } - - if !ip.Is4() { - return fmt.Errorf("eBPF DNS forwarder only supports IPv4, got %s", ip) - } - ip4 := ip.As4() - err = tf.bpfObjs.NbMapDnsIp.Put(mapKeyDNSIP, binary.BigEndian.Uint32(ip4[:])) - if err != nil { - return err - } - - err = tf.bpfObjs.NbMapDnsPort.Put(mapKeyDNSPort, uint16(dnsPort)) - if err != nil { - return err - } - - tf.setFeatureFlag(featureFlagDnsForwarder) - err = tf.bpfObjs.NbFeatures.Put(mapKeyFeatures, tf.featureFlags) - if err != nil { - return err - } - return nil -} - -func (tf *GeneralManager) FreeDNSFwd() error { - log.Debugf("free ebpf DNS forwarder") - return tf.unsetFeatureFlag(featureFlagDnsForwarder) -} - diff --git a/client/internal/ebpf/ebpf/manager_linux.go b/client/internal/ebpf/ebpf/manager_linux.go index 7520a6387..a13f5f19a 100644 --- a/client/internal/ebpf/ebpf/manager_linux.go +++ b/client/internal/ebpf/ebpf/manager_linux.go @@ -15,8 +15,7 @@ import ( const ( mapKeyFeatures uint32 = 0 - featureFlagWGProxy = 0b00000001 - featureFlagDnsForwarder = 0b00000010 + featureFlagWGProxy = 0b00000001 ) var ( @@ -28,9 +27,9 @@ var ( // GeneralManager is used to load multiple eBPF programs with a custom check (if then) done in prog.c // The manager simply adds a feature (byte) of each program to a map that is shared between the userspace and kernel. -// When packet arrives, the C code checks for each feature (if it is set) and executes each enabled program (e.g., dns_fwd.c and wg_proxy.c). +// When packet arrives, the C code checks for each feature (if it is set) and executes each enabled program (e.g., wg_proxy.c). // -//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc clang-14 bpf src/prog.c -- -I /usr/x86_64-linux-gnu/include +//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc clang-14 bpf src/prog.c -- -I /usr/x86_64-linux-gnu/include -include src/bpf_map_def.h type GeneralManager struct { lock sync.Mutex link link.Link diff --git a/client/internal/ebpf/ebpf/manager_linux_test.go b/client/internal/ebpf/ebpf/manager_linux_test.go index 5664a4565..e09fcb977 100644 --- a/client/internal/ebpf/ebpf/manager_linux_test.go +++ b/client/internal/ebpf/ebpf/manager_linux_test.go @@ -7,33 +7,24 @@ import ( func TestManager_setFeatureFlag(t *testing.T) { mgr := GeneralManager{} mgr.setFeatureFlag(featureFlagWGProxy) - if mgr.featureFlags != 1 { + if mgr.featureFlags != featureFlagWGProxy { t.Errorf("invalid feature state") } - mgr.setFeatureFlag(featureFlagDnsForwarder) - if mgr.featureFlags != 3 { - t.Errorf("invalid feature state") + mgr.setFeatureFlag(featureFlagWGProxy) + if mgr.featureFlags != featureFlagWGProxy { + t.Errorf("setting a flag twice must be idempotent, got: %d", mgr.featureFlags) } } func TestManager_unsetFeatureFlag(t *testing.T) { mgr := GeneralManager{} mgr.setFeatureFlag(featureFlagWGProxy) - mgr.setFeatureFlag(featureFlagDnsForwarder) err := mgr.unsetFeatureFlag(featureFlagWGProxy) if err != nil { t.Errorf("unexpected error: %s", err) } - if mgr.featureFlags != 2 { - t.Errorf("invalid feature state, expected: %d, got: %d", 2, mgr.featureFlags) - } - - err = mgr.unsetFeatureFlag(featureFlagDnsForwarder) - if err != nil { - t.Errorf("unexpected error: %s", err) - } if mgr.featureFlags != 0 { t.Errorf("invalid feature state, expected: %d, got: %d", 0, mgr.featureFlags) } diff --git a/client/internal/ebpf/ebpf/src/bpf_map_def.h b/client/internal/ebpf/ebpf/src/bpf_map_def.h new file mode 100644 index 000000000..9528fb592 --- /dev/null +++ b/client/internal/ebpf/ebpf/src/bpf_map_def.h @@ -0,0 +1,16 @@ +// libbpf 1.0 removed struct bpf_map_def, but the programs here keep the legacy +// map definitions: they load on kernels built without BTF, which BTF-style +// (SEC(".maps")) definitions do not. Define the struct ourselves so the +// programs compile against current libbpf headers. +#ifndef NB_BPF_MAP_DEF_H +#define NB_BPF_MAP_DEF_H + +struct bpf_map_def { + unsigned int type; + unsigned int key_size; + unsigned int value_size; + unsigned int max_entries; + unsigned int map_flags; +}; + +#endif diff --git a/client/internal/ebpf/ebpf/src/dns_fwd.c b/client/internal/ebpf/ebpf/src/dns_fwd.c deleted file mode 100644 index 9f8de2001..000000000 --- a/client/internal/ebpf/ebpf/src/dns_fwd.c +++ /dev/null @@ -1,67 +0,0 @@ -const __u32 map_key_dns_ip = 0; -const __u32 map_key_dns_port = 1; - -struct bpf_map_def SEC("maps") nb_map_dns_ip = { - .type = BPF_MAP_TYPE_ARRAY, - .key_size = sizeof(__u32), - .value_size = sizeof(__u32), - .max_entries = 10, -}; - -struct bpf_map_def SEC("maps") nb_map_dns_port = { - .type = BPF_MAP_TYPE_ARRAY, - .key_size = sizeof(__u32), - .value_size = sizeof(__u16), - .max_entries = 10, -}; - -__be32 dns_ip = 0; -__be16 dns_port = 0; - -// 13568 is 53 in big endian -__be16 GENERAL_DNS_PORT = 13568; - -bool read_settings() { - __u16 *port_value; - __u32 *ip_value; - - // read dns ip - ip_value = bpf_map_lookup_elem(&nb_map_dns_ip, &map_key_dns_ip); - if(!ip_value) { - return false; - } - dns_ip = htonl(*ip_value); - - // read dns port - port_value = bpf_map_lookup_elem(&nb_map_dns_port, &map_key_dns_port); - if (!port_value) { - return false; - } - dns_port = htons(*port_value); - return true; -} - -int xdp_dns_fwd(struct iphdr *ip, struct udphdr *udp) { - if (dns_port == 0) { - if(!read_settings()){ - return XDP_PASS; - } - // bpf_printk("dns port: %d", ntohs(dns_port)); - // bpf_printk("dns ip: %d", ntohl(dns_ip)); - } - - if (udp->dest == GENERAL_DNS_PORT && ip->daddr == dns_ip) { - udp->dest = dns_port; - // Clear the now-stale checksum; zero means "not computed" for IPv4. - udp->check = 0; - return XDP_PASS; - } - - if (udp->source == dns_port && ip->saddr == dns_ip) { - udp->source = GENERAL_DNS_PORT; - udp->check = 0; - return XDP_PASS; - } - - return XDP_PASS; -} diff --git a/client/internal/ebpf/ebpf/src/prog.c b/client/internal/ebpf/ebpf/src/prog.c index f32103f28..44ee53458 100644 --- a/client/internal/ebpf/ebpf/src/prog.c +++ b/client/internal/ebpf/ebpf/src/prog.c @@ -5,11 +5,9 @@ #include #include #include -#include "dns_fwd.c" #include "wg_proxy.c" const __u16 flag_feature_wg_proxy = 0b01; -const __u16 flag_feature_dns_fwd = 0b10; const __u32 map_key_features = 0; struct bpf_map_def SEC("maps") nb_features = { @@ -48,10 +46,6 @@ int nb_xdp_prog(struct xdp_md *ctx) { return XDP_PASS; } - if (*features & flag_feature_dns_fwd) { - xdp_dns_fwd(ip, udp); - } - if (*features & flag_feature_wg_proxy) { xdp_wg_proxy(ip, udp); } diff --git a/client/internal/ebpf/ebpf/src/readme.md b/client/internal/ebpf/ebpf/src/readme.md index 0ab393dd4..aa47847da 100644 --- a/client/internal/ebpf/ebpf/src/readme.md +++ b/client/internal/ebpf/ebpf/src/readme.md @@ -1,8 +1,18 @@ -# DNS forwarder +# XDP programs -The agent attach the XDP program to the lo device. We can not use fake address in eBPF because the -traffic does not appear in the eBPF program. The program capture the traffic on wg_ip:53 and -overwrite in it the destination port to 5053. +`prog.c` is attached to the `lo` device and dispatches to the features enabled in the +`nb_features` map. The only feature is the WireGuard proxy (`wg_proxy.c`): it rewrites +loopback UDP sent from the WireGuard listen port so it reaches the userspace relay proxy +port instead, and swaps the peer endpoint port into the source so the proxy can tell +peers apart. + +Maps use the legacy `struct bpf_map_def` form, defined in `bpf_map_def.h` because libbpf +1.0 removed it. They load on kernels built without BTF, which BTF-style (`SEC(".maps")`) +definitions do not. + +Regenerate the objects with `go generate ./client/internal/ebpf/ebpf/`; it needs +`clang-14`. Loading a regenerated object needs root, attaching it needs `bpf_link` +(kernel >= 5.7), and only one XDP program can own `lo` at a time. # Debug diff --git a/client/internal/ebpf/manager/manager.go b/client/internal/ebpf/manager/manager.go index 25a767090..fdc5d8d82 100644 --- a/client/internal/ebpf/manager/manager.go +++ b/client/internal/ebpf/manager/manager.go @@ -1,11 +1,7 @@ package manager -import "net/netip" - -// Manager is used to load multiple eBPF programs. E.g., current DNS programs and WireGuard proxy +// Manager is used to load multiple eBPF programs. E.g., the WireGuard proxy type Manager interface { - LoadDNSFwd(ip netip.Addr, dnsPort int) error - FreeDNSFwd() error LoadWgProxy(proxyPort, wgPort int) error FreeWGProxy() error } From 269cbadfeb43a611423d461b666b65973d67be56 Mon Sep 17 00:00:00 2001 From: Pascal Fischer <32096965+pascal-fischer@users.noreply.github.com> Date: Wed, 9 Sep 2026 13:44:19 +0200 Subject: [PATCH 3/9] [management] expire and disconnect peers while including offline peers (#7467) --- management/server/account.go | 7 +- management/server/account_test.go | 179 +++++++++++++++++++++++++++- management/server/peer.go | 16 ++- management/server/scheduler.go | 9 +- management/server/scheduler_test.go | 89 ++++++++++++++ management/server/types/account.go | 5 +- management/server/user.go | 77 +++++++++--- 7 files changed, 350 insertions(+), 32 deletions(-) diff --git a/management/server/account.go b/management/server/account.go index 3ceef79db..6ccf673f5 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -719,8 +719,10 @@ func (am *DefaultAccountManager) schedulePeerLoginExpiration(ctx context.Context log.WithContext(ctx).Tracef("peer login expiration job for account %s is already scheduled", accountID) return } + // The job outlives the request that arms it, so it must not inherit the request's cancellation. + jobCtx := context.WithoutCancel(ctx) if nextRun, ok := am.getNextPeerExpiration(ctx, accountID); ok { - go am.peerLoginExpiry.Schedule(ctx, nextRun, accountID, am.peerLoginExpirationJob(ctx, accountID)) + go am.peerLoginExpiry.Schedule(jobCtx, nextRun, accountID, am.peerLoginExpirationJob(jobCtx, accountID)) } } @@ -752,8 +754,9 @@ func (am *DefaultAccountManager) peerInactivityExpirationJob(ctx context.Context // checkAndSchedulePeerInactivityExpiration periodically checks for inactive peers to end their sessions func (am *DefaultAccountManager) checkAndSchedulePeerInactivityExpiration(ctx context.Context, accountID string) { am.peerInactivityExpiry.Cancel(ctx, []string{accountID}) + jobCtx := context.WithoutCancel(ctx) if nextRun, ok := am.getNextInactivePeerExpiration(ctx, accountID); ok { - go am.peerInactivityExpiry.Schedule(ctx, nextRun, accountID, am.peerInactivityExpirationJob(ctx, accountID)) + go am.peerInactivityExpiry.Schedule(jobCtx, nextRun, accountID, am.peerInactivityExpirationJob(jobCtx, accountID)) } } diff --git a/management/server/account_test.go b/management/server/account_test.go index b462cc2a6..bd7bf2d97 100644 --- a/management/server/account_test.go +++ b/management/server/account_test.go @@ -1920,6 +1920,154 @@ func TestDefaultAccountManager_MarkPeerConnected_PeerLoginExpiration(t *testing. } } +func TestDefaultAccountManager_SchedulePeerLoginExpiration_IncludesOfflinePeers(t *testing.T) { + manager, updateManager, err := createManager(t) + require.NoError(t, err, "unable to create account manager") + + accountID, err := manager.GetAccountIDByUserID(context.Background(), auth.UserAuth{UserId: userID}) + require.NoError(t, err, "unable to create an account") + + connectedKey, offlineKey := addExpiringPeers(t, manager) + _, err = manager.UpdateAccountSettings(context.Background(), accountID, userID, &types.Settings{ + PeerLoginExpiration: time.Hour, + PeerLoginExpirationEnabled: true, + Extra: &types.ExtraSettings{}, + }) + require.NoError(t, err, "expecting to update account settings successfully but got error") + manager.peerLoginExpiry.CancelAll(context.Background()) + + // The connected peer logged in just now, so a job computed from connected peers alone + // would be armed for an hour. The offline peer's login expires in two seconds; a + // reconnect of that peer must not have to wait for the connected peer's tick. + now := time.Now().UTC() + setPeerLogin(t, manager, accountID, connectedKey, true, now) + setPeerLogin(t, manager, accountID, offlineKey, false, now.Add(-time.Hour+2*time.Second)) + + offlinePeer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, offlineKey) + require.NoError(t, err) + updateManager.CreateChannel(context.Background(), offlinePeer.ID) + + manager.peerLoginExpiry = NewDefaultScheduler() + t.Cleanup(func() { manager.peerLoginExpiry.CancelAll(context.Background()) }) + manager.schedulePeerLoginExpiration(context.Background(), accountID) + + // The flag is committed per peer before the disconnect fans out, so wait for both. + require.Eventually(t, func() bool { + peer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, offlineKey) + return err == nil && peer.Status.LoginExpired && !updateManager.HasChannel(offlinePeer.ID) + }, 10*time.Second, 100*time.Millisecond, "offline peer should be expired and disconnected at its own deadline") + + connectedPeer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, connectedKey) + require.NoError(t, err) + assert.False(t, connectedPeer.Status.LoginExpired, "connected peer with a fresh login must not expire") +} + +func TestDefaultAccountManager_SchedulePeerLoginExpiration_DetachesRequestContext(t *testing.T) { + manager, _, err := createManager(t) + require.NoError(t, err, "unable to create account manager") + + accountID, err := manager.GetAccountIDByUserID(context.Background(), auth.UserAuth{UserId: userID}) + require.NoError(t, err, "unable to create an account") + connectedKey, _ := addExpiringPeers(t, manager) + setPeerLogin(t, manager, accountID, connectedKey, true, time.Now().UTC()) + + scheduled := make(chan context.Context, 1) + manager.peerLoginExpiry = &MockScheduler{ + IsSchedulerRunningFunc: func(string) bool { return false }, + ScheduleFunc: func(ctx context.Context, _ time.Duration, _ string, _ func() (time.Duration, bool)) { + scheduled <- ctx + }, + } + + requestCtx, cancel := context.WithCancel(context.Background()) + manager.schedulePeerLoginExpiration(requestCtx, accountID) + cancel() + + select { + case jobCtx := <-scheduled: + assert.NoError(t, jobCtx.Err(), "the expiration job must outlive the request that armed it") + case <-time.After(time.Second): + t.Fatal("timeout while waiting for the job to be scheduled") + } +} + +func TestDefaultAccountManager_ExpireAndUpdatePeers_SkipsPeerThatLoggedInAgain(t *testing.T) { + manager, updateManager, err := createManager(t) + require.NoError(t, err, "unable to create account manager") + + accountID, err := manager.GetAccountIDByUserID(context.Background(), auth.UserAuth{UserId: userID}) + require.NoError(t, err, "unable to create an account") + + reloggedKey, staleKey := addExpiringPeers(t, manager) + _, err = manager.UpdateAccountSettings(context.Background(), accountID, userID, &types.Settings{ + PeerLoginExpiration: time.Hour, + PeerLoginExpirationEnabled: true, + Extra: &types.ExtraSettings{}, + }) + require.NoError(t, err, "expecting to update account settings successfully but got error") + manager.peerLoginExpiry.CancelAll(context.Background()) + + expiredLogin := time.Now().UTC().Add(-2 * time.Hour) + setPeerLogin(t, manager, accountID, reloggedKey, true, expiredLogin) + setPeerLogin(t, manager, accountID, staleKey, true, expiredLogin) + + expiredPeers, err := manager.getExpiredPeers(context.Background(), accountID) + require.NoError(t, err) + require.Len(t, expiredPeers, 2, "both peers should be due for expiration") + + // The job holds the candidate list while one peer completes a fresh login, which + // moves its deadline into the future and must win over the stale candidate entry. + setPeerLogin(t, manager, accountID, reloggedKey, true, time.Now().UTC()) + + reloggedPeer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, reloggedKey) + require.NoError(t, err) + stalePeer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, staleKey) + require.NoError(t, err) + updateManager.CreateChannel(context.Background(), reloggedPeer.ID) + updateManager.CreateChannel(context.Background(), stalePeer.ID) + + err = manager.expireAndUpdatePeers(context.Background(), accountID, expiredPeers, peerExpirationSessionExpired) + require.NoError(t, err) + + reloggedPeer, err = manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, reloggedKey) + require.NoError(t, err) + assert.False(t, reloggedPeer.Status.LoginExpired, "a peer that logged in again must not be flagged from the stale candidate list") + assert.True(t, reloggedPeer.Status.Connected, "the re-logged peer must keep its connected status") + assert.True(t, updateManager.HasChannel(reloggedPeer.ID), "the re-logged peer's update channel must stay open") + + stalePeer, err = manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, staleKey) + require.NoError(t, err) + assert.True(t, stalePeer.Status.LoginExpired, "a peer that is still due must be flagged") + assert.False(t, updateManager.HasChannel(stalePeer.ID), "the expired peer's update channel must be closed") +} + +// addExpiringPeers registers two SSO peers with login expiration enabled and returns their public keys. +func addExpiringPeers(t *testing.T, manager *DefaultAccountManager) (string, string) { + t.Helper() + keys := make([]string, 0, 2) + for _, hostname := range []string{"connected-peer", "offline-peer"} { + key, err := wgtypes.GenerateKey() + require.NoError(t, err, "unable to generate WireGuard key") + _, _, _, _, err = manager.AddPeer(context.Background(), "", "", userID, &nbpeer.Peer{ + Key: key.PublicKey().String(), + Meta: nbpeer.PeerSystemMeta{Hostname: hostname}, + LoginExpirationEnabled: true, + }, false) + require.NoError(t, err, "unable to add peer") + keys = append(keys, key.PublicKey().String()) + } + return keys[0], keys[1] +} + +func setPeerLogin(t *testing.T, manager *DefaultAccountManager, accountID, peerKey string, connected bool, lastLogin time.Time) { + t.Helper() + peer, err := manager.Store.GetPeerByPeerPubKey(context.Background(), store.LockingStrengthNone, peerKey) + require.NoError(t, err) + peer.Status.Connected = connected + peer.LastLogin = &lastLogin + require.NoError(t, manager.Store.SavePeer(context.Background(), accountID, peer)) +} + func TestDefaultAccountManager_MarkPeerDisconnected_SchedulesInactivityExpiration(t *testing.T) { manager, _, err := createManager(t) require.NoError(t, err, "unable to create account manager") @@ -2702,7 +2850,7 @@ func TestAccount_GetNextPeerExpiration(t *testing.T) { expectedNextExpiration: time.Duration(0), }, { - name: "No connected peers, no expiration", + name: "Offline peer with expiration, return expiration", peers: map[string]*nbpeer.Peer{ "peer-1": { Status: &nbpeer.PeerStatus{ @@ -2721,8 +2869,33 @@ func TestAccount_GetNextPeerExpiration(t *testing.T) { }, expiration: time.Second, expirationEnabled: false, - expectedNextRun: false, - expectedNextExpiration: time.Duration(0), + expectedNextRun: true, + expectedNextExpiration: time.Second, + }, + { + name: "Offline peer with the earliest deadline defines the next run", + peers: map[string]*nbpeer.Peer{ + "peer-1": { + Status: &nbpeer.PeerStatus{ + Connected: true, + }, + LoginExpirationEnabled: true, + LastLogin: util.ToPtr(time.Now().UTC()), + UserID: userID, + }, + "peer-2": { + Status: &nbpeer.PeerStatus{ + Connected: false, + }, + LoginExpirationEnabled: true, + LastLogin: util.ToPtr(time.Now().UTC().Add(-50 * time.Minute)), + UserID: userID, + }, + }, + expiration: time.Hour, + expirationEnabled: true, + expectedNextRun: true, + expectedNextExpiration: 10 * time.Minute, }, { name: "Connected peers with disabled expiration, no expiration", diff --git a/management/server/peer.go b/management/server/peer.go index 07619f51e..9f5572252 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -1494,9 +1494,12 @@ func checkAuth(ctx context.Context, loginUserID string, peer *nbpeer.Peer) error func peerLoginExpired(ctx context.Context, peer *nbpeer.Peer, settings *types.Settings) bool { expired, expiresIn := peer.LoginExpired(settings.PeerLoginExpiration) - expired = settings.PeerLoginExpirationEnabled && expired - if expired || peer.Status.LoginExpired { - log.WithContext(ctx).Debugf("peer's %s login expired %v ago", peer.ID, expiresIn) + if settings.PeerLoginExpirationEnabled && expired { + log.WithContext(ctx).Debugf("peer's %s login expired %v ago", peer.ID, -expiresIn) + return true + } + if peer.Status.LoginExpired { + log.WithContext(ctx).Debugf("peer's %s login is marked as expired", peer.ID) return true } return false @@ -1643,7 +1646,9 @@ func (am *DefaultAccountManager) UpdateAccountPeer(ctx context.Context, accountI // getNextPeerExpiration returns the minimum duration in which the next peer of the account will expire if it was found. // If there is no peer that expires this function returns false and a duration of 0. -// This function only considers peers that haven't been expired yet and that are connected. +// This function only considers peers that haven't been expired yet. Offline peers count too: +// a running job is never re-armed on connect, so a peer that reconnects with an old login +// must already be part of the scheduled run. func (am *DefaultAccountManager) getNextPeerExpiration(ctx context.Context, accountID string) (time.Duration, bool) { peersWithExpiry, err := am.Store.GetAccountPeersWithExpiration(ctx, store.LockingStrengthNone, accountID) if err != nil { @@ -1663,8 +1668,7 @@ func (am *DefaultAccountManager) getNextPeerExpiration(ctx context.Context, acco var nextExpiry *time.Duration for _, peer := range peersWithExpiry { - // consider only connected peers because others will require login on connecting to the management server - if peer.Status.LoginExpired || !peer.Status.Connected { + if peer.Status.LoginExpired { continue } _, duration := peer.LoginExpired(settings.PeerLoginExpiration) diff --git a/management/server/scheduler.go b/management/server/scheduler.go index b61643295..1daea4295 100644 --- a/management/server/scheduler.go +++ b/management/server/scheduler.go @@ -117,6 +117,7 @@ func (wm *DefaultScheduler) Schedule(ctx context.Context, in time.Duration, ID s } ticker := time.NewTicker(in) + period := in wm.jobs[ID] = cancel log.WithContext(ctx).Debugf("scheduled a job %s to run in %s. There are %d total jobs scheduled.", ID, in.String(), len(wm.jobs)) @@ -136,14 +137,18 @@ func (wm *DefaultScheduler) Schedule(ctx context.Context, in time.Duration, ID s if !reschedule { wm.mu.Lock() defer wm.mu.Unlock() - delete(wm.jobs, ID) + // A Cancel during job() may have registered a replacement under this ID. + if current, ok := wm.jobs[ID]; ok && current == cancel { + delete(wm.jobs, ID) + } log.WithContext(ctx).Debugf("job %s is not scheduled to run again", ID) ticker.Stop() return } // we need this comparison to avoid resetting the ticker with the same duration and missing the current elapsesed time - if runIn != in { + if runIn != period { ticker.Reset(runIn) + period = runIn } case <-cancel: log.WithContext(ctx).Debugf("job %s was canceled, stopping timer", ID) diff --git a/management/server/scheduler_test.go b/management/server/scheduler_test.go index e3af551ad..9dd13ce6b 100644 --- a/management/server/scheduler_test.go +++ b/management/server/scheduler_test.go @@ -6,10 +6,12 @@ import ( "math/rand" "runtime" "sync" + "sync/atomic" "testing" "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestScheduler_Performance(t *testing.T) { @@ -150,3 +152,90 @@ func TestScheduler_Schedule(t *testing.T) { scheduler.cancel(context.Background(), jobID) } + +func TestScheduler_Schedule_ResetsTickerAfterReturningInitialInterval(t *testing.T) { + jobID := "test-scheduler-job-2" + scheduler := NewDefaultScheduler() + defer scheduler.Cancel(context.Background(), []string{jobID}) + + initial := 30 * time.Millisecond + stretched := 400 * time.Millisecond + runs := make(chan time.Time, 3) + count := 0 + // The first run stretches the period; the second returns the initial interval again, + // which must shrink the period back instead of keeping the stretched one. + job := func() (nextRunIn time.Duration, reschedule bool) { + count++ + runs <- time.Now() + switch count { + case 1: + return stretched, true + case 2: + return initial, true + default: + return 0, false + } + } + scheduler.Schedule(context.Background(), initial, jobID, job) + + var stamps []time.Time + for len(stamps) < 3 { + select { + case ts := <-runs: + stamps = append(stamps, ts) + case <-time.After(2 * time.Second): + t.Fatalf("timed out after %d runs", len(stamps)) + } + } + assert.Less(t, stamps[2].Sub(stamps[1]), stretched/2, "returning the initial interval must reset the stretched ticker") +} + +func TestScheduler_Schedule_StaleCompletionKeepsReplacement(t *testing.T) { + jobID := "test-scheduler-job-3" + scheduler := NewDefaultScheduler() + defer scheduler.Cancel(context.Background(), []string{jobID}) + + started := make(chan struct{}) + release := make(chan struct{}) + staleJob := func() (nextRunIn time.Duration, reschedule bool) { + close(started) + <-release + return 0, false + } + scheduler.Schedule(context.Background(), 10*time.Millisecond, jobID, staleJob) + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("timed out waiting for the first job to start") + } + + // Cancel the job while it is still executing and register a replacement under the + // same ID, as the expiration paths do on a settings change. + scheduler.Cancel(context.Background(), []string{jobID}) + var replacementRuns atomic.Int32 + scheduler.Schedule(context.Background(), 20*time.Millisecond, jobID, func() (nextRunIn time.Duration, reschedule bool) { + replacementRuns.Add(1) + return 20 * time.Millisecond, true + }) + require.True(t, scheduler.IsSchedulerRunning(jobID), "replacement must be registered") + + // The stale job now completes without rescheduling; its cleanup must leave the + // replacement's entry in place. + close(release) + assert.Never(t, func() bool { return !scheduler.IsSchedulerRunning(jobID) }, 200*time.Millisecond, 10*time.Millisecond, + "stale completion must not drop the replacement job") + + var duplicateRuns atomic.Int32 + scheduler.Schedule(context.Background(), 10*time.Millisecond, jobID, func() (nextRunIn time.Duration, reschedule bool) { + duplicateRuns.Add(1) + return 10 * time.Millisecond, true + }) + assert.Never(t, func() bool { return duplicateRuns.Load() > 0 }, 100*time.Millisecond, 10*time.Millisecond, + "a duplicate schedule must be refused while the replacement is registered") + + scheduler.Cancel(context.Background(), []string{jobID}) + assert.False(t, scheduler.IsSchedulerRunning(jobID), "cancel must find and remove the replacement") + runsAfterCancel := replacementRuns.Load() + assert.Never(t, func() bool { return replacementRuns.Load() > runsAfterCancel+1 }, 150*time.Millisecond, 10*time.Millisecond, + "the replacement must stop after cancel") +} diff --git a/management/server/types/account.go b/management/server/types/account.go index d689b0175..d0688d1ee 100644 --- a/management/server/types/account.go +++ b/management/server/types/account.go @@ -404,7 +404,7 @@ func (a *Account) GetExpiredPeers() []*nbpeer.Peer { // GetNextPeerExpiration returns the minimum duration in which the next peer of the account will expire if it was found. // If there is no peer that expires this function returns false and a duration of 0. -// This function only considers peers that haven't been expired yet and that are connected. +// This function only considers peers that haven't been expired yet, whether connected or not. func (a *Account) GetNextPeerExpiration() (time.Duration, bool) { peersWithExpiry := a.GetPeersWithExpiration() if len(peersWithExpiry) == 0 { @@ -412,8 +412,7 @@ func (a *Account) GetNextPeerExpiration() (time.Duration, bool) { } var nextExpiry *time.Duration for _, peer := range peersWithExpiry { - // consider only connected peers because others will require login on connecting to the management server - if peer.Status.LoginExpired || !peer.Status.Connected { + if peer.Status.LoginExpired { continue } _, duration := peer.LoginExpired(a.Settings.PeerLoginExpiration) diff --git a/management/server/user.go b/management/server/user.go index 0a711389a..823c1b2e4 100644 --- a/management/server/user.go +++ b/management/server/user.go @@ -1177,28 +1177,35 @@ func (am *DefaultAccountManager) expireAndUpdatePeers(ctx context.Context, accou dnsDomain := am.networkMapController.GetDNSDomain(settings) var peerIDs []string - for _, peer := range peers { + defer func() { + if len(peerIDs) == 0 { + return + } + // this will trigger peer disconnect from the management service + log.Debugf("Expiring %d peers for account %s", len(peerIDs), accountID) + am.networkMapController.DisconnectPeers(ctx, accountID, peerIDs) + }() + for _, candidate := range peers { // nolint:staticcheck - ctx = context.WithValue(ctx, nbcontext.PeerIDKey, peer.Key) + peerCtx := context.WithValue(ctx, nbcontext.PeerIDKey, candidate.Key) - if peer.UserID == "" { + if candidate.UserID == "" { // we do not want to expire peers that are added via setup key continue } - if peer.Status.LoginExpired { + peer, err := am.expirePeerIfStillDue(peerCtx, accountID, candidate.ID, settings, reason) + if err != nil { + return err + } + if peer == nil { continue } peerIDs = append(peerIDs, peer.ID) - peer.MarkLoginExpired(true) - - if err := am.Store.SavePeerStatus(ctx, accountID, peer.ID, *peer.Status); err != nil { - return err - } meta := peer.EventMeta(dnsDomain) meta["reason"] = string(reason) am.StoreEvent( - ctx, + peerCtx, peer.UserID, peer.ID, accountID, activity.PeerLoginExpired, meta, ) @@ -1215,15 +1222,53 @@ func (am *DefaultAccountManager) expireAndUpdatePeers(ctx context.Context, accou if err != nil { return fmt.Errorf("notify network map controller of peer update: %w", err) } - - if len(peerIDs) != 0 { - // this will trigger peer disconnect from the management service - log.Debugf("Expiring %d peers for account %s", len(peerIDs), accountID) - am.networkMapController.DisconnectPeers(ctx, accountID, peerIDs) - } return nil } +// expirePeerIfStillDue flags the peer as login-expired and returns its fresh copy, or nil +// when it no longer qualifies. The candidate list is read without a lock, so a login that +// landed in between would otherwise be overwritten with a stale expired status. +func (am *DefaultAccountManager) expirePeerIfStillDue(ctx context.Context, accountID, peerID string, settings *types.Settings, reason peerExpirationReason) (*nbpeer.Peer, error) { + var expired *nbpeer.Peer + err := am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + peer, err := transaction.GetPeerByID(ctx, store.LockingStrengthUpdate, accountID, peerID) + if err != nil { + if s, ok := status.FromError(err); ok && s.Type() == status.NotFound { + return nil + } + return err + } + if peer.Status.LoginExpired || !peerExpirationDue(peer, settings, reason) { + return nil + } + peer.MarkLoginExpired(true) + if err := transaction.SavePeerStatus(ctx, accountID, peer.ID, *peer.Status); err != nil { + return err + } + expired = peer + return nil + }) + if err != nil { + return nil, err + } + return expired, nil +} + +// peerExpirationDue re-evaluates a time-based expiry against the peer's current state. +// Administrative reasons expire the peer unconditionally. +func peerExpirationDue(peer *nbpeer.Peer, settings *types.Settings, reason peerExpirationReason) bool { + switch reason { + case peerExpirationSessionExpired: + expired, _ := peer.LoginExpired(settings.PeerLoginExpiration) + return settings.PeerLoginExpirationEnabled && expired + case peerExpirationInactivity: + expired, _ := peer.SessionExpired(settings.PeerInactivityExpiration) + return settings.PeerInactivityExpirationEnabled && expired + default: + return true + } +} + func (am *DefaultAccountManager) deleteUserFromIDP(ctx context.Context, targetUserID, accountID string) error { if am.userDeleteFromIDPEnabled { log.WithContext(ctx).Debugf("user %s deleted from IdP", targetUserID) From 27991aab984e5aa8fc19a307b13ea4ed0f1dc6e4 Mon Sep 17 00:00:00 2001 From: Brad Ison Date: Wed, 9 Sep 2026 15:36:37 +0200 Subject: [PATCH 4/9] [management] Let embedding binaries extend the command tree (#7483) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The management binary is embedded by downstream builds that override server construction via SetNewServer, but the cobra command tree itself was closed: rootCmd is unexported and fully assembled in init, with no way to attach additional subcommands. Customize hands the built root command to a caller-supplied function before Execute, so an embedding binary can add its own commands next to — or under — the built-in ones, such as extra administrative helpers beneath the existing admin group. --- management/cmd/root.go | 9 ++++++++ management/cmd/root_test.go | 42 +++++++++++++++++++++++++++++++++++++ 2 files changed, 51 insertions(+) create mode 100644 management/cmd/root_test.go diff --git a/management/cmd/root.go b/management/cmd/root.go index 969dd60dd..ae03a09e8 100644 --- a/management/cmd/root.go +++ b/management/cmd/root.go @@ -54,6 +54,15 @@ func Execute() error { return rootCmd.Execute() } +// Customize hands the fully built root command to fn so an embedding binary +// can extend or adjust the command tree — most commonly attaching its own +// subcommands next to (or under) the built-in ones — before calling Execute. +// The root command is constructed in this package's init, so Customize may be +// called from the embedding binary's main at any point before Execute. +func Customize(fn func(root *cobra.Command)) { + fn(rootCmd) +} + func init() { mgmtCmd.Flags().IntVar(&mgmtPort, "port", 80, "server port to listen on (defaults to 443 if TLS is enabled, 80 otherwise") mgmtCmd.Flags().BoolVar(&disableLegacyManagementPort, "disable-legacy-port", false, "disabling the old legacy port (33073)") diff --git a/management/cmd/root_test.go b/management/cmd/root_test.go new file mode 100644 index 000000000..826fd2d50 --- /dev/null +++ b/management/cmd/root_test.go @@ -0,0 +1,42 @@ +package cmd + +import ( + "testing" + + "github.com/spf13/cobra" +) + +// TestCustomize verifies an embedding binary can extend the command tree: a +// top-level command attached through the hook, and a subcommand attached under +// the built-in admin group, are both resolvable exactly as Execute would +// resolve them. +func TestCustomize(t *testing.T) { + topLevel := &cobra.Command{Use: "some-extra", RunE: func(*cobra.Command, []string) error { return nil }} + nested := &cobra.Command{Use: "cluster", RunE: func(*cobra.Command, []string) error { return nil }} + + Customize(func(root *cobra.Command) { + root.AddCommand(topLevel) + for _, c := range root.Commands() { + if c.Name() == "admin" { + c.AddCommand(nested) + return + } + } + t.Fatal("admin command not found in the root tree") + }) + t.Cleanup(func() { + rootCmd.RemoveCommand(topLevel) + for _, c := range rootCmd.Commands() { + if c.Name() == "admin" { + c.RemoveCommand(nested) + } + } + }) + + if found, _, err := rootCmd.Find([]string{"some-extra"}); err != nil || found != topLevel { + t.Fatalf("top-level command not resolvable: found=%v err=%v", found, err) + } + if found, _, err := rootCmd.Find([]string{"admin", "cluster"}); err != nil || found != nested { + t.Fatalf("nested admin subcommand not resolvable: found=%v err=%v", found, err) + } +} From 21b4a83cea05cb2a3a54d2357c875f85ec2dd10f Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Thu, 10 Sep 2026 11:57:14 +0200 Subject: [PATCH 5/9] [management] Refuse services on unvalidated custom domains (#7341) Require validated custom domains when creating or updating reverse proxy services. Propagate validation errors during updates and return HTTP 409 for duplicate domain claims. Add regression tests for domain validation, ownership, and service creation and updates. --- .../domain/manager/domain_test.go | 50 ++- .../reverseproxy/domain/manager/manager.go | 64 +++- .../domain/manager/manager_realstore_test.go | 326 ++++++++++++++++++ .../domain/manager/manager_test.go | 4 + .../service/manager/domain_validation_test.go | 127 +++++++ .../reverseproxy/service/manager/manager.go | 19 +- management/server/store/sql_store.go | 29 ++ management/server/store/store.go | 1 + management/server/store/store_mock.go | 15 + 9 files changed, 612 insertions(+), 23 deletions(-) create mode 100644 management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go create mode 100644 management/internals/modules/reverseproxy/service/manager/domain_validation_test.go diff --git a/management/internals/modules/reverseproxy/domain/manager/domain_test.go b/management/internals/modules/reverseproxy/domain/manager/domain_test.go index 523920a99..38d5a923b 100644 --- a/management/internals/modules/reverseproxy/domain/manager/domain_test.go +++ b/management/internals/modules/reverseproxy/domain/manager/domain_test.go @@ -66,8 +66,8 @@ func TestExtractClusterFromFreeDomain(t *testing.T) { func TestExtractClusterFromCustomDomains(t *testing.T) { customDomains := []*domain.Domain{ - {Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io"}, - {Domain: "proxy.corp.io", TargetCluster: "us1.proxy.netbird.io"}, + {Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io", Validated: true}, + {Domain: "proxy.corp.io", TargetCluster: "us1.proxy.netbird.io", Validated: true}, } tests := []struct { @@ -120,19 +120,49 @@ func TestExtractClusterFromCustomDomains(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - cluster, ok := extractClusterFromCustomDomains(tc.domain, customDomains) - assert.Equal(t, tc.wantOK, ok) - if ok { - assert.Equal(t, tc.wantVal, cluster) + cluster, match := extractClusterFromCustomDomains(tc.domain, customDomains) + if !tc.wantOK { + assert.Equal(t, customDomainNoMatch, match, "unrelated domain should not match any custom domain") + return } + assert.Equal(t, customDomainValidated, match, "validated custom domain should resolve a cluster") + assert.Equal(t, tc.wantVal, cluster) }) } } +// An unvalidated row must never yield a cluster: the account has not shown it +// controls the name, so no service may be bound to it. +func TestExtractClusterFromCustomDomains_UnvalidatedDomainRefused(t *testing.T) { + customDomains := []*domain.Domain{ + {Domain: "example.com", TargetCluster: "eu1.proxy.netbird.io", Validated: false}, + } + + for _, serviceDomain := range []string{"example.com", "app.example.com"} { + t.Run(serviceDomain, func(t *testing.T) { + cluster, match := extractClusterFromCustomDomains(serviceDomain, customDomains) + assert.Equal(t, customDomainUnvalidated, match, "unvalidated row must be reported as such") + assert.Empty(t, cluster, "unvalidated row must not resolve a cluster") + }) + } +} + +// A more specific unvalidated row must not shadow a validated parent domain. +func TestExtractClusterFromCustomDomains_ValidatedParentWinsOverUnvalidatedChild(t *testing.T) { + customDomains := []*domain.Domain{ + {Domain: "example.com", TargetCluster: "cluster-generic", Validated: true}, + {Domain: "app.example.com", TargetCluster: "cluster-app", Validated: false}, + } + + cluster, match := extractClusterFromCustomDomains("app.example.com", customDomains) + assert.Equal(t, customDomainValidated, match) + assert.Equal(t, "cluster-generic", cluster, "validated parent domain should provide the cluster") +} + func TestExtractClusterFromCustomDomains_OverlappingDomains(t *testing.T) { customDomains := []*domain.Domain{ - {Domain: "example.com", TargetCluster: "cluster-generic"}, - {Domain: "app.example.com", TargetCluster: "cluster-app"}, + {Domain: "example.com", TargetCluster: "cluster-generic", Validated: true}, + {Domain: "app.example.com", TargetCluster: "cluster-app", Validated: true}, } tests := []struct { @@ -164,8 +194,8 @@ func TestExtractClusterFromCustomDomains_OverlappingDomains(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { - cluster, ok := extractClusterFromCustomDomains(tc.domain, customDomains) - assert.True(t, ok) + cluster, match := extractClusterFromCustomDomains(tc.domain, customDomains) + assert.Equal(t, customDomainValidated, match) assert.Equal(t, tc.wantVal, cluster) }) } diff --git a/management/internals/modules/reverseproxy/domain/manager/manager.go b/management/internals/modules/reverseproxy/domain/manager/manager.go index a9774d0e9..46e4ced83 100644 --- a/management/internals/modules/reverseproxy/domain/manager/manager.go +++ b/management/internals/modules/reverseproxy/domain/manager/manager.go @@ -26,6 +26,7 @@ type store interface { GetAgentNetworkSettings(ctx context.Context, lockStrength nbstore.LockingStrength, accountID string) (*agentnetworkTypes.Settings, error) GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error) + GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) ListFreeDomains(ctx context.Context, accountID string) ([]string, error) ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error) @@ -150,6 +151,10 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName return nil, fmt.Errorf("target cluster %s is not available", targetCluster) } + if err := m.checkDomainAvailable(ctx, domainName); err != nil { + return nil, err + } + // Attempt an initial validation against the specified cluster only var validated bool if m.validator.IsValid(ctx, domainName, []string{targetCluster}) { @@ -166,6 +171,23 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName return d, nil } +// checkDomainAvailable reports whether the domain is free to claim. The unique +// index on the column is the real guard; this turns the violation into a +// conflict the caller can act on instead of a database error, and says nothing +// about which account holds the domain. +func (m Manager) checkDomainAvailable(ctx context.Context, domainName string) error { + _, err := m.store.GetCustomDomainByName(ctx, domainName) + if err == nil { + return status.Errorf(status.AlreadyExists, "domain %s is already registered", domainName) + } + + if sErr, ok := status.FromError(err); ok && sErr.Type() == status.NotFound { + return nil + } + + return fmt.Errorf("look up domain: %w", err) +} + func (m Manager) DeleteDomain(ctx context.Context, accountID, userID, domainID string) error { ok, ctx, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Services, operations.Delete) if err != nil { @@ -203,7 +225,9 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID log.WithFields(log.Fields{ "accountID": accountID, "domainID": domainID, - }).WithError(err).Error("validate domain") + "userID": userID, + }).Error("validate domain: permission denied") + return } log.WithFields(log.Fields{ @@ -298,9 +322,12 @@ func (m Manager) DeriveClusterFromDomain(ctx context.Context, accountID, domain return "", fmt.Errorf("list custom domains: %w", err) } - targetCluster, valid := extractClusterFromCustomDomains(domain, customDomains) - if valid { + targetCluster, match := extractClusterFromCustomDomains(domain, customDomains) + switch match { + case customDomainValidated: return targetCluster, nil + case customDomainUnvalidated: + return "", status.Errorf(status.PreconditionFailed, "domain %s is not validated", domain) } return "", fmt.Errorf("domain %s does not match any available proxy cluster", domain) @@ -363,19 +390,46 @@ func (m Manager) reservedGatewayAddress(ctx context.Context, accountID string) ( return settings.ProxyAddress, nil } -func extractClusterFromCustomDomains(serviceDomain string, customDomains []*domain.Domain) (string, bool) { +// customDomainMatch describes how a service domain relates to the account's +// custom domain rows. +type customDomainMatch int + +const ( + customDomainNoMatch customDomainMatch = iota + customDomainUnvalidated + customDomainValidated +) + +// extractClusterFromCustomDomains finds the longest custom domain covering the +// service domain and reports its target cluster. Only a validated row yields a +// cluster: until the CNAME check has passed the account has not shown it +// controls the name, so no traffic may be routed for it. +func extractClusterFromCustomDomains(serviceDomain string, customDomains []*domain.Domain) (string, customDomainMatch) { bestCluster := "" bestLen := -1 + matched := false for _, cd := range customDomains { if serviceDomain != cd.Domain && !strings.HasSuffix(serviceDomain, "."+cd.Domain) { continue } + matched = true + if !cd.Validated { + continue + } if l := len(cd.Domain); l > bestLen { bestLen = l bestCluster = cd.TargetCluster } } - return bestCluster, bestLen >= 0 + + switch { + case bestLen >= 0: + return bestCluster, customDomainValidated + case matched: + return "", customDomainUnvalidated + default: + return "", customDomainNoMatch + } } // ExtractClusterFromFreeDomain extracts the cluster address from a free domain. diff --git a/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go b/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go new file mode 100644 index 000000000..8a0b56171 --- /dev/null +++ b/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go @@ -0,0 +1,326 @@ +package manager + +import ( + "context" + "fmt" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/metric/noop" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager" + "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/mock_server" + "github.com/netbirdio/netbird/management/server/permissions" + nbstore "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + "github.com/netbirdio/netbird/shared/management/status" +) + +const ( + testCluster = "eu.proxy.test" + accountA = "account-a" + accountAUser = "account-a-admin" + accountB = "account-b" + accountBUser = "account-b-admin" + accountAMember = "account-a-member" +) + +// stubResolver answers CNAME lookups from a table the test controls, so a +// domain can point at the cluster or nowhere without touching a real resolver. +type stubResolver struct { + mu sync.Mutex + cnames map[string]string +} + +func (r *stubResolver) LookupCNAME(_ context.Context, host string) (string, error) { + r.mu.Lock() + defer r.mu.Unlock() + + cname, ok := r.cnames[host] + if !ok { + return "", fmt.Errorf("lookup %s: no such host", host) + } + return cname + ".", nil +} + +func (r *stubResolver) set(host, cname string) { + r.mu.Lock() + defer r.mu.Unlock() + r.cnames[host] = cname +} + +type domainTestEnv struct { + manager Manager + store nbstore.Store + resolver *stubResolver +} + +// setupDomainTest builds the domain manager on a real SQLite store with two +// accounts and one active public proxy cluster. +func setupDomainTest(t *testing.T) *domainTestEnv { + t.Helper() + + ctx := context.Background() + testStore, cleanup, err := nbstore.NewTestStoreFromSQL(ctx, "", t.TempDir()) + require.NoError(t, err) + t.Cleanup(cleanup) + + for accountID, userID := range map[string]string{accountA: accountAUser, accountB: accountBUser} { + users := map[string]*types.User{ + userID: { + Id: userID, + AccountID: accountID, + Role: types.UserRoleAdmin, + }, + } + if accountID == accountA { + // A real member of the account whose role denies Services:Create, so + // permission denial is exercised as ok=false rather than as a lookup + // error for a user who is not in the account at all. + users[accountAMember] = &types.User{ + Id: accountAMember, + AccountID: accountID, + Role: types.UserRoleUser, + } + } + + require.NoError(t, testStore.SaveAccount(ctx, &types.Account{ + Id: accountID, + CreatedBy: userID, + Settings: &types.Settings{}, + Users: users, + })) + } + + proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter("")) + require.NoError(t, err) + + _, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", nil, nil) + require.NoError(t, err) + + resolver := &stubResolver{cnames: make(map[string]string)} + + mgr := Manager{ + store: testStore, + proxyManager: proxyMgr, + validator: domain.Validator{Resolver: resolver}, + permissionsManager: permissions.NewManager(testStore), + accountManager: &mock_server.MockAccountManager{ + StoreEventFunc: func(context.Context, string, string, string, activity.ActivityDescriber, map[string]any) {}, + }, + } + + return &domainTestEnv{manager: mgr, store: testStore, resolver: resolver} +} + +// storedDomain reads a domain row back through the store so assertions are made +// on what was persisted rather than on the value the manager returned. +func storedDomain(t *testing.T, s nbstore.Store, accountID, domainName string) *domain.Domain { + t.Helper() + + domains, err := s.ListCustomDomains(context.Background(), accountID) + require.NoError(t, err) + for _, d := range domains { + if d.Domain == domainName { + return d + } + } + return nil +} + +// A domain whose CNAME check fails is stored unvalidated and must not resolve a +// cluster, which is what service creation gates on. +func TestCreateDomain_FailedLookupIsNotServable(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + + created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "apps.example.com", testCluster) + require.NoError(t, err) + assert.False(t, created.Validated, "a domain whose CNAME lookup fails must not be created validated") + + stored := storedDomain(t, env.store, accountA, "apps.example.com") + require.NotNil(t, stored, "domain row should exist") + assert.False(t, stored.Validated, "persisted row must be unvalidated") + + cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "apps.example.com") + require.Error(t, err, "an unvalidated domain must not resolve a cluster") + assert.Empty(t, cluster) + assert.Contains(t, err.Error(), "not validated", "error should tell the caller what to fix") + + sErr, ok := status.FromError(err) + require.True(t, ok, "error should be a typed status error") + assert.Equal(t, status.PreconditionFailed, sErr.Type()) + + _, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "sub.apps.example.com") + assert.Error(t, err, "subdomains of an unvalidated custom domain are not servable either") +} + +// A second account claiming a registered domain gets a clean conflict, not a +// database error surfaced as a 500. +func TestCreateDomain_DuplicateIsAConflict(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + + _, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "shared.example.com", testCluster) + require.NoError(t, err) + + _, err = env.manager.CreateDomain(ctx, accountB, accountBUser, "shared.example.com", testCluster) + require.Error(t, err) + + sErr, ok := status.FromError(err) + require.True(t, ok, "conflict must be a typed status error, not a raw database error") + assert.Equal(t, status.AlreadyExists, sErr.Type(), "conflict should map to 409, not 500") + assert.NotContains(t, sErr.Message, accountA, "the response must not reveal the holding account") + + assert.Nil(t, storedDomain(t, env.store, accountB, "shared.example.com"), "no row should be written on conflict") +} + +// The same account re-adding one of its own domains is a conflict too. +func TestCreateDomain_SameAccountDuplicateIsAConflict(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + + _, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "dup.example.com", testCluster) + require.NoError(t, err) + + _, err = env.manager.CreateDomain(ctx, accountA, accountAUser, "dup.example.com", testCluster) + require.Error(t, err) + + sErr, ok := status.FromError(err) + require.True(t, ok) + assert.Equal(t, status.AlreadyExists, sErr.Type()) +} + +// The negative control: a validated domain still derives its cluster, for the +// bare name and for subdomains, exactly as before. +func TestCreateDomain_ValidatedDomainDerivesCluster(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + env.resolver.set("validation.valid.example.com", testCluster) + + created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "valid.example.com", testCluster) + require.NoError(t, err) + require.True(t, created.Validated, "a matching CNAME should validate on create") + + cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "valid.example.com") + require.NoError(t, err) + assert.Equal(t, testCluster, cluster) + + cluster, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "app.valid.example.com") + require.NoError(t, err) + assert.Equal(t, testCluster, cluster, "subdomains of a validated custom domain resolve too") +} + +// Validating a domain flips the gate: the same lookup that failed before now +// resolves a cluster. +func TestValidateDomain_UnlocksClusterDerivation(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + + created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "later.example.com", testCluster) + require.NoError(t, err) + require.False(t, created.Validated) + + _, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "later.example.com") + require.Error(t, err) + + env.resolver.set("validation.later.example.com", testCluster) + env.manager.ValidateDomain(ctx, accountA, accountAUser, created.ID) + + require.True(t, storedDomain(t, env.store, accountA, "later.example.com").Validated) + + cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "later.example.com") + require.NoError(t, err) + assert.Equal(t, testCluster, cluster) +} + +// Free cluster domains are unaffected by the custom domain gate. +func TestDeriveClusterFromDomain_FreeDomainUnaffected(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + + cluster, err := env.manager.DeriveClusterFromDomain(ctx, accountA, "myapp.abc123."+testCluster) + require.NoError(t, err) + assert.Equal(t, testCluster, cluster) +} + +// The manager pre-check exists to turn a conflict into a 409, but the unique +// index on the column is what actually guarantees the domain is claimed once. +// +// Two requests can clear the pre-check concurrently and race to the insert. +// Inserting twice through the store reaches the same code path the loser of +// that race takes, without the nondeterminism of driving it from goroutines, +// and the loser must still see a conflict rather than an internal error. +func TestStore_DuplicateDomainRejectedByIndexAsConflict(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + + _, err := env.store.CreateCustomDomain(ctx, accountA, "indexed.example.com", testCluster, false) + require.NoError(t, err) + + _, err = env.store.CreateCustomDomain(ctx, accountB, "indexed.example.com", testCluster, false) + require.Error(t, err, "the unique index must reject the same domain in a second account") + + sErr, ok := status.FromError(err) + require.True(t, ok, "the losing insert must return a typed status error") + assert.Equal(t, status.AlreadyExists, sErr.Type(), "a lost race is a 409, not a 500") +} + +// Validation is what decides whether a domain routes traffic, so a caller +// without permission to it must not be able to flip the flag. The check logged +// the denial and then carried on, which was inert while nothing read Validated +// and is not once cluster derivation gates on it. +func TestValidateDomain_PermissionDeniedDoesNotValidate(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + + created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "guarded.example.com", testCluster) + require.NoError(t, err) + require.False(t, created.Validated) + + // The CNAME is in place, so the only thing standing between this caller and + // a validated domain is the permission check. + env.resolver.set("validation.guarded.example.com", testCluster) + + env.manager.ValidateDomain(ctx, accountA, accountAMember, created.ID) + + stored := storedDomain(t, env.store, accountA, "guarded.example.com") + require.NotNil(t, stored) + assert.False(t, stored.Validated, "a caller without permission must not validate the domain") + + _, err = env.manager.DeriveClusterFromDomain(ctx, accountA, "guarded.example.com") + assert.Error(t, err, "the domain must still be unservable") +} + +// Validation runs asynchronously, so it can finish after the domain was +// deleted and then write a stale row back. gorm's Save falls back to an insert +// when an update affects no rows, which would resurrect the domain as +// validated; UpdateCustomDomain avoids that by selecting explicit columns. +// This pins that behaviour, since dropping the Select would reintroduce it. +func TestUpdateCustomDomain_DoesNotResurrectDeletedDomain(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + + created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "racy.example.com", testCluster) + require.NoError(t, err) + + stale := storedDomain(t, env.store, accountA, "racy.example.com") + require.NotNil(t, stale) + + require.NoError(t, env.manager.DeleteDomain(ctx, accountA, accountAUser, created.ID)) + require.Nil(t, storedDomain(t, env.store, accountA, "racy.example.com"), "the domain should be gone") + + // What an in-flight validation would write once its CNAME check succeeded. + // The write has to succeed for the assertion below to mean anything: a + // rejected write would leave the domain absent for the wrong reason. + stale.Validated = true + _, err = env.store.UpdateCustomDomain(ctx, accountA, stale) + require.NoError(t, err, "the update itself must succeed, so absence is not just a failed write") + + assert.Nil(t, storedDomain(t, env.store, accountA, "racy.example.com"), + "a late validation write must not recreate a deleted domain") +} diff --git a/management/internals/modules/reverseproxy/domain/manager/manager_test.go b/management/internals/modules/reverseproxy/domain/manager/manager_test.go index 12281b447..519f5efeb 100644 --- a/management/internals/modules/reverseproxy/domain/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/domain/manager/manager_test.go @@ -184,6 +184,10 @@ func (s *stubStore) GetCustomDomain(context.Context, string, string) (*domain.Do panic("not used in allow-list tests") } +func (s *stubStore) GetCustomDomainByName(context.Context, string) (*domain.Domain, error) { + panic("not used in allow-list tests") +} + func (s *stubStore) ListFreeDomains(context.Context, string) ([]string, error) { panic("not used in allow-list tests") } diff --git a/management/internals/modules/reverseproxy/service/manager/domain_validation_test.go b/management/internals/modules/reverseproxy/service/manager/domain_validation_test.go new file mode 100644 index 000000000..ccb955cd8 --- /dev/null +++ b/management/internals/modules/reverseproxy/service/manager/domain_validation_test.go @@ -0,0 +1,127 @@ +package manager + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/metric/noop" + + domainmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain/manager" + proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/mock_server" + "github.com/netbirdio/netbird/management/server/permissions" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/shared/management/status" +) + +const validationTestCluster = "eu.proxy.test" + +// withRealDomainManager swaps the stub cluster deriver for the real domain +// manager backed by the same store, so service creation is gated by the actual +// domain rows rather than by a test double that always agrees. +func withRealDomainManager(t *testing.T, mgr *Manager, testStore store.Store) { + t.Helper() + + ctx := context.Background() + proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter("")) + require.NoError(t, err) + + _, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", nil, nil) + require.NoError(t, err) + + accountMgr := &mock_server.MockAccountManager{ + StoreEventFunc: func(context.Context, string, string, string, activity.ActivityDescriber, map[string]any) {}, + } + mgr.clusterDeriver = domainmanager.NewManager(testStore, proxyMgr, permissions.NewManager(testStore), accountMgr) +} + +func newTestService(domain string) *rpservice.Service { + return &rpservice.Service{ + Name: "test-service", + Domain: domain, + Enabled: true, + Mode: rpservice.ModeHTTP, + Targets: []*rpservice.Target{{ + Host: "10.0.0.1", + Port: 8080, + Protocol: "http", + TargetId: testPeerID, + TargetType: "peer", + Enabled: true, + }}, + } +} + +// A service must not bind to a domain the account has not validated, and +// nothing may be persisted for the attempt. +func TestCreateService_RefusesUnvalidatedDomain(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupIntegrationTest(t) + withRealDomainManager(t, mgr, testStore) + + _, err := testStore.CreateCustomDomain(ctx, testAccountID, "unproven.example.com", validationTestCluster, false) + require.NoError(t, err) + + _, err = mgr.CreateService(ctx, testAccountID, testUserID, newTestService("unproven.example.com")) + require.Error(t, err, "an unvalidated domain must not bind a service") + assert.Contains(t, err.Error(), "not validated", "the API error should name the actual problem") + + sErr, ok := status.FromError(err) + require.True(t, ok, "error should be a typed status error") + assert.Equal(t, status.PreconditionFailed, sErr.Type()) + + services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, err) + assert.Empty(t, services, "no service row should be written for a refused domain") +} + +// The negative control: a validated domain still binds a service and derives +// its cluster exactly as before. +func TestCreateService_ValidatedDomainBindsService(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupIntegrationTest(t) + withRealDomainManager(t, mgr, testStore) + + _, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true) + require.NoError(t, err) + + created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.proven.example.com")) + require.NoError(t, err) + assert.Equal(t, validationTestCluster, created.ProxyCluster, "service should bind to the domain's target cluster") + + services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, err) + require.Len(t, services, 1, "the service should be persisted") + assert.Equal(t, "app.proven.example.com", services[0].Domain) +} + +// An update must not be a way around the creation gate: moving a live service +// onto an unvalidated domain has to fail rather than silently keep the old +// cluster and start serving the new hostname. +func TestUpdateService_RefusesMoveToUnvalidatedDomain(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupIntegrationTest(t) + withRealDomainManager(t, mgr, testStore) + + _, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true) + require.NoError(t, err) + _, err = testStore.CreateCustomDomain(ctx, testAccountID, "unproven.example.com", validationTestCluster, false) + require.NoError(t, err) + + created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.proven.example.com")) + require.NoError(t, err) + + moved := *created + moved.Domain = "app.unproven.example.com" + _, err = mgr.UpdateService(ctx, testAccountID, testUserID, &moved) + require.Error(t, err, "moving to an unvalidated domain must fail") + assert.Contains(t, err.Error(), "not validated") + + stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID) + require.NoError(t, err) + assert.Equal(t, "app.proven.example.com", stored.Domain, "the service must keep its original domain") +} diff --git a/management/internals/modules/reverseproxy/service/manager/manager.go b/management/internals/modules/reverseproxy/service/manager/manager.go index 365fbab40..9c7f95eb4 100644 --- a/management/internals/modules/reverseproxy/service/manager/manager.go +++ b/management/internals/modules/reverseproxy/service/manager/manager.go @@ -606,16 +606,19 @@ func (m *Manager) resolveEffectiveCluster(ctx context.Context, accountID string, return existing.ProxyCluster, nil } - if m.clusterDeriver != nil { - derived, err := m.clusterDeriver.DeriveClusterFromDomain(ctx, accountID, svc.Domain) - if err != nil { - log.WithError(err).Warnf("could not derive cluster from domain %s", svc.Domain) - } else { - return derived, nil - } + if m.clusterDeriver == nil { + return existing.ProxyCluster, nil } - return existing.ProxyCluster, nil + // Falling back to the old cluster here would let an update move a service + // onto a domain the account has not validated, bypassing the check that + // creation makes. + derived, err := m.clusterDeriver.DeriveClusterFromDomain(ctx, accountID, svc.Domain) + if err != nil { + return "", status.Errorf(status.PreconditionFailed, "could not derive cluster from domain %s: %v", svc.Domain, err) + } + + return derived, nil } func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.Store, accountID string, service *service.Service, updateInfo *serviceUpdateInfo, customPorts *bool, effectiveCluster string) error { diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 6337ebf1a..ef353ea83 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -5686,6 +5686,23 @@ func (s *SqlStore) ListCustomDomains(ctx context.Context, accountID string) ([]* return domains, nil } +// GetCustomDomainByName returns the custom domain row holding the given name, +// regardless of which account owns it. +func (s *SqlStore) GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) { + customDomain := &domain.Domain{} + result := s.db.Take(customDomain, "domain = ?", domainName) + if result.Error != nil { + if errors.Is(result.Error, gorm.ErrRecordNotFound) { + return nil, status.Errorf(status.NotFound, "custom domain %s not found", domainName) + } + + log.WithContext(ctx).Errorf("failed to get custom domain by name from store: %v", result.Error) + return nil, status.Errorf(status.Internal, "failed to get custom domain from store") + } + + return customDomain, nil +} + func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error) { newDomain := &domain.Domain{ ID: xid.New().String(), // Generate our own ID because gorm doesn't always configure the database to handle this for us. @@ -5697,6 +5714,18 @@ func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, dom } result := s.db.Create(newDomain) if result.Error != nil { + // The unique index is the last guard when two requests clear the + // manager's availability check at the same time. The one that loses the + // insert is a conflict, not an internal failure. + var count int64 + if err := s.db.Model(&domain.Domain{}).Where("domain = ?", domainName).Count(&count).Error; err == nil && count > 0 { + // The insert error is logged even on this path: the name being taken + // is what the caller has to act on, but if the insert also failed for + // an unrelated reason the operator still needs to see it. + log.WithContext(ctx).Warnf("create reverse proxy custom domain %s rejected, name already registered: %v", domainName, result.Error) + return nil, status.Errorf(status.AlreadyExists, "domain %s is already registered", domainName) + } + log.WithContext(ctx).Errorf("failed to create reverse proxy custom domain to store: %v", result.Error) return nil, status.Errorf(status.Internal, "failed to create reverse proxy custom domain to store") } diff --git a/management/server/store/store.go b/management/server/store/store.go index 7daeb28a9..da2b3c6e0 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -302,6 +302,7 @@ type Store interface { GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error) ListFreeDomains(ctx context.Context, accountID string) ([]string, error) ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) + GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error) UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error) DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 70acb9f58..9bf49f076 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -1941,6 +1941,21 @@ func (mr *MockStoreMockRecorder) GetCustomDomain(ctx, accountID, domainID any) * return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetCustomDomain", reflect.TypeOf((*MockStore)(nil).GetCustomDomain), ctx, accountID, domainID) } +// GetCustomDomainByName mocks base method. +func (m *MockStore) GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetCustomDomainByName", ctx, domainName) + ret0, _ := ret[0].(*domain.Domain) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetCustomDomainByName indicates an expected call of GetCustomDomainByName. +func (mr *MockStoreMockRecorder) GetCustomDomainByName(ctx, domainName any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetCustomDomainByName", reflect.TypeOf((*MockStore)(nil).GetCustomDomainByName), ctx, domainName) +} + // GetCustomDomainsCounts mocks base method. func (m *MockStore) GetCustomDomainsCounts(ctx context.Context) (int64, int64, error) { m.ctrl.T.Helper() From 15a684248c7e33556fe6535662ec8ded719b7db9 Mon Sep 17 00:00:00 2001 From: Nicolas Frati Date: Thu, 10 Sep 2026 12:02:19 +0200 Subject: [PATCH 6/9] [client] Support arbitrary UIDs in rootless image (#7440) * [client] Support arbitrary UIDs in rootless image * [client] Keep rootless executables root-owned * [client] Harden arbitrary UID image validation * [client] Preserve executable access in rootless image Keep the binary and entrypoint executable when deployments override the runtime group. Retain root ownership so non-root users cannot modify either file. * [client] Verify rootless state reuse with a stable UID Persisted profiles remain scoped to the creating UID. Verify same-UID container recreation without broadening application permissions, and document the Kubernetes volume permission behavior observed on OpenShift. Remove unused synthetic-user home metadata. * [client] Separate image changes from invoking user fix Keep this PR limited to resolving unmapped non-root invoking users. Move container permissions and their smoke test to a dependent image branch so they can be reviewed separately. * [client] Restore invoking process user test Retain coverage for successful current-user lookup without sudo. Numeric-identity fallback tests do not cover this existing behavior. --- .../internal/profilemanager/invoking_user.go | 35 ++++++++-- .../profilemanager/invoking_user_test.go | 70 ++++++++++++++++++- 2 files changed, 97 insertions(+), 8 deletions(-) diff --git a/client/internal/profilemanager/invoking_user.go b/client/internal/profilemanager/invoking_user.go index c86a6ce43..7ba612ffb 100644 --- a/client/internal/profilemanager/invoking_user.go +++ b/client/internal/profilemanager/invoking_user.go @@ -6,6 +6,7 @@ import ( "os/user" "path/filepath" "runtime" + "strconv" log "github.com/sirupsen/logrus" ) @@ -13,17 +14,21 @@ import ( const envSudoUser = "SUDO_USER" var ( - geteuid = os.Geteuid - lookupUser = user.Lookup + currentUser = user.Current + getegid = os.Getegid + geteuid = os.Geteuid + lookupUser = user.Lookup ) // InvokingUser returns the user a CLI invocation acts for. Under sudo that is // the user who ran sudo, not root: privileged flags force commands through // sudo, and resolving profiles as root would silently switch the daemon to -// root's (default) profile instead of the invoking user's. Privilege decisions -// are not made here — those stay on the kernel credentials of the daemon -// connection, which SUDO_USER (a plain environment variable) can never -// influence; a forged value only selects a profile root could select anyway. +// root's (default) profile instead of the invoking user's. An unmapped positive +// process UID uses its numeric kernel identity; root, sudo lookup failures, and +// unavailable platform identities still fail closed. Privilege decisions stay +// on the kernel credentials of the daemon connection, which SUDO_USER (a plain +// environment variable) can never influence; a forged value only selects a +// profile root could select anyway. func InvokingUser() (*user.User, error) { if u, ok := sudoInvokingUser(); ok { return u, nil @@ -35,7 +40,23 @@ func InvokingUser() (*user.User, error) { if sudoActive() { return nil, fmt.Errorf("resolve sudo invoking user %q: refusing to fall back to root", os.Getenv(envSudoUser)) } - return user.Current() + u, err := currentUser() + if err == nil { + return u, nil + } + + uid := geteuid() + if uid <= 0 { + return nil, err + } + + log.Debugf("current user lookup for UID %d: %v; using numeric UID", uid, err) + uidString := strconv.Itoa(uid) + return &user.User{ + Username: uidString, + Uid: uidString, + Gid: strconv.Itoa(getegid()), + }, nil } // IsPlainRoot reports that the process runs as root with no usable sudo diff --git a/client/internal/profilemanager/invoking_user_test.go b/client/internal/profilemanager/invoking_user_test.go index 54c8ad8fd..159d2616b 100644 --- a/client/internal/profilemanager/invoking_user_test.go +++ b/client/internal/profilemanager/invoking_user_test.go @@ -2,6 +2,7 @@ package profilemanager import ( "errors" + "fmt" "io/fs" "os" "os/user" @@ -21,7 +22,51 @@ func TestInvokingUserFallsBackToProcessUser(t *testing.T) { current, err := user.Current() require.NoError(t, err) - assert.Equal(t, current.Username, got.Username) + assert.Equal(t, current.Username, got.Username, "invoking user should match the process user without sudo") +} + +func TestInvokingUserFailsClosedWithoutPositiveUID(t *testing.T) { + for _, uid := range []int{0, -1} { + t.Run(fmt.Sprintf("UID%d", uid), func(t *testing.T) { + t.Setenv(envSudoUser, "") + lookupErr := errors.New("current user unavailable") + fakeUnmappedUser(t, uid, 0, lookupErr) + + got, err := InvokingUser() + require.ErrorIs(t, err, lookupErr) + assert.Nil(t, got, "root or unavailable UID must not become a synthetic identity") + }) + } +} + +func TestProfileFilePathUsesNumericIdentityForUnmappedNonRoot(t *testing.T) { + t.Setenv(envSudoUser, "") + fakeUnmappedUser(t, 1001230000, 0, errors.New("user: unknown userid 1001230000")) + + profilesRoot := t.TempDir() + origDir := DefaultConfigPathDir + origOverride := ConfigDirOverride + DefaultConfigPathDir = profilesRoot + ConfigDirOverride = "" + t.Cleanup(func() { + DefaultConfigPathDir = origDir + ConfigDirOverride = origOverride + }) + + profileID := ID("0123456789abcdef0123456789abcdef") + got, err := (&Profile{ID: profileID}).FilePath() + require.NoError(t, err) + assert.Equal(t, + filepath.Join(profilesRoot, "1001230000", profileID.String()+".json"), + got, + "profile path should use the numeric UID namespace", + ) + + entries, err := os.ReadDir(profilesRoot) + require.NoError(t, err) + require.Len(t, entries, 1, "only the numeric UID directory should be created") + assert.Equal(t, "1001230000", entries[0].Name(), "profile namespace should be numeric") + assert.True(t, entries[0].IsDir(), "profile namespace should be a directory") } func TestSudoInvokingUserInactiveWithoutSudoContext(t *testing.T) { @@ -60,6 +105,13 @@ func TestInvokingUserFailsClosedWhenSudoLookupFails(t *testing.T) { fakeSudo(t, filepath.Join("/home", "misha")) lookupUser = func(string) (*user.User, error) { return nil, errors.New("nss unavailable") } + origCurrentUser := currentUser + currentUser = func() (*user.User, error) { + t.Fatal("currentUser must not be called after a sudo lookup failure") + return nil, errors.New("currentUser called unexpectedly") + } + t.Cleanup(func() { currentUser = origCurrentUser }) + got, err := InvokingUser() require.Error(t, err) assert.Nil(t, got, "must not resolve to the root process user") @@ -215,6 +267,22 @@ func fakeSudo(t *testing.T, home string) { }) } +func fakeUnmappedUser(t *testing.T, uid, gid int, lookupErr error) { + t.Helper() + + origCurrentUser := currentUser + origEuid := geteuid + origEgid := getegid + currentUser = func() (*user.User, error) { return nil, lookupErr } + geteuid = func() int { return uid } + getegid = func() int { return gid } + t.Cleanup(func() { + currentUser = origCurrentUser + geteuid = origEuid + getegid = origEgid + }) +} + func assertNoEntries(t *testing.T, root string) { t.Helper() err := filepath.WalkDir(root, func(path string, _ fs.DirEntry, err error) error { From 9615d2ab162e7badee5c8b4e84048ce49e1db6e7 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Thu, 10 Sep 2026 12:06:30 +0200 Subject: [PATCH 7/9] [client] Report the remote jobs key in the MDM UI snapshot (#7485) * [client] Report the remote jobs key in the MDM UI snapshot Co-Authored-By: Claude Opus 5 (1M context) * [client] Align the remote jobs snapshot key with the policy key The snapshot field carried the JSON tag remoteJobsAllowed while the policy key is allowRemoteJobs. GetConfigResponse.mDMManagedFields reports the raw policy keys, and applyMDMRestrictions matches them against the struct's JSON tags, so the field never turned true for a policy that set the key. Every other field in Fields already uses its policy key as the JSON tag; this was the only divergence. Co-Authored-By: Claude Opus 5 (1M context) --------- Co-authored-by: Claude Opus 5 (1M context) --- client/mdm/restrictions.go | 2 ++ 1 file changed, 2 insertions(+) diff --git a/client/mdm/restrictions.go b/client/mdm/restrictions.go index c8e443395..200756b78 100644 --- a/client/mdm/restrictions.go +++ b/client/mdm/restrictions.go @@ -20,6 +20,7 @@ type Fields struct { DisableMetricsCollection bool `json:"disableMetricsCollection"` SplitTunnelMode bool `json:"splitTunnelMode"` SplitTunnelApps bool `json:"splitTunnelApps"` + RemoteJobsAllowed bool `json:"allowRemoteJobs"` DisableAdvancedView *bool `json:"disableAdvancedView"` } @@ -60,6 +61,7 @@ func BuildRestrictions(policy *Policy) Restrictions { r.MDM.DisableMetricsCollection = policy.HasKey(KeyDisableMetricsCollection) r.MDM.SplitTunnelMode = policy.HasKey(KeySplitTunnelMode) r.MDM.SplitTunnelApps = policy.HasKey(KeySplitTunnelApps) + r.MDM.RemoteJobsAllowed = policy.HasKey(KeyRemoteJobsAllowed) if v, ok := policy.GetBool(KeyAllowServerSSH); ok { r.MDM.AllowServerSSH = &v } From 0fac1ee638d89a5b2a5cfef38b9c638d57bfad24 Mon Sep 17 00:00:00 2001 From: dmitri-netbird Date: Thu, 10 Sep 2026 16:37:43 +0200 Subject: [PATCH 8/9] [management] cleanup resources when ws-grpc proxy connection goes away (#7484) * ws to grpc connection adapter Signed-off-by: Dmitri Dolguikh * support for timeouts on reading h2 stream headers Signed-off-by: Dmitri Dolguikh * cleanups Signed-off-by: Dmitri Dolguikh * we can't always expect a DATA frame, as not all http methods send it Signed-off-by: Dmitri Dolguikh * set default headers read timeout to 10s Signed-off-by: Dmitri Dolguikh * fix a race in tests Signed-off-by: Dmitri Dolguikh * remove frame interceptor Signed-off-by: Dmitri Dolguikh * cleanup test cleanup Signed-off-by: Dmitri Dolguikh * make linter happy Signed-off-by: Dmitri Dolguikh * removed unused consts Signed-off-by: Dmitri Dolguikh * set 5s ReadTimeout Signed-off-by: Dmitri Dolguikh * making linter happy Signed-off-by: Dmitri Dolguikh * making linter happy Signed-off-by: Dmitri Dolguikh * updated comments Signed-off-by: Dmitri Dolguikh * fix spelling Signed-off-by: Dmitri Dolguikh * disabled all http server read timeouts Signed-off-by: Dmitri Dolguikh * Revert "disabled all http server read timeouts" This reverts commit adf5005ba44d42f620d3a0a0851df9854cfd180d. Signed-off-by: Dmitri Dolguikh * clarify comment re: ReadTimeout/WriteTimeout issues Signed-off-by: Dmitri Dolguikh --------- Signed-off-by: Dmitri Dolguikh --- util/wsproxy/server/proxy.go | 163 +++++----------- util/wsproxy/server/ws_conn_adapter.go | 126 ++++++++++++ util/wsproxy/server/ws_conn_adapter_test.go | 204 ++++++++++++++++++++ 3 files changed, 373 insertions(+), 120 deletions(-) create mode 100644 util/wsproxy/server/ws_conn_adapter.go create mode 100644 util/wsproxy/server/ws_conn_adapter_test.go diff --git a/util/wsproxy/server/proxy.go b/util/wsproxy/server/proxy.go index ffb622200..0618beb91 100644 --- a/util/wsproxy/server/proxy.go +++ b/util/wsproxy/server/proxy.go @@ -1,12 +1,8 @@ package server import ( - "context" - "io" - "net" "net/http" - "sync" - "time" + "sync/atomic" "github.com/coder/websocket" log "github.com/sirupsen/logrus" @@ -15,11 +11,6 @@ import ( "github.com/netbirdio/netbird/util/wsproxy" ) -const ( - bufferSize = 32 * 1024 - ioTimeout = 5 * time.Second -) - // Config contains the configuration for the WebSocket proxy. type Config struct { Handler http.Handler @@ -53,14 +44,23 @@ func New(handler http.Handler, opts ...Option) *Proxy { // Handler returns an http.Handler that proxies WebSocket connections to the local gRPC server. func (p *Proxy) Handler() http.Handler { - return http.HandlerFunc(p.handleWebSocket) + return &proxyHandler{ + metrics: p.config.MetricsRecorder, + handler: p.config.Handler, + } } -func (p *Proxy) handleWebSocket(w http.ResponseWriter, r *http.Request) { +type proxyHandler struct { + metrics MetricsRecorder + handler http.Handler + conn atomic.Pointer[wsConnAdapter] +} + +func (ph *proxyHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { ctx := r.Context() - p.metrics.RecordConnection(ctx) - defer p.metrics.RecordDisconnection(ctx) + ph.metrics.RecordConnection(ctx) + defer ph.metrics.RecordDisconnection(ctx) log.Debugf("WebSocket proxy handling connection from %s, forwarding to internal gRPC handler", r.RemoteAddr) acceptOptions := &websocket.AcceptOptions{ @@ -69,121 +69,44 @@ func (p *Proxy) handleWebSocket(w http.ResponseWriter, r *http.Request) { wsConn, err := websocket.Accept(w, r, acceptOptions) if err != nil { - p.metrics.RecordError(ctx, "websocket_accept_failed") + ph.metrics.RecordError(ctx, "websocket_accept_failed") log.Errorf("WebSocket upgrade failed from %s: %v", r.RemoteAddr, err) return } - defer func() { - _ = wsConn.Close(websocket.StatusNormalClosure, "") - }() + serverConn := (&wsConnAdapter{ + ctx: ctx, + conn: wsConn, + metrics: ph.metrics, + clientAddr: r.RemoteAddr, + }) - clientConn, serverConn := net.Pipe() defer func() { - _ = clientConn.Close() _ = serverConn.Close() }() + ph.conn.Store(serverConn) // used in tests only + log.Debugf("WebSocket proxy established: %s -> gRPC handler", r.RemoteAddr) - go func() { - (&http2.Server{}).ServeConn(serverConn, &http2.ServeConnOpts{ - Context: ctx, - Handler: p.config.Handler, - }) - }() + (&http2.Server{ + // TODO (dmitri) we should limit the number of concurrent streams per connection (peer) + // and idle timeouts + // MaxConcurrentStreams: 20, + // IdleTimeout: 10 * time.Second, + }).ServeConn(serverConn, &http2.ServeConnOpts{ + Context: ctx, + Handler: ph.handler, + BaseConfig: &http.Server{ + // b/c we are wrapping a ws connection, read and write connection deadlines normally set + // via ReadTimeout and WriteTimeout http.Server fields aren't available to us. The ws + // library doesn't expose connection deadline timer config, and we ignore these calls in "wsConnAdapter". + // + // Another issue is that Server.ServeConn() call bypasses setting of connection deadlines altogether, + // ReadTimeout and Writetimeout set here would only apply to h2 streams, i.e. after a HEADERS frame + // arrival and processing, turning ReadTimeout into a request body read deadline, and WriteTimeout into + // a response deadline (the latter not useful for streaming requests). + }, + }) - p.proxyData(ctx, wsConn, clientConn, r.RemoteAddr) -} - -func (p *Proxy) proxyData(ctx context.Context, wsConn *websocket.Conn, pipeConn net.Conn, clientAddr string) { - proxyCtx, cancel := context.WithCancel(ctx) - defer cancel() - - var wg sync.WaitGroup - wg.Add(2) - - go p.wsToPipe(proxyCtx, cancel, &wg, wsConn, pipeConn, clientAddr) - go p.pipeToWS(proxyCtx, cancel, &wg, wsConn, pipeConn, clientAddr) - - wg.Wait() -} - -func (p *Proxy) wsToPipe(ctx context.Context, cancel context.CancelFunc, wg *sync.WaitGroup, wsConn *websocket.Conn, pipeConn net.Conn, clientAddr string) { - defer wg.Done() - defer cancel() - - for { - msgType, data, err := wsConn.Read(ctx) - if err != nil { - switch { - case ctx.Err() != nil: - log.Debugf("WebSocket from %s terminating due to context cancellation", clientAddr) - case websocket.CloseStatus(err) != -1: - log.Debugf("WebSocket from %s disconnected", clientAddr) - default: - p.metrics.RecordError(ctx, "websocket_read_error") - log.Debugf("WebSocket read error from %s: %v", clientAddr, err) - } - return - } - - if msgType != websocket.MessageBinary { - log.Warnf("Unexpected WebSocket message type from %s: %v", clientAddr, msgType) - continue - } - - if ctx.Err() != nil { - log.Tracef("wsToPipe goroutine terminating due to context cancellation before pipe write") - return - } - - if err := pipeConn.SetWriteDeadline(time.Now().Add(ioTimeout)); err != nil { - log.Debugf("Failed to set pipe write deadline: %v", err) - } - - n, err := pipeConn.Write(data) - if err != nil { - p.metrics.RecordError(ctx, "pipe_write_error") - log.Warnf("Pipe write error for %s: %v", clientAddr, err) - return - } - - p.metrics.RecordBytesTransferred(ctx, "ws_to_grpc", int64(n)) - } -} - -func (p *Proxy) pipeToWS(ctx context.Context, cancel context.CancelFunc, wg *sync.WaitGroup, wsConn *websocket.Conn, pipeConn net.Conn, clientAddr string) { - defer wg.Done() - defer cancel() - - buf := make([]byte, bufferSize) - for { - n, err := pipeConn.Read(buf) - if err != nil { - if ctx.Err() != nil { - log.Tracef("pipeToWS goroutine terminating due to context cancellation") - return - } - - if err != io.EOF { - log.Debugf("Pipe read error for %s: %v", clientAddr, err) - } - return - } - - if ctx.Err() != nil { - log.Tracef("pipeToWS goroutine terminating due to context cancellation before WebSocket write") - return - } - - if n > 0 { - if err := wsConn.Write(ctx, websocket.MessageBinary, buf[:n]); err != nil { - p.metrics.RecordError(ctx, "websocket_write_error") - log.Warnf("WebSocket write error for %s: %v", clientAddr, err) - return - } - - p.metrics.RecordBytesTransferred(ctx, "grpc_to_ws", int64(n)) - } - } + log.Debugf("WebSocket proxy closing: %s -> gRPC handler", r.RemoteAddr) } diff --git a/util/wsproxy/server/ws_conn_adapter.go b/util/wsproxy/server/ws_conn_adapter.go new file mode 100644 index 000000000..eb29ab0cb --- /dev/null +++ b/util/wsproxy/server/ws_conn_adapter.go @@ -0,0 +1,126 @@ +package server + +import ( + "context" + "net" + "sync/atomic" + "time" + + "github.com/coder/websocket" + log "github.com/sirupsen/logrus" +) + +type wsConnAdapter struct { + prefix string + ctx context.Context + conn *websocket.Conn + metrics MetricsRecorder + clientAddr string + closed atomic.Bool + bufferedRead []byte +} + +var _ net.Conn = &wsConnAdapter{} + +type wsAddr struct{ prefix string } + +func (wa wsAddr) Network() string { return wa.prefix + "ws-proxy" } +func (wa wsAddr) String() string { return wa.prefix + "ws-proxy" } + +func (ws *wsConnAdapter) Read(b []byte) (int, error) { + if len(ws.bufferedRead) > 0 { + return ws.readFromBuffer(b) + } + + msgType, data, err := ws.conn.Read(ws.ctx) + if err != nil { + switch { + case ws.ctx.Err() != nil: + log.Debugf("WebSocket from %s terminating due to context cancellation", ws.clientAddr) + case websocket.CloseStatus(err) != -1: + log.Debugf("WebSocket from %s disconnected", ws.clientAddr) + default: + ws.recordError(ws.ctx, "websocket_read_error") + log.Debugf("WebSocket read error from %s: %v", ws.clientAddr, err) + } + return copy(b, data), err + } + if msgType != websocket.MessageBinary { + log.Warnf("Unexpected WebSocket message type from %s: %v", ws.clientAddr, msgType) + return 0, nil + } + + ws.bufferedRead = data + return ws.readFromBuffer(b) +} + +func (ws *wsConnAdapter) readFromBuffer(b []byte) (int, error) { + n := copy(b, ws.bufferedRead) + + ws.recordBytesTransferred(ws.ctx, "ws_to_grpc", n) + if n == len(ws.bufferedRead) { + ws.bufferedRead = nil + return n, nil + } else { + ws.bufferedRead = ws.bufferedRead[n:] + } + return n, nil +} + +func (ws *wsConnAdapter) Write(b []byte) (int, error) { + maybeErr := ws.ctx.Err() + + n := len(b) + if n == 0 { + return n, maybeErr + } + if maybeErr != nil { + return 0, maybeErr + } + if err := ws.conn.Write(ws.ctx, websocket.MessageBinary, b[:n]); err != nil { + ws.recordError(ws.ctx, "websocket_write_error") + log.Warnf("WebSocket write error for %s: %v", ws.clientAddr, err) + return 0, err // we don't know how many bytes have been written + } + + ws.recordBytesTransferred(ws.ctx, "grpc_to_ws", n) + return n, nil +} + +func (ws *wsConnAdapter) Close() error { + ws.closed.Store(true) + return ws.conn.Close(websocket.StatusNormalClosure, "") +} + +func (ws *wsConnAdapter) LocalAddr() net.Addr { return wsAddr{ws.prefix} } +func (ws *wsConnAdapter) RemoteAddr() net.Addr { return wsAddr{ws.prefix} } + +func (ws *wsConnAdapter) SetDeadline(t time.Time) error { + return nil +} + +func (ws *wsConnAdapter) SetReadDeadline(t time.Time) error { + return nil +} + +func (ws *wsConnAdapter) SetWriteDeadline(t time.Time) error { + return nil +} + +func (ws *wsConnAdapter) recordError(ctx context.Context, errorType string) { + if ws.metrics == nil { + return + } + ws.metrics.RecordError(ctx, errorType) +} + +func (ws *wsConnAdapter) recordBytesTransferred(ctx context.Context, direction string, bytes int) { + if ws.metrics == nil { + return + } + ws.metrics.RecordBytesTransferred(ctx, direction, int64(bytes)) +} + +func (ws *wsConnAdapter) IsClosed() bool { + return ws.closed.Load() +} diff --git a/util/wsproxy/server/ws_conn_adapter_test.go b/util/wsproxy/server/ws_conn_adapter_test.go new file mode 100644 index 000000000..5369b2362 --- /dev/null +++ b/util/wsproxy/server/ws_conn_adapter_test.go @@ -0,0 +1,204 @@ +package server + +import ( + "bytes" + "context" + "crypto/tls" + "io" + "math/rand/v2" + "net" + "net/http" + "os" + "path/filepath" + "strconv" + "strings" + "testing" + "time" + + "github.com/coder/websocket" + "github.com/stretchr/testify/assert" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +func TestAdapterHandlingConnectionClosures(t *testing.T) { + var cases = []struct { + description string + casenum int + }{ + {"client-side ws connection is closed", 0}, + {"server-side ws connection is closed", 1}, + {"client-side context is cancelled", 2}, + {"server-side context is cancelled", 3}, + } + + for _, c := range cases { + t.Run(c.description, func(t *testing.T) { + serversock := filepath.Join("/tmp", "http-server-"+strconv.FormatInt(rand.Int64(), 10)+".sock") + t.Cleanup(func() { os.Remove(serversock) }) + + l, err := net.Listen("unix", serversock) + assert.NoError(t, err) + + proxy := New(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + buf, _ := io.ReadAll(r.Body) + defer r.Body.Close() + w.Write([]byte("echo: " + string(buf))) //nolint:errcheck + })) + + handler, ok := proxy.Handler().(*proxyHandler) + assert.True(t, ok) + + protocols := new(http.Protocols) + protocols.SetHTTP1(true) + protocols.SetUnencryptedHTTP2(true) + httpServer := http.Server{ + Handler: handler, + } + go httpServer.Serve(l) //nolint:errcheck + t.Cleanup(func() { httpServer.Close() }) + + clientconn, _, err := websocket.Dial(context.Background(), "http://whatever", //nolint:bodyclose + &websocket.DialOptions{HTTPClient: &http.Client{ + Transport: &http.Transport{ + DialContext: func(_ context.Context, _, _ string) (net.Conn, error) { + return net.Dial("unix", serversock) + }, + }}}) + assert.NoError(t, err) + + clientCtx, cancel := context.WithCancel(context.Background()) //nolint:govet + h2client := &http.Client{ + Transport: &http2.Transport{ + AllowHTTP: true, + DialTLSContext: func(_ context.Context, _, _ string, _ *tls.Config) (net.Conn, error) { + return &wsConnAdapter{ + prefix: "test-client", + ctx: clientCtx, + conn: clientconn, + }, nil + }, + }} + + resp, err := h2client.Post("http://whatever", "text/html", strings.NewReader("g'day")) + assert.NoError(t, err) + + body, err := io.ReadAll(resp.Body) + defer resp.Body.Close() + + assert.NoError(t, err) + assert.Equal(t, "echo: g'day", string(body)) + + switch c.casenum { + case 0: + clientconn.Close(websocket.StatusNormalClosure, "") + case 1: + handler.conn.Load().Close() + case 2: + cancel() + case 3: + resp.Body.Close() + h2client.CloseIdleConnections() + } + + assert.EventuallyWithT(t, func(c *assert.CollectT) { + assert.True(c, handler.conn.Load().IsClosed()) + }, 3*time.Second, 100*time.Millisecond) + }) //nolint:govet + } +} + +func TestAdapterHandlingHttpConnection_NoHeadersSent(t *testing.T) { + t.Skip("currently disabled as it requires idle timeout to be set") + + serversock := filepath.Join("/tmp", "http-server-"+strconv.FormatInt(rand.Int64(), 10)+".sock") + defer os.Remove(serversock) + + l, err := net.Listen("unix", serversock) + assert.NoError(t, err) + + proxy := New(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + buf, _ := io.ReadAll(r.Body) + defer r.Body.Close() //nolint:errcheck + w.Write([]byte("echo: " + string(buf))) //nolint:errcheck + })) + + handler, ok := proxy.Handler().(*proxyHandler) + assert.True(t, ok) + + protocols := new(http.Protocols) + protocols.SetHTTP1(true) + protocols.SetUnencryptedHTTP2(true) + httpServer := http.Server{ + Handler: handler, + } + go httpServer.Serve(l) //nolint:errcheck + + clientconn, _, err := websocket.Dial(context.Background(), "http://whatever", //nolint:bodyclose + &websocket.DialOptions{HTTPClient: &http.Client{ + Transport: &http.Transport{ + DialContext: func(_ context.Context, _, _ string) (net.Conn, error) { + return net.Dial("unix", serversock) + }, + }}}) + assert.NoError(t, err) + + h2client := &http.Client{ + Transport: &http2.Transport{ + AllowHTTP: true, + DialTLSContext: func(_ context.Context, _, _ string, _ *tls.Config) (net.Conn, error) { + return &h2ConnectionSnooper{wrappedConn: &wsConnAdapter{ + prefix: "test-client", + ctx: context.Background(), + conn: clientconn, + }, shouldDropFrame: func(f http2.FrameType) bool { return f == http2.FrameHeaders || f == http2.FrameData }}, nil + }, + }} + + _, err = h2client.Post("http://whatever", "text/html", strings.NewReader("g'day")) + assert.Error(t, err) + + assert.EventuallyWithT(t, func(c *assert.CollectT) { + assert.True(c, handler.conn.Load().IsClosed()) + }, 3*time.Second, 100*time.Millisecond) +} + +type h2ConnectionSnooper struct { + wrappedConn net.Conn + shouldDropFrame func(f http2.FrameType) bool +} + +func (hs *h2ConnectionSnooper) Read(b []byte) (n int, err error) { + return hs.wrappedConn.Read(b) +} + +func (hs *h2ConnectionSnooper) Write(b []byte) (n int, err error) { + fr := http2.NewFramer(nil, bytes.NewReader(b)) + fr.ReadMetaHeaders = hpack.NewDecoder(0, nil) + f, err := fr.ReadFrame() + if err != nil { + return hs.wrappedConn.Write(b) + } + + if hs.shouldDropFrame != nil && hs.shouldDropFrame(f.Header().Type) { + return len(b), nil + } + + return hs.wrappedConn.Write(b) +} + +func (hs *h2ConnectionSnooper) Close() error { return hs.wrappedConn.Close() } + +func (hs *h2ConnectionSnooper) LocalAddr() net.Addr { return hs.wrappedConn.LocalAddr() } + +func (hs *h2ConnectionSnooper) RemoteAddr() net.Addr { return hs.wrappedConn.RemoteAddr() } + +func (hs *h2ConnectionSnooper) SetDeadline(t time.Time) error { return hs.wrappedConn.SetDeadline(t) } + +func (hs *h2ConnectionSnooper) SetReadDeadline(t time.Time) error { + return hs.wrappedConn.SetReadDeadline(t) +} + +func (hs *h2ConnectionSnooper) SetWriteDeadline(t time.Time) error { + return hs.wrappedConn.SetWriteDeadline(t) +} From e704203927fcd40ee2b4687a134d513c3909f3bf Mon Sep 17 00:00:00 2001 From: dmitri-netbird Date: Thu, 10 Sep 2026 20:05:42 +0200 Subject: [PATCH 9/9] [management] do not hard-code tmp dir path in ws_conn_adapter_test (#7503) * do not hard-code tmp dir path Signed-off-by: Dmitri Dolguikh * use os.TempDir to get tmp dir Signed-off-by: Dmitri Dolguikh --------- Signed-off-by: Dmitri Dolguikh --- util/wsproxy/server/ws_conn_adapter_test.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/util/wsproxy/server/ws_conn_adapter_test.go b/util/wsproxy/server/ws_conn_adapter_test.go index 5369b2362..d46e4830b 100644 --- a/util/wsproxy/server/ws_conn_adapter_test.go +++ b/util/wsproxy/server/ws_conn_adapter_test.go @@ -34,7 +34,7 @@ func TestAdapterHandlingConnectionClosures(t *testing.T) { for _, c := range cases { t.Run(c.description, func(t *testing.T) { - serversock := filepath.Join("/tmp", "http-server-"+strconv.FormatInt(rand.Int64(), 10)+".sock") + serversock := filepath.Join(os.TempDir(), "http-server-"+strconv.FormatInt(rand.Int64(), 10)+".sock") t.Cleanup(func() { os.Remove(serversock) }) l, err := net.Listen("unix", serversock) @@ -111,7 +111,7 @@ func TestAdapterHandlingConnectionClosures(t *testing.T) { func TestAdapterHandlingHttpConnection_NoHeadersSent(t *testing.T) { t.Skip("currently disabled as it requires idle timeout to be set") - serversock := filepath.Join("/tmp", "http-server-"+strconv.FormatInt(rand.Int64(), 10)+".sock") + serversock := filepath.Join(os.TempDir(), "http-server-"+strconv.FormatInt(rand.Int64(), 10)+".sock") defer os.Remove(serversock) l, err := net.Listen("unix", serversock)