From f400f4bee8d110c1bae499a42e52a505792877c5 Mon Sep 17 00:00:00 2001 From: Brad Ison Date: Fri, 2 Oct 2026 12:02:06 +0200 Subject: [PATCH 01/14] [proxy] Optionally refuse private addresses on direct-upstream dials (#7913) Direct-upstream targets are dialled on the proxy host's network stack, outside the embedded client's LAN blocking. A proxy that serves untrusted accounts lets them reach the host's loopback, its LAN or cluster, and the cloud metadata service through such a target. NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE adds a dialer control that refuses addresses that are not globally reachable. It checks each socket's resolved address just before connect, so hostnames and DNS rebinding are covered, and IPv4 embedded in IPv6 addresses is checked as IPv4. Refused dials are served as a 502. The setting defaults to off for private and self-hosted proxies; an unparsable value turns it on. --- proxy/internal/proxy/reverseproxy.go | 6 + proxy/internal/proxy/reverseproxy_test.go | 11 ++ proxy/internal/roundtrip/dialguard.go | 90 ++++++++++ proxy/internal/roundtrip/dialguard_test.go | 189 +++++++++++++++++++++ proxy/internal/roundtrip/multi.go | 7 +- proxy/internal/roundtrip/transport.go | 27 +++ 6 files changed, 329 insertions(+), 1 deletion(-) create mode 100644 proxy/internal/roundtrip/dialguard.go create mode 100644 proxy/internal/roundtrip/dialguard_test.go diff --git a/proxy/internal/proxy/reverseproxy.go b/proxy/internal/proxy/reverseproxy.go index 7b0acd1b3..7583b2e01 100644 --- a/proxy/internal/proxy/reverseproxy.go +++ b/proxy/internal/proxy/reverseproxy.go @@ -809,6 +809,12 @@ func classifyProxyError(err error) (title, message string, code int, status web. http.StatusBadGateway, web.ErrorStatus{Proxy: false, Destination: false} + case errors.Is(err, roundtrip.ErrDirectUpstreamBlocked): + return "Destination Not Allowed", + "This proxy does not connect to private or internal addresses. Please contact your administrator.", + http.StatusBadGateway, + web.ErrorStatus{Proxy: false, Destination: false} + case errors.Is(err, roundtrip.ErrTooManyInflight): return "Service Overloaded", "The service is currently handling too many requests. Please try again shortly.", diff --git a/proxy/internal/proxy/reverseproxy_test.go b/proxy/internal/proxy/reverseproxy_test.go index 83afee387..c0724ce84 100644 --- a/proxy/internal/proxy/reverseproxy_test.go +++ b/proxy/internal/proxy/reverseproxy_test.go @@ -1053,6 +1053,17 @@ func TestClassifyProxyError(t *testing.T) { wantCode: http.StatusBadGateway, wantStatus: web.ErrorStatus{Proxy: true, Destination: false}, }, + { + name: "direct upstream blocked by dial guard", + err: &net.OpError{ + Op: "dial", + Net: "tcp", + Err: roundtrip.ErrDirectUpstreamBlocked, + }, + wantTitle: "Destination Not Allowed", + wantCode: http.StatusBadGateway, + wantStatus: web.ErrorStatus{Proxy: false, Destination: false}, + }, { name: "unknown error falls to default", err: errors.New("something unexpected"), diff --git a/proxy/internal/roundtrip/dialguard.go b/proxy/internal/roundtrip/dialguard.go new file mode 100644 index 000000000..ac01263b4 --- /dev/null +++ b/proxy/internal/roundtrip/dialguard.go @@ -0,0 +1,90 @@ +package roundtrip + +import ( + "context" + "errors" + "net/netip" + "syscall" +) + +// ErrDirectUpstreamBlocked is returned when a direct-upstream dial targets +// an address that is not globally reachable while +// NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE is set. +var ErrDirectUpstreamBlocked = errors.New("direct upstream address is not allowed") + +// blockedUpstreamPrefixes are the ranges that reach the proxy host, its +// cluster or its cloud provider rather than the public internet. NAT64 +// and 6to4 addresses are matched by the IPv4 address they embed. +var blockedUpstreamPrefixes = []netip.Prefix{ + // IPv4 + netip.MustParsePrefix("0.0.0.0/8"), // "this network", including 0.0.0.0 + netip.MustParsePrefix("10.0.0.0/8"), // RFC1918 + netip.MustParsePrefix("100.64.0.0/10"), // CGNAT + netip.MustParsePrefix("127.0.0.0/8"), // loopback + netip.MustParsePrefix("169.254.0.0/16"), // link-local, cloud metadata services + netip.MustParsePrefix("172.16.0.0/12"), // RFC1918 + netip.MustParsePrefix("192.0.0.0/24"), // IETF protocol assignments + netip.MustParsePrefix("192.0.2.0/24"), // documentation + netip.MustParsePrefix("192.88.99.0/24"), // 6to4 relay anycast (deprecated) + netip.MustParsePrefix("192.168.0.0/16"), // RFC1918 + netip.MustParsePrefix("198.18.0.0/15"), // benchmarking + netip.MustParsePrefix("198.51.100.0/24"), // documentation + netip.MustParsePrefix("203.0.113.0/24"), // documentation + netip.MustParsePrefix("224.0.0.0/4"), // multicast + netip.MustParsePrefix("240.0.0.0/4"), // reserved, including broadcast + + // IPv6 + netip.MustParsePrefix("::/96"), // unspecified, loopback, IPv4-compatible + netip.MustParsePrefix("64:ff9b:1::/48"), // local-use NAT64 + netip.MustParsePrefix("100::/64"), // discard-only + netip.MustParsePrefix("2001::/32"), // Teredo + netip.MustParsePrefix("2001:2::/48"), // benchmarking + netip.MustParsePrefix("2001:db8::/32"), // documentation + netip.MustParsePrefix("3fff::/20"), // documentation + netip.MustParsePrefix("5f00::/16"), // SRv6 SIDs + netip.MustParsePrefix("fc00::/7"), // unique local, including AWS IMDS fd00:ec2::254 + netip.MustParsePrefix("fe80::/10"), // link-local + netip.MustParsePrefix("fec0::/10"), // site-local (deprecated) + netip.MustParsePrefix("ff00::/8"), // multicast +} + +var ( + nat64Prefix = netip.MustParsePrefix("64:ff9b::/96") + sixToFour = netip.MustParsePrefix("2002::/16") +) + +// isBlockedUpstreamAddr reports whether a guarded direct-upstream dial +// must refuse addr. +func isBlockedUpstreamAddr(addr netip.Addr) bool { + addr = addr.Unmap().WithZone("") + if !addr.IsValid() { + return true + } + + if nat64Prefix.Contains(addr) { + b := addr.As16() + return isBlockedUpstreamAddr(netip.AddrFrom4([4]byte(b[12:16]))) + } + if sixToFour.Contains(addr) { + b := addr.As16() + return isBlockedUpstreamAddr(netip.AddrFrom4([4]byte(b[2:6]))) + } + + for _, p := range blockedUpstreamPrefixes { + if p.Contains(addr) { + return true + } + } + return false +} + +// guardUpstreamDial is a net.Dialer ControlContext that refuses blocked +// addresses. It sees the resolved address of each socket just before +// connect, so DNS rebinding cannot swap the target after the check. +func guardUpstreamDial(_ context.Context, _, address string, _ syscall.RawConn) error { + ap, err := netip.ParseAddrPort(address) + if err != nil || isBlockedUpstreamAddr(ap.Addr()) { + return ErrDirectUpstreamBlocked + } + return nil +} diff --git a/proxy/internal/roundtrip/dialguard_test.go b/proxy/internal/roundtrip/dialguard_test.go new file mode 100644 index 000000000..79453d8c3 --- /dev/null +++ b/proxy/internal/roundtrip/dialguard_test.go @@ -0,0 +1,189 @@ +package roundtrip + +import ( + "context" + "io" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "net/url" + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestIsBlockedUpstreamAddr(t *testing.T) { + blocked := []string{ + "0.0.0.0", + "0.1.2.3", + "10.1.2.3", + "100.64.0.1", + "100.127.255.254", + "127.0.0.1", + "127.255.255.255", + "169.254.169.254", + "172.16.0.1", + "172.31.255.255", + "192.0.0.170", + "192.168.1.1", + "192.88.99.1", + "198.18.0.1", + "224.0.0.1", + "255.255.255.255", + "::", + "::1", + "::169.254.169.254", + "::ffff:127.0.0.1", + "::ffff:169.254.169.254", + "::ffff:10.0.0.1", + "64:ff9b::a9fe:a9fe", // NAT64 of 169.254.169.254 + "64:ff9b::a00:1", // NAT64 of 10.0.0.1 + "64:ff9b:1::1", + "2001::1", + "2001:0:4136:e378:8000:63bf:3fff:fdd2", + "2001:2::1", + "3fff::1", + "5f00::1", + "2002:a9fe:a9fe::1", // 6to4 of 169.254.169.254 + "2002:7f00:1::", // 6to4 of 127.0.0.1 + "fc00::1", + "fd00:ec2::254", + "fe80::1", + "fe80::1%eth0", + "fec0::1", + "ff02::1", + } + for _, s := range blocked { + t.Run("blocks "+s, func(t *testing.T) { + assert.True(t, isBlockedUpstreamAddr(netip.MustParseAddr(s))) + }) + } + + allowed := []string{ + "1.1.1.1", + "8.8.8.8", + "100.63.255.255", + "100.128.0.0", + "172.15.255.255", + "172.32.0.0", + "169.253.255.255", + "2606:4700:4700::1111", + "2001:4860:4860::8888", + "2001:1::1", + "4000::1", + "::ffff:8.8.8.8", + "64:ff9b::808:808", // NAT64 of 8.8.8.8 + "2002:808:808::1", // 6to4 of 8.8.8.8 + } + for _, s := range allowed { + t.Run("allows "+s, func(t *testing.T) { + assert.False(t, isBlockedUpstreamAddr(netip.MustParseAddr(s))) + }) + } + + assert.True(t, isBlockedUpstreamAddr(netip.Addr{}), "the zero Addr must be refused") +} + +func TestGuardUpstreamDial_RejectsUnparsableAddress(t *testing.T) { + err := guardUpstreamDial(context.Background(), "tcp", "not-an-address", nil) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked, "an address the guard cannot parse must fail closed") +} + +// TestMultiTransport_BlockPrivateUpstreams exercises the guard end to end +// against a loopback test server: by IP literal and by a hostname that +// resolves to loopback, on both direct branches, and confirms the +// embedded branch is not affected. +func TestMultiTransport_BlockPrivateUpstreams(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, "reached") + })) + defer srv.Close() + + _, port, err := net.SplitHostPort(srv.Listener.Addr().String()) + require.NoError(t, err) + byName := (&url.URL{Scheme: "http", Host: net.JoinHostPort("localhost", port)}).String() + + directCtx := WithDirectUpstream(context.Background()) + insecureCtx := WithSkipTLSVerify(directCtx) + + // roundTrip returns the response body, so callers never hold one open. + roundTrip := func(t *testing.T, mt *MultiTransport, ctx context.Context, target string) (string, error) { + t.Helper() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil) + require.NoError(t, err) + resp, err := mt.RoundTrip(req) + if err != nil { + return "", err + } + defer func() { _ = resp.Body.Close() }() + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return string(body), nil + } + + t.Run("enabled refuses loopback", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "true") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + cases := []struct { + name string + ctx context.Context + target string + }{ + {"direct by IP", directCtx, srv.URL}, + {"direct by hostname", directCtx, byName}, + {"insecure by IP", insecureCtx, srv.URL}, + {"insecure by hostname", insecureCtx, byName}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := roundTrip(t, mt, tc.ctx, tc.target) + require.Error(t, err) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked) + }) + } + }) + + t.Run("invalid value enables the guard", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "yes please") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + _, err := roundTrip(t, mt, directCtx, srv.URL) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked, "a value that does not parse must fail closed") + }) + + t.Run("explicit false disables the guard", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "false") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + body, err := roundTrip(t, mt, directCtx, srv.URL) + require.NoError(t, err) + assert.Equal(t, "reached", body) + }) + + t.Run("enabled leaves embedded branch alone", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "true") + embedded := &stubRoundTripper{body: "embedded"} + mt := NewMultiTransport(embedded, nil) + + body, err := roundTrip(t, mt, context.Background(), srv.URL) + require.NoError(t, err) + assert.Equal(t, "embedded", body) + assert.True(t, embedded.called, "the guard must not change dispatch to the embedded transport") + }) + + t.Run("disabled by default", func(t *testing.T) { + // Register the restore first so an exported value comes back after + // the test, then exercise a genuinely absent variable. + t.Setenv(EnvDirectUpstreamBlockPrivate, "") + require.NoError(t, os.Unsetenv(EnvDirectUpstreamBlockPrivate)) + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + body, err := roundTrip(t, mt, directCtx, srv.URL) + require.NoError(t, err, "private and self-hosted proxies must keep reaching local upstreams") + assert.Equal(t, "reached", body) + }) +} diff --git a/proxy/internal/roundtrip/multi.go b/proxy/internal/roundtrip/multi.go index d50ad1fc9..a430d45bd 100644 --- a/proxy/internal/roundtrip/multi.go +++ b/proxy/internal/roundtrip/multi.go @@ -41,7 +41,9 @@ var errNoEmbeddedTransport = errors.New("multitransport: embedded roundtripper n // MultiTransport that only ever uses the direct branch. The direct // branches honour the same NB_PROXY_* tuning env vars as the embedded // transport (see loadTransportConfig) plus a dial-timeout wrapper that -// respects types.WithDialTimeout. +// respects types.WithDialTimeout. With NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE +// set, the direct branches refuse addresses that are not globally reachable +// (see guardUpstreamDial). func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTransport { if logger == nil { logger = log.StandardLogger() @@ -51,6 +53,9 @@ func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTra Timeout: 30 * time.Second, KeepAlive: 30 * time.Second, } + if cfg.blockPrivateUpstreams { + dialer.ControlContext = guardUpstreamDial + } direct := &http.Transport{ DialContext: dialWithTimeout(dialer.DialContext), MaxIdleConns: cfg.maxIdleConns, diff --git a/proxy/internal/roundtrip/transport.go b/proxy/internal/roundtrip/transport.go index 9e872e447..6383079c4 100644 --- a/proxy/internal/roundtrip/transport.go +++ b/proxy/internal/roundtrip/transport.go @@ -25,6 +25,12 @@ const ( EnvDisableCompression = "NB_PROXY_DISABLE_COMPRESSION" EnvMaxInflight = "NB_PROXY_MAX_INFLIGHT" EnvUpstreamHTTPVersion = "NB_PROXY_UPSTREAM_HTTP_VERSION" + // EnvDirectUpstreamBlockPrivate refuses direct-upstream dials to + // addresses that are not globally reachable (loopback, private, + // link-local, CGNAT, ...). Off by default: private and self-hosted + // proxies use direct_upstream to reach LAN and localhost services. + // Proxies that serve untrusted accounts must turn it on. + EnvDirectUpstreamBlockPrivate = "NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE" ) // upstreamHTTPVersion selects the HTTP version the proxy uses towards an @@ -69,6 +75,9 @@ type transportConfig struct { // explicit values are for backends whose advertised h2 support is // unusable and whose failure mode the negotiation cannot see. upstreamHTTPVersion upstreamHTTPVersion + // blockPrivateUpstreams guards the direct branches' dialer with + // guardUpstreamDial. It has no effect on the embedded branch. + blockPrivateUpstreams bool } func defaultTransportConfig() transportConfig { @@ -122,6 +131,7 @@ func loadTransportConfig(logger *log.Logger) transportConfig { if v, ok := envUpstreamHTTPVersion(EnvUpstreamHTTPVersion, logger); ok { cfg.upstreamHTTPVersion = v } + cfg.blockPrivateUpstreams = envGuardBool(EnvDirectUpstreamBlockPrivate, logger) logger.WithFields(log.Fields{ "max_idle_conns": cfg.maxIdleConns, @@ -136,6 +146,7 @@ func loadTransportConfig(logger *log.Logger) transportConfig { "disable_compression": cfg.disableCompression, "max_inflight": cfg.maxInflight, "upstream_http_version": cfg.upstreamHTTPVersion, + "block_private_upstreams": cfg.blockPrivateUpstreams, }).Debug("backend transport configuration") return cfg @@ -246,6 +257,22 @@ func envDuration(key string, logger *log.Logger) (time.Duration, bool) { return v, true } +// envGuardBool reads a bool that turns a security guard on. Unset means +// off, but a value that does not parse turns the guard on: a typo must not +// leave a proxy that was meant to be guarded without the guard. +func envGuardBool(key string, logger *log.Logger) bool { + s := os.Getenv(key) + if s == "" { + return false + } + v, err := strconv.ParseBool(s) + if err != nil { + logger.Warnf("failed to parse %s=%q as bool, enabling it: %v", key, s, err) + return true + } + return v +} + func envBool(key string, logger *log.Logger) (bool, bool) { s := os.Getenv(key) if s == "" { From 0712a5a5b92e17392ddd809d1fe0c0188175f611 Mon Sep 17 00:00:00 2001 From: Bethuel Mmbaga Date: Fri, 2 Oct 2026 15:55:51 +0300 Subject: [PATCH 02/14] [management,proxy] Rename the OIDC session code query parameter (#7981) --- management/server/http/handlers/proxy/auth.go | 4 ++-- .../proxy/auth_callback_integration_test.go | 8 ++++---- proxy/auth/auth.go | 8 ++++++++ proxy/internal/auth/middleware.go | 10 +++++----- proxy/internal/auth/middleware_test.go | 13 ++++++++++--- proxy/internal/auth/oidc.go | 6 +++--- proxy/internal/proxy/reverseproxy.go | 6 +++--- proxy/internal/proxy/reverseproxy_test.go | 11 +++++++++++ proxy/web/web.go | 6 ++++-- 9 files changed, 50 insertions(+), 22 deletions(-) diff --git a/management/server/http/handlers/proxy/auth.go b/management/server/http/handlers/proxy/auth.go index 298fb503e..133236401 100644 --- a/management/server/http/handlers/proxy/auth.go +++ b/management/server/http/handlers/proxy/auth.go @@ -125,9 +125,9 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ http.Error(w, "Failed to create session", http.StatusInternalServerError) return } - query.Set("session_code", code) + query.Set(auth.SessionCodeQueryParam, code) } else { - query.Set("session_token", sessionToken) + query.Set(auth.SessionTokenQueryParam, sessionToken) } redirectURL.RawQuery = query.Encode() diff --git a/management/server/http/handlers/proxy/auth_callback_integration_test.go b/management/server/http/handlers/proxy/auth_callback_integration_test.go index 964841a63..862d5d5f2 100644 --- a/management/server/http/handlers/proxy/auth_callback_integration_test.go +++ b/management/server/http/handlers/proxy/auth_callback_integration_test.go @@ -532,8 +532,8 @@ func TestAuthCallback_UserAllowedToLogin(t *testing.T) { wantParam string absentParam string }{ - {name: "legacy proxy", manager: testSessionCodeManager{}, wantParam: "session_token", absentParam: "session_code"}, - {name: "compatible proxy", manager: testSessionCodeManager{supported: true}, wantParam: "session_code", absentParam: "session_token"}, + {name: "legacy proxy", manager: testSessionCodeManager{}, wantParam: "session_token", absentParam: "nb_session_code"}, + {name: "compatible proxy", manager: testSessionCodeManager{supported: true}, wantParam: "nb_session_code", absentParam: "session_token"}, } for _, tt := range tests { @@ -555,8 +555,8 @@ func TestAuthCallback_UserAllowedToLogin(t *testing.T) { require.Empty(t, location.Query().Get(tt.absentParam)) require.Empty(t, location.Query().Get("error")) - if tt.wantParam == "session_code" { - code := location.Query().Get("session_code") + if tt.wantParam == "nb_session_code" { + code := location.Query().Get("nb_session_code") response, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{ Domain: location.Hostname(), SessionCode: code, diff --git a/proxy/auth/auth.go b/proxy/auth/auth.go index 084046c49..605780959 100644 --- a/proxy/auth/auth.go +++ b/proxy/auth/auth.go @@ -30,6 +30,14 @@ const ( SessionJWTIssuer = "netbird-management" ) +// Query parameters management uses to hand the OIDC session to the proxy. The +// proxy strips them before forwarding, so they must not collide with names the +// proxied service uses itself. +const ( + SessionCodeQueryParam = "nb_session_code" + SessionTokenQueryParam = "session_token" +) + // HeaderUserID is the synthetic user id recorded for header-authenticated // requests. Header auth validates a per-service secret and resolves no user // record, so proxy access logs and management-minted session tokens both diff --git a/proxy/internal/auth/middleware.go b/proxy/internal/auth/middleware.go index 672286748..647741139 100644 --- a/proxy/internal/auth/middleware.go +++ b/proxy/internal/auth/middleware.go @@ -583,7 +583,7 @@ func (mw *Middleware) authenticateWithSchemes(w http.ResponseWriter, r *http.Req // handleAuthenticatedToken validates the token, handles denied access, and on // success sets a session cookie and redirects to the original URL. func (mw *Middleware) handleAuthenticatedToken(w http.ResponseWriter, r *http.Request, host, token string, config DomainConfig, scheme Scheme) { - isCode := scheme.Type() == auth.MethodOIDC && r.URL.Query().Get("session_code") != "" + isCode := scheme.Type() == auth.MethodOIDC && r.URL.Query().Get(auth.SessionCodeQueryParam) != "" result, err := mw.validateSessionToken(r.Context(), host, token, isCode, config.SessionPublicKey, scheme.Type()) if err != nil { if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil { @@ -661,7 +661,7 @@ func wasCredentialSubmitted(r *http.Request, method auth.Method) bool { case auth.MethodPassword: return credentialFormValue(r, passwordFormId) != "" case auth.MethodOIDC: - return r.URL.Query().Get("session_token") != "" || r.URL.Query().Get("session_code") != "" + return r.URL.Query().Get(auth.SessionTokenQueryParam) != "" || r.URL.Query().Get(auth.SessionCodeQueryParam) != "" } return false } @@ -806,11 +806,11 @@ func sessionGroupsAllowed(allowed map[string]struct{}, method auth.Method, group // or history. func stripSessionTokenParam(u *url.URL) string { q := u.Query() - if !q.Has("session_token") && !q.Has("session_code") { + if !q.Has(auth.SessionTokenQueryParam) && !q.Has(auth.SessionCodeQueryParam) { return u.RequestURI() } - q.Del("session_token") - q.Del("session_code") + q.Del(auth.SessionTokenQueryParam) + q.Del(auth.SessionCodeQueryParam) clean := *u clean.RawQuery = q.Encode() return clean.RequestURI() diff --git a/proxy/internal/auth/middleware_test.go b/proxy/internal/auth/middleware_test.go index 88c900f97..cce35ae35 100644 --- a/proxy/internal/auth/middleware_test.go +++ b/proxy/internal/auth/middleware_test.go @@ -786,9 +786,15 @@ func TestWasCredentialSubmitted(t *testing.T) { { name: "OIDC code in query", method: auth.MethodOIDC, - query: url.Values{"session_code": {"abc123"}}, + query: url.Values{"nb_session_code": {"abc123"}}, expected: true, }, + { + name: "OIDC backend session_code in query", + method: auth.MethodOIDC, + query: url.Values{"session_code": {"abc123"}}, + expected: false, + }, { name: "OIDC token not in query", method: auth.MethodOIDC, @@ -1585,8 +1591,9 @@ func TestStripSessionTokenParam(t *testing.T) { want string }{ {"strips session_token", "https://ex.com/p?a=1&session_token=tok", "/p?a=1"}, - {"strips session_code", "https://ex.com/p?a=1&session_code=code", "/p?a=1"}, - {"strips both", "https://ex.com/p?session_token=tok&session_code=code&a=1", "/p?a=1"}, + {"strips nb_session_code", "https://ex.com/p?a=1&nb_session_code=code", "/p?a=1"}, + {"strips both", "https://ex.com/p?session_token=tok&nb_session_code=code&a=1", "/p?a=1"}, + {"keeps backend session_code", "https://ex.com/p?a=1&session_code=backend", "/p?a=1&session_code=backend"}, {"no-op when absent", "https://ex.com/p?a=1", "/p?a=1"}, } for _, tc := range cases { diff --git a/proxy/internal/auth/oidc.go b/proxy/internal/auth/oidc.go index 739777924..0215fddc3 100644 --- a/proxy/internal/auth/oidc.go +++ b/proxy/internal/auth/oidc.go @@ -43,12 +43,12 @@ func (o OIDC) Authenticate(r *http.Request) (string, string, error) { // Check for the session credential returned by the OIDC callback. The management // server passes it in the URL because it cannot set a cookie for the proxy's // domain (cookies are domain-scoped per RFC 6265). The current flow uses a - // single-use session_code to keep the durable token out of the URL. + // single-use session code to keep the durable token out of the URL. // session_token remains supported for backward compatibility. - if code := r.URL.Query().Get("session_code"); code != "" { + if code := r.URL.Query().Get(auth.SessionCodeQueryParam); code != "" { return code, "", nil } - if token := r.URL.Query().Get("session_token"); token != "" { + if token := r.URL.Query().Get(auth.SessionTokenQueryParam); token != "" { return token, "", nil } diff --git a/proxy/internal/proxy/reverseproxy.go b/proxy/internal/proxy/reverseproxy.go index 7583b2e01..a3987fe5a 100644 --- a/proxy/internal/proxy/reverseproxy.go +++ b/proxy/internal/proxy/reverseproxy.go @@ -725,9 +725,9 @@ func stripSessionCookie(r *httputil.ProxyRequest) { // from the outgoing URL to prevent credential leakage to backends. func stripSessionTokenQuery(r *httputil.ProxyRequest) { q := r.Out.URL.Query() - if q.Has("session_token") || q.Has("session_code") { - q.Del("session_token") - q.Del("session_code") + if q.Has(auth.SessionTokenQueryParam) || q.Has(auth.SessionCodeQueryParam) { + q.Del(auth.SessionTokenQueryParam) + q.Del(auth.SessionCodeQueryParam) r.Out.URL.RawQuery = q.Encode() } } diff --git a/proxy/internal/proxy/reverseproxy_test.go b/proxy/internal/proxy/reverseproxy_test.go index c0724ce84..b26ca1f9f 100644 --- a/proxy/internal/proxy/reverseproxy_test.go +++ b/proxy/internal/proxy/reverseproxy_test.go @@ -236,6 +236,17 @@ func TestRewriteFunc_SessionTokenQueryStripping(t *testing.T) { "other query parameters must be preserved") }) + t.Run("strips nb_session_code query parameter", func(t *testing.T) { + pr := newProxyRequest(t, "http://example.com/callback?nb_session_code=code123&other=keep", "1.2.3.4:5000") + + rewrite(pr) + + assert.Empty(t, pr.Out.URL.Query().Get("nb_session_code"), + "OIDC session code must be stripped from backend request") + assert.Equal(t, "keep", pr.Out.URL.Query().Get("other"), + "other query parameters must be preserved") + }) + t.Run("preserves query when no session_token present", func(t *testing.T) { pr := newProxyRequest(t, "http://example.com/api?foo=bar&baz=qux", "1.2.3.4:5000") diff --git a/proxy/web/web.go b/proxy/web/web.go index a45fc8730..de3e4771a 100644 --- a/proxy/web/web.go +++ b/proxy/web/web.go @@ -10,6 +10,8 @@ import ( "net/url" "path/filepath" "strings" + + "github.com/netbirdio/netbird/proxy/auth" ) // PathPrefix is the unique URL prefix for serving the proxy's own web assets. @@ -180,8 +182,8 @@ func ServeAccessDeniedPage(w http.ResponseWriter, r *http.Request, code int, tit // stripAuthParams returns the request URI with auth-related query parameters removed. func stripAuthParams(u *url.URL) string { q := u.Query() - q.Del("session_token") - q.Del("session_code") + q.Del(auth.SessionTokenQueryParam) + q.Del(auth.SessionCodeQueryParam) q.Del("error") q.Del("error_description") clean := *u From e2678d4e05e8eda40a74915a4d376b0c4ea49a67 Mon Sep 17 00:00:00 2001 From: Allan ELKAIM Date: Fri, 2 Oct 2026 15:00:26 +0200 Subject: [PATCH 03/14] [management] expose peer MAC addresses and make peers searchable by MAC (#6553) --- .../network_map/controller/repository.go | 2 +- management/internals/modules/peers/manager.go | 2 +- management/server/account.go | 10 +-- management/server/account/manager.go | 2 +- management/server/account/manager_mock.go | 8 +- management/server/account_test.go | 18 ++--- .../handlers/accounts/accounts_handler.go | 2 +- .../http/handlers/groups/groups_handler.go | 10 +-- .../handlers/groups/groups_handler_test.go | 2 +- .../http/handlers/peers/peers_handler.go | 16 +++- .../http/handlers/peers/peers_handler_test.go | 46 ++++++++++- management/server/integrated_validator.go | 2 +- management/server/mock_server/account_mock.go | 6 +- management/server/peer.go | 4 +- management/server/peer_test.go | 76 ++++++++++++++++++- management/server/store/sql_store_peer.go | 7 +- .../server/store/sql_store_peer_test.go | 46 ++++++++++- management/server/store/store.go | 2 +- management/server/store/store_mock.go | 8 +- shared/management/http/api/openapi.yml | 24 ++++++ shared/management/http/api/types.gen.go | 18 +++++ 21 files changed, 265 insertions(+), 46 deletions(-) diff --git a/management/internals/controllers/network_map/controller/repository.go b/management/internals/controllers/network_map/controller/repository.go index 5c3195f16..c11af0b69 100644 --- a/management/internals/controllers/network_map/controller/repository.go +++ b/management/internals/controllers/network_map/controller/repository.go @@ -44,7 +44,7 @@ func (r *repository) GetAccountNetwork(ctx context.Context, accountID string) (* } func (r *repository) GetAccountPeers(ctx context.Context, accountID string) ([]*peer.Peer, error) { - return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") } func (r *repository) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) { diff --git a/management/internals/modules/peers/manager.go b/management/internals/modules/peers/manager.go index 3274ec524..e944be291 100644 --- a/management/internals/modules/peers/manager.go +++ b/management/internals/modules/peers/manager.go @@ -97,7 +97,7 @@ func (m *managerImpl) GetAllPeers(ctx context.Context, accountID, userID string) return m.store.GetUserPeers(ctx, store.LockingStrengthNone, accountID, userID) } - return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") } func (m *managerImpl) GetPeerAccountID(ctx context.Context, peerID string) (string, error) { diff --git a/management/server/account.go b/management/server/account.go index 340bcc84b..1c09c8252 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -2391,7 +2391,7 @@ func (am *DefaultAccountManager) reallocateAccountPeerIPs(ctx context.Context, t return err } - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "") if err != nil { return err } @@ -2428,7 +2428,7 @@ func (am *DefaultAccountManager) reallocateAccountPeerIPs(ctx context.Context, t // v6 address get one allocated. When disabled, all v6 addresses are cleared. // When the v6 range changes, all v6 addresses are reallocated. func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transaction store.Store, accountID, peerID string, newIPv6 netip.Addr) error { - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "", "") if err != nil { return fmt.Errorf("get peers: %w", err) } @@ -2441,7 +2441,7 @@ func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transac } func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) error { - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "") if err != nil { return fmt.Errorf("get peers: %w", err) } @@ -2602,7 +2602,7 @@ func (am *DefaultAccountManager) buildIPv6AllowedPeers(ctx context.Context, tran // Embedded proxy peers sit outside regular group membership but must // participate in any v6-enabled overlay to reach v6-only peers. - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") if err != nil { return nil, fmt.Errorf("get peers: %w", err) } @@ -2673,7 +2673,7 @@ func (am *DefaultAccountManager) updatePeerIPInTransaction(ctx context.Context, return nil } - peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "") + peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "", "") if err != nil { return fmt.Errorf("get account peers: %w", err) } diff --git a/management/server/account/manager.go b/management/server/account/manager.go index 154c9ab18..2ac8584f4 100644 --- a/management/server/account/manager.go +++ b/management/server/account/manager.go @@ -62,7 +62,7 @@ type Manager interface { GetUserByID(ctx context.Context, id string) (*types.User, error) GetUserFromUserAuth(ctx context.Context, userAuth auth.UserAuth) (*types.User, error) ListUsers(ctx context.Context, accountID string) ([]*types.User, error) - GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) + GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) MarkPeerConnected(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error MarkPeerDisconnected(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error DeletePeer(ctx context.Context, accountID, peerID, userID string) error diff --git a/management/server/account/manager_mock.go b/management/server/account/manager_mock.go index f31f63d0e..60075b169 100644 --- a/management/server/account/manager_mock.go +++ b/management/server/account/manager_mock.go @@ -982,18 +982,18 @@ func (mr *MockManagerMockRecorder) GetPeerNetwork(ctx, peerID any) *gomock.Call } // GetPeers mocks base method. -func (m *MockManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*peer.Peer, error) { +func (m *MockManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*peer.Peer, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetPeers", ctx, accountID, userID, nameFilter, ipFilter) + ret := m.ctrl.Call(m, "GetPeers", ctx, accountID, userID, nameFilter, ipFilter, macFilter) ret0, _ := ret[0].([]*peer.Peer) ret1, _ := ret[1].(error) return ret0, ret1 } // GetPeers indicates an expected call of GetPeers. -func (mr *MockManagerMockRecorder) GetPeers(ctx, accountID, userID, nameFilter, ipFilter any) *gomock.Call { +func (mr *MockManagerMockRecorder) GetPeers(ctx, accountID, userID, nameFilter, ipFilter, macFilter any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeers", reflect.TypeOf((*MockManager)(nil).GetPeers), ctx, accountID, userID, nameFilter, ipFilter) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeers", reflect.TypeOf((*MockManager)(nil).GetPeers), ctx, accountID, userID, nameFilter, ipFilter, macFilter) } // GetPolicy mocks base method. diff --git a/management/server/account_test.go b/management/server/account_test.go index 8c735b28e..881ad19d7 100644 --- a/management/server/account_test.go +++ b/management/server/account_test.go @@ -2557,7 +2557,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_PeerApproval(t *testing.T) _, err = manager.UpdateAccountSettings(ctx, accountID, userID, newSettings) require.NoError(t, err) - accountPeers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + accountPeers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) for _, peer := range accountPeers { @@ -4557,7 +4557,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te }) require.NoError(t, err) - peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "") + peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "") require.NoError(t, err) require.Len(t, peers, len(before)) for _, p := range peers { @@ -4575,7 +4575,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te }) require.NoError(t, err) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "") require.NoError(t, err) for _, p := range peers { assert.Equal(t, before[p.ID], p.IP, "peer %s IP should not change for host-bit-set equivalent range", p.ID) @@ -4589,7 +4589,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te }) require.NoError(t, err) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "") require.NoError(t, err) for _, p := range peers { assert.Equal(t, before[p.ID], p.IP, "peer %s IP should not change when NetworkRange omitted", p.ID) @@ -4605,7 +4605,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te }) require.NoError(t, err) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "") require.NoError(t, err) for _, p := range peers { assert.True(t, newRange.Contains(p.IP), "peer %s should be in new range %s, got %s", p.ID, newRange, p.IP) @@ -4623,7 +4623,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin require.NoError(t, err) require.NotEmpty(t, settings.IPv6EnabledGroups, "new account should have IPv6 enabled for All group") - peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) for _, p := range peers { assert.True(t, p.IPv6.IsValid(), "peer %s should have IPv6 with All group enabled", p.ID) @@ -4651,7 +4651,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin assert.Equal(t, []string{partialGroup.ID}, updatedSettings.IPv6EnabledGroups) // peer1 and peer2 should have IPv6; peer3 should not. - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) peerMap := make(map[string]*nbpeer.Peer, len(peers)) for _, p := range peers { @@ -4671,7 +4671,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin require.NoError(t, err) assert.Empty(t, updatedSettings.IPv6EnabledGroups) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) for _, p := range peers { assert.False(t, p.IPv6.IsValid(), "peer %s should have no IPv6 when groups cleared", p.ID) @@ -4686,7 +4686,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin }) require.NoError(t, err) - peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) peerMap = make(map[string]*nbpeer.Peer, len(peers)) for _, p := range peers { diff --git a/management/server/http/handlers/accounts/accounts_handler.go b/management/server/http/handlers/accounts/accounts_handler.go index c4cba5962..795214c31 100644 --- a/management/server/http/handlers/accounts/accounts_handler.go +++ b/management/server/http/handlers/accounts/accounts_handler.go @@ -127,7 +127,7 @@ func (h *handler) validateNetworkRange(ctx context.Context, accountID, userID st } func (h *handler) validateCapacity(ctx context.Context, accountID, userID string, prefix netip.Prefix) error { - peers, err := h.accountManager.GetPeers(ctx, accountID, userID, "", "") + peers, err := h.accountManager.GetPeers(ctx, accountID, userID, "", "", "") if err != nil { return status.Errorf(status.Internal, "get peer count: %v", err) } diff --git a/management/server/http/handlers/groups/groups_handler.go b/management/server/http/handlers/groups/groups_handler.go index ed01e7c3d..1a7753a57 100644 --- a/management/server/http/handlers/groups/groups_handler.go +++ b/management/server/http/handlers/groups/groups_handler.go @@ -58,7 +58,7 @@ func (h *handler) getAllGroups(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -77,7 +77,7 @@ func (h *handler) getAllGroups(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -169,7 +169,7 @@ func (h *handler) updateGroup(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -226,7 +226,7 @@ func (h *handler) createGroup(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return @@ -287,7 +287,7 @@ func (h *handler) getGroup(w http.ResponseWriter, r *http.Request) { return } - accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "") + accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "") if err != nil { util.WriteError(r.Context(), err, w) return diff --git a/management/server/http/handlers/groups/groups_handler_test.go b/management/server/http/handlers/groups/groups_handler_test.go index 78e4a2578..3e322db4e 100644 --- a/management/server/http/handlers/groups/groups_handler_test.go +++ b/management/server/http/handlers/groups/groups_handler_test.go @@ -78,7 +78,7 @@ func initGroupTestData(initGroups ...*types.Group) *handler { return nil, status.Errorf(status.NotFound, "unknown group name") }, - GetPeersFunc: func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { + GetPeersFunc: func(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { return maps.Values(TestPeers), nil }, DeleteGroupFunc: func(_ context.Context, accountID, userId, groupID string) error { diff --git a/management/server/http/handlers/peers/peers_handler.go b/management/server/http/handlers/peers/peers_handler.go index 773b640e0..8a9bf1f70 100644 --- a/management/server/http/handlers/peers/peers_handler.go +++ b/management/server/http/handlers/peers/peers_handler.go @@ -317,10 +317,11 @@ func (h *Handler) GetAllPeers(w http.ResponseWriter, r *http.Request) { nameFilter := r.URL.Query().Get("name") ipFilter := r.URL.Query().Get("ip") + macFilter := r.URL.Query().Get("mac") accountID, userID := userAuth.AccountId, userAuth.UserId - peers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, nameFilter, ipFilter) + peers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, nameFilter, ipFilter, macFilter) if err != nil { util.WriteError(r.Context(), err, w) return @@ -571,6 +572,17 @@ func peerToAccessiblePeer(peer *nbpeer.Peer, dnsDomain string) api.AccessiblePee } } +func toNetworkAddresses(addrs []nbpeer.NetworkAddress) *[]api.NetworkAddress { + if len(addrs) == 0 { + return nil + } + out := make([]api.NetworkAddress, 0, len(addrs)) + for _, a := range addrs { + out = append(out, api.NetworkAddress{NetIp: a.NetIP.String(), Mac: a.Mac}) + } + return &out +} + func toSinglePeerResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dnsDomain string, approved bool, reason string) *api.Peer { osVersion := peer.Meta.OSVersion if osVersion == "" { @@ -583,6 +595,7 @@ func toSinglePeerResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dnsD Name: peer.Name, Ip: peer.IP.String(), Ipv6: peerIPv6String(peer), + NetworkAddresses: toNetworkAddresses(peer.Meta.NetworkAddresses), ConnectionIp: peer.Location.ConnectionIP.String(), Connected: peer.Status.Connected, LastSeen: peer.Status.LastSeen, @@ -639,6 +652,7 @@ func toPeerListItemResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dn Name: peer.Name, Ip: peer.IP.String(), Ipv6: peerIPv6String(peer), + NetworkAddresses: toNetworkAddresses(peer.Meta.NetworkAddresses), ConnectionIp: peer.Location.ConnectionIP.String(), Connected: peer.Status.Connected, LastSeen: peer.Status.LastSeen, diff --git a/management/server/http/handlers/peers/peers_handler_test.go b/management/server/http/handlers/peers/peers_handler_test.go index 592d64d1a..7054082cc 100644 --- a/management/server/http/handlers/peers/peers_handler_test.go +++ b/management/server/http/handlers/peers/peers_handler_test.go @@ -173,7 +173,7 @@ func initTestMetaData(t *testing.T, peers ...*nbpeer.Peer) *Handler { return nil, fmt.Errorf("user not found") } }, - GetPeersFunc: func(_ context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { + GetPeersFunc: func(_ context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { return peers, nil }, GetPeerGroupsFunc: func(ctx context.Context, accountID, peerID string) ([]*types.Group, error) { @@ -364,6 +364,50 @@ func TestGetPeers(t *testing.T) { } } +func TestPeerResponseNetworkAddresses(t *testing.T) { + tests := []struct { + name string + addresses []nbpeer.NetworkAddress + wantJSON string + }{ + {name: "not reported"}, + {name: "empty", addresses: []nbpeer.NetworkAddress{}}, + { + name: "multiple interfaces", + addresses: []nbpeer.NetworkAddress{ + {NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"}, + {NetIP: netip.MustParsePrefix("2001:db8::123/64"), Mac: "00:93:37:bd:83:10"}, + }, + wantJSON: `[{"net_ip":"192.168.0.11/24","mac":"00:93:37:bd:83:0f"},{"net_ip":"2001:db8::123/64","mac":"00:93:37:bd:83:10"}]`, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peer := &nbpeer.Peer{ + Status: &nbpeer.PeerStatus{}, + Meta: nbpeer.PeerSystemMeta{NetworkAddresses: tt.addresses}, + } + responses := map[string]any{ + "single peer": toSinglePeerResponse(peer, nil, "example.com", true, ""), + "peer list": toPeerListItemResponse(peer, nil, "example.com", 0), + } + for name, response := range responses { + t.Run(name, func(t *testing.T) { + body, err := json.Marshal(response) + require.NoError(t, err) + var fields map[string]json.RawMessage + require.NoError(t, json.Unmarshal(body, &fields)) + if tt.wantJSON == "" { + assert.NotContains(t, fields, "network_addresses", "unreported interfaces should be omitted") + return + } + assert.JSONEq(t, tt.wantJSON, string(fields["network_addresses"]), "response should preserve interface addresses and MACs") + }) + } + }) + } +} + func TestGetAccessiblePeers(t *testing.T) { peer1 := &nbpeer.Peer{ ID: "peer1", diff --git a/management/server/integrated_validator.go b/management/server/integrated_validator.go index 9ec1f491e..5928a8ed2 100644 --- a/management/server/integrated_validator.go +++ b/management/server/integrated_validator.go @@ -100,7 +100,7 @@ func (am *DefaultAccountManager) GetValidatedPeers(ctx context.Context, accountI return nil, nil, err } - peers, err = am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "") + peers, err = am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "") if err != nil { return nil, nil, err } diff --git a/management/server/mock_server/account_mock.go b/management/server/mock_server/account_mock.go index 2f871c3e2..3313bf99c 100644 --- a/management/server/mock_server/account_mock.go +++ b/management/server/mock_server/account_mock.go @@ -39,7 +39,7 @@ type MockAccountManager struct { GetAccountIDByUserIdFunc func(ctx context.Context, userAuth auth.UserAuth) (string, error) GetUserFromUserAuthFunc func(ctx context.Context, userAuth auth.UserAuth) (*types.User, error) ListUsersFunc func(ctx context.Context, accountID string) ([]*types.User, error) - GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) + GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) MarkPeerConnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error MarkPeerDisconnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error SyncAndMarkPeerFunc func(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error) @@ -807,9 +807,9 @@ func (am *MockAccountManager) GetAccountIDFromUserAuth(ctx context.Context, user } // GetPeers mocks GetPeers of the AccountManager interface -func (am *MockAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { +func (am *MockAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { if am.GetPeersFunc != nil { - return am.GetPeersFunc(ctx, accountID, userID, nameFilter, ipFilter) + return am.GetPeersFunc(ctx, accountID, userID, nameFilter, ipFilter, macFilter) } return nil, status.Errorf(codes.Unimplemented, "method GetPeers is not implemented") } diff --git a/management/server/peer.go b/management/server/peer.go index 9f5572252..5d5863fa7 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -47,7 +47,7 @@ const ( // GetPeers returns peers visible to the user within an account. // Users with "peers:read" see all peers. Otherwise, users see only their own peers, or none if restricted by account settings. -func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { +func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { user, err := am.Store.GetUserByUserID(ctx, store.LockingStrengthNone, userID) if err != nil { return nil, err @@ -59,7 +59,7 @@ func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID } if allowed { - return am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, nameFilter, ipFilter) + return am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, nameFilter, ipFilter, macFilter) } settings, err := am.Store.GetAccountSettings(ctx, store.LockingStrengthNone, accountID) diff --git a/management/server/peer_test.go b/management/server/peer_test.go index 22f2b9b6f..5c3e02af5 100644 --- a/management/server/peer_test.go +++ b/management/server/peer_test.go @@ -4,10 +4,14 @@ import ( "context" "crypto/sha256" b64 "encoding/base64" + "encoding/json" "fmt" "io" "net" + "net/http" + "net/http/httptest" "net/netip" + "net/url" "os" "runtime" "strconv" @@ -33,12 +37,15 @@ import ( "github.com/netbirdio/netbird/management/internals/server/config" "github.com/netbirdio/netbird/management/internals/shared/grpc" nbcache "github.com/netbirdio/netbird/management/server/cache" + nbcontext "github.com/netbirdio/netbird/management/server/context" + peershandler "github.com/netbirdio/netbird/management/server/http/handlers/peers" "github.com/netbirdio/netbird/management/server/http/testing/testing_tools" "github.com/netbirdio/netbird/management/server/integrations/port_forwarding" "github.com/netbirdio/netbird/management/server/job" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/settings" "github.com/netbirdio/netbird/shared/auth" + "github.com/netbirdio/netbird/shared/management/http/api" "github.com/netbirdio/netbird/shared/management/status" "github.com/netbirdio/netbird/management/server/util" @@ -718,7 +725,7 @@ func TestDefaultAccountManager_GetPeers(t *testing.T) { return } - peers, err := manager.GetPeers(context.Background(), accountID, someUser, "", "") + peers, err := manager.GetPeers(context.Background(), accountID, someUser, "", "", "") if err != nil { t.Fatal(err) return @@ -731,6 +738,71 @@ func TestDefaultAccountManager_GetPeers(t *testing.T) { } } +func TestDefaultAccountManager_GetPeers_FilterByMac(t *testing.T) { + ctx := context.Background() + manager, _, err := createManager(t) + require.NoError(t, err) + account := newAccountWithId(ctx, "mac-account", "mac-admin", "", "", "", false) + account.Peers["matching"] = &nbpeer.Peer{ + ID: "matching", Key: "matching-key", Name: "laptop", DNSLabel: "laptop", + IP: netip.MustParseAddr("100.64.0.10"), Status: &nbpeer.PeerStatus{LastSeen: time.Now()}, + Meta: nbpeer.PeerSystemMeta{NetworkAddresses: []nbpeer.NetworkAddress{ + {NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"}, + {NetIP: netip.MustParsePrefix("192.168.1.11/24"), Mac: "aa:bb:cc:dd:ee:ff"}, + }}, + } + account.Peers["other"] = &nbpeer.Peer{ + ID: "other", Key: "other-key", Name: "desktop", DNSLabel: "desktop", + IP: netip.MustParseAddr("100.64.0.20"), Status: &nbpeer.PeerStatus{LastSeen: time.Now()}, + } + require.NoError(t, manager.Store.SaveAccount(ctx, account)) + otherAccount := newAccountWithId(ctx, "other-account", "other-admin", "", "", "", false) + otherPeer := account.Peers["matching"].Copy() + otherPeer.ID, otherPeer.Key = "outside-account", "outside-key" + otherAccount.Peers[otherPeer.ID] = otherPeer + require.NoError(t, manager.Store.SaveAccount(ctx, otherAccount)) + handler := peershandler.NewHandler(manager, manager.networkMapController, manager.permissionsManager) + + tests := []struct { + name, nameFilter, ipFilter, macFilter string + wantIDs []string + }{ + {name: "no filter", wantIDs: []string{"matching", "other"}}, + {name: "full MAC", macFilter: "00:93:37:bd:83:0f", wantIDs: []string{"matching"}}, + {name: "partial MAC", macFilter: "93:37:bd", wantIDs: []string{"matching"}}, + {name: "second interface", macFilter: "aa:bb:cc:dd:ee:ff", wantIDs: []string{"matching"}}, + {name: "unknown MAC", macFilter: "11:22:33:44:55:66"}, + {name: "combined filters", nameFilter: "laptop", ipFilter: "100.64.0.10", macFilter: "00:93:37", wantIDs: []string{"matching"}}, + {name: "name mismatch", nameFilter: "desktop", macFilter: "00:93:37"}, + {name: "IP mismatch", ipFilter: "100.64.0.20", macFilter: "00:93:37"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := manager.GetPeers(ctx, account.Id, "mac-admin", tt.nameFilter, tt.ipFilter, tt.macFilter) + require.NoError(t, err) + ids := make([]string, 0, len(peers)) + for _, peer := range peers { + ids = append(ids, peer.ID) + } + assert.ElementsMatch(t, tt.wantIDs, ids, "filters should return only matching peers in the account") + + query := url.Values{"name": {tt.nameFilter}, "ip": {tt.ipFilter}, "mac": {tt.macFilter}} + req := httptest.NewRequest(http.MethodGet, "/api/peers?"+query.Encode(), nil) + req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: account.Id, UserId: "mac-admin"}) + recorder := httptest.NewRecorder() + handler.GetAllPeers(recorder, req) + require.Equal(t, http.StatusOK, recorder.Code, "peer listing should succeed: %s", recorder.Body.String()) + var response []api.PeerBatch + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response)) + responseIDs := make([]string, 0, len(response)) + for _, peer := range response { + responseIDs = append(responseIDs, peer.Id) + } + assert.ElementsMatch(t, tt.wantIDs, responseIDs, "HTTP query filters should reach the store") + }) + } +} + func setupTestAccountManager(b testing.TB, peers int, groups int) (*DefaultAccountManager, *update_channel.PeersUpdateManager, string, string, error) { b.Helper() @@ -934,7 +1006,7 @@ func BenchmarkGetPeers(b *testing.B) { b.ResetTimer() for i := 0; i < b.N; i++ { - _, err := manager.GetPeers(context.Background(), accountID, userID, "", "") + _, err := manager.GetPeers(context.Background(), accountID, userID, "", "", "") if err != nil { b.Fatalf("GetPeers failed: %v", err) } diff --git a/management/server/store/sql_store_peer.go b/management/server/store/sql_store_peer.go index e5086b6db..1b0e23cec 100644 --- a/management/server/store/sql_store_peer.go +++ b/management/server/store/sql_store_peer.go @@ -492,7 +492,7 @@ func (s *SqlStore) GetPeerByPeerPubKey(ctx context.Context, lockStrength Locking } // GetAccountPeers retrieves peers for an account. -func (s *SqlStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) { +func (s *SqlStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) { var peers []*nbpeer.Peer tx := s.db if lockStrength != LockingStrengthNone { @@ -506,6 +506,11 @@ func (s *SqlStore) GetAccountPeers(ctx context.Context, lockStrength LockingStre if ipFilter != "" { query = query.Where("ip LIKE ? OR ipv6 LIKE ?", "%"+ipFilter+"%", "%"+ipFilter+"%") } + // MAC addresses live in the JSON-serialized meta_network_addresses column, + // so we match the raw JSON text rather than a dedicated column. + if macFilter != "" { + query = query.Where("meta_network_addresses LIKE ?", "%"+macFilter+"%") + } if err := query.Find(&peers).Error; err != nil { log.WithContext(ctx).Errorf("failed to get peers from the store: %s", err) diff --git a/management/server/store/sql_store_peer_test.go b/management/server/store/sql_store_peer_test.go index b49e04f2f..1432b5d96 100644 --- a/management/server/store/sql_store_peer_test.go +++ b/management/server/store/sql_store_peer_test.go @@ -512,7 +512,7 @@ func TestSqlStore_GetAccountPeers(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - peers, err := store.GetAccountPeers(context.Background(), LockingStrengthNone, tt.accountID, tt.nameFilter, tt.ipFilter) + peers, err := store.GetAccountPeers(context.Background(), LockingStrengthNone, tt.accountID, tt.nameFilter, tt.ipFilter, "") require.NoError(t, err) require.Len(t, peers, tt.expectedCount) }) @@ -520,6 +520,48 @@ func TestSqlStore_GetAccountPeers(t *testing.T) { } +func TestSqlStore_GetAccountPeers_FilterByMac(t *testing.T) { + ctx := context.Background() + store, cleanup, err := NewTestStoreFromSQL(ctx, "", t.TempDir()) + require.NoError(t, err) + t.Cleanup(cleanup) + + accountID := "test-account-mac" + userID := "test-user-mac" + account := newAccountWithId(ctx, accountID, userID, "example.com") + account.Peers["peer-mac-1"] = &nbpeer.Peer{ + ID: "peer-mac-1", + AccountID: accountID, + Key: "peer-mac-key-1", + Name: "macpeer", + IP: netip.MustParseAddr("100.64.0.10"), + Meta: nbpeer.PeerSystemMeta{ + NetworkAddresses: []nbpeer.NetworkAddress{ + {NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"}, + }, + }, + } + require.NoError(t, store.SaveAccount(ctx, account)) + + tests := []struct { + name string + macFilter string + expectedCount int + }{ + {name: "full mac matches", macFilter: "00:93:37:bd:83:0f", expectedCount: 1}, + {name: "mac prefix matches", macFilter: "00:93:37", expectedCount: 1}, + {name: "unknown mac does not match", macFilter: "11:22:33:44:55:66", expectedCount: 0}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + peers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "", tt.macFilter) + require.NoError(t, err) + require.Len(t, peers, tt.expectedCount) + }) + } +} + func TestSqlStore_GetAccountPeersWithExpiration(t *testing.T) { store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir()) t.Cleanup(cleanup) @@ -878,7 +920,7 @@ func TestSqlStore_ApproveAccountPeers(t *testing.T) { require.NoError(t, err) assert.Equal(t, 2, count) - allPeers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "") + allPeers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "", "") require.NoError(t, err) for _, peer := range allPeers { diff --git a/management/server/store/store.go b/management/server/store/store.go index 465f84413..177c2a47c 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -160,7 +160,7 @@ type Store interface { RemoveResourceFromGroup(ctx context.Context, accountId string, groupID string, resourceID string) error AddPeerToAccount(ctx context.Context, peer *nbpeer.Peer) error GetPeerByPeerPubKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (*nbpeer.Peer, error) - GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) + GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) GetUserPeers(ctx context.Context, lockStrength LockingStrength, accountID, userID string) ([]*nbpeer.Peer, error) GetPeerByID(ctx context.Context, lockStrength LockingStrength, accountID string, peerID string) (*nbpeer.Peer, error) GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*nbpeer.Peer, error) diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index cd9e7334d..53e35b866 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -1330,18 +1330,18 @@ func (mr *MockStoreMockRecorder) GetAccountOwner(ctx, lockStrength, accountID an } // GetAccountPeers mocks base method. -func (m *MockStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*peer.Peer, error) { +func (m *MockStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*peer.Peer, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetAccountPeers", ctx, lockStrength, accountID, nameFilter, ipFilter) + ret := m.ctrl.Call(m, "GetAccountPeers", ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter) ret0, _ := ret[0].([]*peer.Peer) ret1, _ := ret[1].(error) return ret0, ret1 } // GetAccountPeers indicates an expected call of GetAccountPeers. -func (mr *MockStoreMockRecorder) GetAccountPeers(ctx, lockStrength, accountID, nameFilter, ipFilter any) *gomock.Call { +func (mr *MockStoreMockRecorder) GetAccountPeers(ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountPeers", reflect.TypeOf((*MockStore)(nil).GetAccountPeers), ctx, lockStrength, accountID, nameFilter, ipFilter) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountPeers", reflect.TypeOf((*MockStore)(nil).GetAccountPeers), ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter) } // GetAccountPeersWithExpiration mocks base method. diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index 3dd9f41f1..4b7077cac 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -826,6 +826,20 @@ components: - ssh_enabled - login_expiration_enabled - inactivity_expiration_enabled + NetworkAddress: + type: object + properties: + net_ip: + description: IP address with CIDR of the interface + type: string + example: 192.168.0.11/24 + mac: + description: MAC address of the interface + type: string + example: "00:93:37:bd:83:0f" + required: + - net_ip + - mac Peer: allOf: - $ref: '#/components/schemas/PeerMinimum' @@ -845,6 +859,11 @@ components: type: string format: ipv6 example: "fd00:4e42:ab12::1" + network_addresses: + description: Network interfaces (IP + MAC) reported by the peer + type: array + items: + $ref: '#/components/schemas/NetworkAddress' connection_ip: description: Peer's public connection IP address type: string @@ -7516,6 +7535,11 @@ paths: schema: type: string description: Filter peers by IP address + - in: query + name: mac + schema: + type: string + description: Filter peers by MAC address of a network interface security: - BearerAuth: [ ] - TokenAuth: [ ] diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index 9a90a72d3..009a9a7a7 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -3829,6 +3829,15 @@ type Network struct { RoutingPeersCount int `json:"routing_peers_count"` } +// NetworkAddress defines model for NetworkAddress. +type NetworkAddress struct { + // Mac MAC address of the interface + Mac string `json:"mac"` + + // NetIp IP address with CIDR of the interface + NetIp string `json:"net_ip"` +} + // NetworkRequest defines model for NetworkRequest. type NetworkRequest struct { // Description Network description @@ -4278,6 +4287,9 @@ type Peer struct { // Name Peer's hostname Name string `json:"name"` + // NetworkAddresses Network interfaces (IP + MAC) reported by the peer + NetworkAddresses *[]NetworkAddress `json:"network_addresses,omitempty"` + // Os Peer's operating system and version Os string `json:"os"` @@ -4372,6 +4384,9 @@ type PeerBatch struct { // Name Peer's hostname Name string `json:"name"` + // NetworkAddresses Network interfaces (IP + MAC) reported by the peer + NetworkAddresses *[]NetworkAddress `json:"network_addresses,omitempty"` + // Os Peer's operating system and version Os string `json:"os"` @@ -6294,6 +6309,9 @@ type GetApiPeersParams struct { // Ip Filter peers by IP address Ip *string `form:"ip,omitempty" json:"ip,omitempty"` + + // Mac Filter peers by MAC address of a network interface + Mac *string `form:"mac,omitempty" json:"mac,omitempty"` } // GetApiPeersPeerIdIngressPortsParams defines parameters for GetApiPeersPeerIdIngressPorts. From 1b89880e30fa542ae6ff93fd1423e8a923b98e8d Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Fri, 2 Oct 2026 15:11:38 +0200 Subject: [PATCH 04/14] [misc] Move the FreeBSD port test to release 15.1 (#7999) * [misc] Move the FreeBSD port test to release 15.1 FreeBSD 15.0 reached end of life on 2026-09-30 and the ports tree marks it unsupported since freebsd/freebsd-ports@ed90b23fe9 (2026-10-01), so `make package` refuses to run on the 15.0 VM and the FreeBSD Port job fails on every PR. The pinned vmactions/freebsd-vm v1.4.8 ships a 15.1 image, so only the release needs to move. * [misc] Run the FreeBSD unit tests on release 15.1 too The job installs binary packages instead of building from the ports tree, so it kept passing on the EOL 15.0 image, but the client should be tested on the same supported release the port is built on, and the EOL image is only served from the archive mirror from now on. --- .github/workflows/golang-test-freebsd.yml | 2 +- .github/workflows/release.yml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/golang-test-freebsd.yml b/.github/workflows/golang-test-freebsd.yml index 65c39147a..7bd48e3d0 100644 --- a/.github/workflows/golang-test-freebsd.yml +++ b/.github/workflows/golang-test-freebsd.yml @@ -33,7 +33,7 @@ jobs: with: usesh: true copyback: false - release: "15.0" + release: "15.1" envs: "GO_VERSION" prepare: | pkg install -y curl pkgconf xorg diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 673fcc281..9c9af5f17 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -69,7 +69,7 @@ jobs: with: usesh: true copyback: false - release: "15.0" + release: "15.1" envs: "GO_VERSION" prepare: | # Install required packages From 9f8ddc71315bc40f5c98dca82cd902ec1a9f59dd Mon Sep 17 00:00:00 2001 From: Riccardo Manfrin <3090891+riccardomanfrin@users.noreply.github.com> Date: Fri, 2 Oct 2026 15:24:39 +0200 Subject: [PATCH 05/14] [client] Discover interfaces lazily in stdnet instead of at construction (#7346) * [client] Discover interfaces lazily in stdnet instead of at construction stdnet.NewNet and NewNetWithDiscover ended with return n, n.UpdateInterfaces() handing back a non-nil *Net together with the discovery error. Three of the five call sites (Engine.newWgIface, ice.NewAgent, SingleSocketUDPMux) logged the error and kept using the instance, which is only safe as long as the instance still works after a failed discovery. That stopped being true when Interfaces() gained a lazily refreshed cache: updateInterfaces sets lastUpdate only on success, so after a failed construction the 30s cache guard never holds and Interfaces() returns an error rather than the empty list it used to return. Feeding such an instance to pion is worse than passing nothing at all - ice.NewAgent falls back to its own stdnet when Net is nil, and the interface blacklist is applied separately through AgentConfig.InterfaceFilter, so the fallback loses nothing. Instead, a transient discovery failure (the Android bridge at boot, or an interface disappearing between net.Interfaces() and Interface.Addrs()) turned into a hard "error getting local interfaces" from ice.NewAgent, and aborted the STUN and TURN probes, which never even need the interface list. Since the accessors already refresh a stale cache on demand, the eager discovery in the constructors is redundant: drop it, make both constructors infallible, and let the discovery error surface at the call that actually needs the interfaces. UpdateInterfaces had no callers left and is not part of transport.Net, so it is removed along with it. InterfaceByIndex and InterfaceByName read the cached slice directly and never refreshed it, so they would have kept reporting ErrInterfaceNotFound forever on an instance whose first discovery failed. They now go through the same refresh path as Interfaces(). * [client] Warm the stdnet interface cache at construction Moving discovery to first use regressed the privileged suites on the three platforms that always build an ICE bind: Darwin, FreeBSD and Windows time out in TestWGIface_UpdateAddr, TestRecreation, TestEngine_SSH and TestEngine_MultiplePeers, while Linux stays green because a host with the WireGuard kernel module takes the kernel-device branch and never drives the mux that asks for interfaces. interfaceFilter probes with wgctrl every interface the disallow list does not already exclude. Discovering at construction ran that probe before the caller had an overlay interface of its own; discovering at first use runs it after, so on a userspace WireGuard platform the probe reaches the UAPI socket of the same process. The tests reach it because they construct with a nil disallow list, where the client passes DefaultInterfaceBlacklist and its own interface is excluded by prefix. Restore the original timing with an explicit warm-up. The constructors stay infallible and the error is still reported by the accessor that needs the interfaces, so the contract this branch is about is unchanged. * Revert "[client] Warm the stdnet interface cache at construction" This reverts commit 947e25288f78b1afb8b2cb5a0d21925ba694fe0b. * [client] Give the privileged tests the interface blacklist the client uses The suites that create a WireGuard interface construct stdnet with a nil disallow list, which the client never does: Engine passes profilemanager.DefaultInterfaceBlacklist, whose "wt" and "utun" prefixes exclude the overlay interface before the filter reaches its wgctrl probe. With an empty list every interface reaches that probe, the one the test has just created included, and on a userspace WireGuard platform the probe talks to the UAPI socket of the same process. That is why Darwin, FreeBSD and Windows timed out here while Linux, which takes the kernel-device branch on a host with the module loaded, stayed green. Pass the blacklist in both suites so they exercise the configuration the client ships. client/iface declares the prefixes locally because profilemanager imports it. Also cover the constructors directly: the existing tests build the struct literal, so nothing asserted that NewNet and NewNetWithDiscover leave the cache cold. * [client] Pass the blacklist in the remaining tests that build an interface Same reason as the previous commit, four call sites it missed: engine_test, the route manager and systemops suites, and the privileged DNS server suite all construct stdnet with a nil disallow list and then create a WireGuard interface. TestAddVPNRoute surfaced it on FreeBSD once the earlier two files stopped timing out first. client/internal/dns declares the prefixes locally; profilemanager imports that package, so it cannot import profilemanager back. --- client/iface/iface_test.go | 52 +++---- client/iface/udpmux/mux.go | 5 +- client/internal/dns/server_privileged_test.go | 17 +-- client/internal/dns/server_test.go | 6 +- client/internal/engine.go | 5 +- client/internal/engine_privileged_test.go | 26 ++-- client/internal/engine_stdnet.go | 2 +- client/internal/engine_stdnet_android.go | 2 +- client/internal/engine_test.go | 11 +- client/internal/peer/ice/agent.go | 5 +- client/internal/peer/ice/stdnet.go | 2 +- client/internal/peer/ice/stdnet_android.go | 2 +- client/internal/relay/relay.go | 12 +- client/internal/routemanager/manager_test.go | 6 +- .../systemops/systemops_generic_test.go | 4 +- client/internal/stdnet/stdnet.go | 85 ++++++----- client/internal/stdnet/stdnet_test.go | 136 ++++++++++++++++++ 17 files changed, 237 insertions(+), 141 deletions(-) create mode 100644 client/internal/stdnet/stdnet_test.go diff --git a/client/iface/iface_test.go b/client/iface/iface_test.go index fff0d4e30..cb50ca4a1 100644 --- a/client/iface/iface_test.go +++ b/client/iface/iface_test.go @@ -40,14 +40,18 @@ func init() { peerPubKey = peerPrivateKey.PublicKey().String() } +// testIFaceBlackList mirrors the prefixes profilemanager.DefaultInterfaceBlacklist +// carries for the overlay interface. These tests create their own utun device, and +// stdnet's filter probes with wgctrl every interface it is not told to skip, which +// on a userspace WireGuard platform reaches the UAPI socket of this same process. +// Declared here rather than imported because profilemanager imports this package. +var testIFaceBlackList = []string{"wt", "utun", "tun0"} + func TestWGIface_UpdateAddr(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) addr := "100.64.0.1/8" wgPort := 33100 - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -127,10 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) { func Test_CreateInterface(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1) wgIP := "10.99.99.1/32" - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, Address: wgaddr.MustParseWGAddress(wgIP), @@ -170,10 +171,7 @@ func Test_Close(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) wgIP := "10.99.99.2/32" wgPort := 33100 - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -215,10 +213,7 @@ func TestRecreation(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2) wgIP := "10.99.99.2/32" wgPort := 33100 - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -288,10 +283,7 @@ func Test_ConfigureInterface(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3) wgIP := "10.99.99.5/30" wgPort := 33100 - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, Address: wgaddr.MustParseWGAddress(wgIP), @@ -343,10 +335,7 @@ func Test_ConfigureInterface(t *testing.T) { func Test_UpdatePeer(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) wgIP := "10.99.99.9/30" - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -413,10 +402,7 @@ func Test_UpdatePeer(t *testing.T) { func Test_RemovePeer(t *testing.T) { ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4) wgIP := "10.99.99.13/30" - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := WGIFaceOpts{ IFaceName: ifaceName, @@ -477,10 +463,7 @@ func Test_ConnectPeers(t *testing.T) { peer2wgPort := 33200 keepAlive := 1 * time.Second - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) guid := fmt.Sprintf("{%s}", uuid.New().String()) device.CustomWindowsGUIDString = strings.ToLower(guid) @@ -516,10 +499,7 @@ func Test_ConnectPeers(t *testing.T) { guid = fmt.Sprintf("{%s}", uuid.New().String()) device.CustomWindowsGUIDString = strings.ToLower(guid) - newNet, err = stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet = stdnet.NewNet(context.Background(), testIFaceBlackList) optsPeer2 := WGIFaceOpts{ IFaceName: peer2ifaceName, diff --git a/client/iface/udpmux/mux.go b/client/iface/udpmux/mux.go index c5d2de4a5..68cecc953 100644 --- a/client/iface/udpmux/mux.go +++ b/client/iface/udpmux/mux.go @@ -200,10 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() { } if len(networks) > 0 { if m.params.Net == nil { - var err error - if m.params.Net, err = stdnet.NewNet(context.Background(), nil); err != nil { - m.params.Logger.Errorf("failed to get create network: %v", err) - } + m.params.Net = stdnet.NewNet(context.Background(), nil) } ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true) diff --git a/client/internal/dns/server_privileged_test.go b/client/internal/dns/server_privileged_test.go index a17044cf5..270e3bf91 100644 --- a/client/internal/dns/server_privileged_test.go +++ b/client/internal/dns/server_privileged_test.go @@ -9,9 +9,9 @@ import ( "os" "testing" - "go.uber.org/mock/gomock" "github.com/miekg/dns" "github.com/stretchr/testify/assert" + "go.uber.org/mock/gomock" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "github.com/netbirdio/netbird/client/iface" @@ -24,6 +24,10 @@ import ( nbdns "github.com/netbirdio/netbird/dns" ) +// testIFaceBlackList mirrors the overlay prefixes profilemanager.DefaultInterfaceBlacklist +// carries. Declared here rather than imported because profilemanager imports this package. +var testIFaceBlackList = []string{"wt", "utun", "tun0"} + func TestUpdateDNSServer(t *testing.T) { nameServers := []nbdns.NameServer{ @@ -243,10 +247,7 @@ func TestUpdateDNSServer(t *testing.T) { for n, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { privKey, _ := wgtypes.GenerateKey() - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), testIFaceBlackList) opts := iface.WGIFaceOpts{ IFaceName: fmt.Sprintf("utun230%d", n), @@ -348,11 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) { defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) t.Setenv("NB_WG_KERNEL_DISABLED", "true") - newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"}) - if err != nil { - t.Errorf("create stdnet: %v", err) - return - } + newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}) privKey, _ := wgtypes.GeneratePrivateKey() opts := iface.WGIFaceOpts{ diff --git a/client/internal/dns/server_test.go b/client/internal/dns/server_test.go index 0144a4a8b..414890158 100644 --- a/client/internal/dns/server_test.go +++ b/client/internal/dns/server_test.go @@ -394,11 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) { defer t.Setenv("NB_WG_KERNEL_DISABLED", ov) t.Setenv("NB_WG_KERNEL_DISABLED", "true") - newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"}) - if err != nil { - t.Fatalf("create stdnet: %v", err) - return nil, err - } + newNet := stdnet.NewNet(context.Background(), []string{"utun2301"}) privKey, _ := wgtypes.GeneratePrivateKey() diff --git a/client/internal/engine.go b/client/internal/engine.go index fc7dce869..7e9375771 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -2178,10 +2178,7 @@ func (e *Engine) close() { } func (e *Engine) newWgIface() (*iface.WGIface, error) { - transportNet, err := e.newStdNet() - if err != nil { - log.Errorf("failed to create pion's stdnet: %s", err) - } + transportNet := e.newStdNet() opts := iface.WGIFaceOpts{ IFaceName: e.config.WgIfaceName, diff --git a/client/internal/engine_privileged_test.go b/client/internal/engine_privileged_test.go index 1b047e017..2db0cd5ed 100644 --- a/client/internal/engine_privileged_test.go +++ b/client/internal/engine_privileged_test.go @@ -12,12 +12,12 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/google/uuid" log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.opentelemetry.io/otel" + "go.uber.org/mock/gomock" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" "google.golang.org/grpc" "google.golang.org/grpc/keepalive" @@ -27,6 +27,7 @@ import ( "github.com/netbirdio/netbird/client/iface/wgaddr" "github.com/netbirdio/netbird/client/internal/dns" "github.com/netbirdio/netbird/client/internal/peer" + "github.com/netbirdio/netbird/client/internal/profilemanager" nbssh "github.com/netbirdio/netbird/client/ssh" "github.com/netbirdio/netbird/client/system" nbdns "github.com/netbirdio/netbird/dns" @@ -81,6 +82,7 @@ func TestEngine_SSH(t *testing.T) { WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), WgPrivateKey: key, WgPort: 33100, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, ServerSSHAllowed: true, MTU: iface.DefaultMTU, SSHKey: sshKey, @@ -204,11 +206,12 @@ func TestEngine_Sync(t *testing.T) { } relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU) engine := NewEngine(ctx, cancel, &EngineConfig{ - WgIfaceName: "utun103", - WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), - WgPrivateKey: key, - WgPort: 33100, - MTU: iface.DefaultMTU, + WgIfaceName: "utun103", + WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"), + WgPrivateKey: key, + WgPort: 33100, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, + MTU: iface.DefaultMTU, }, EngineServices{ SignalClient: &signal.MockClient{}, MgmClient: &mgmt.MockClient{SyncFunc: syncFunc}, @@ -412,11 +415,12 @@ func createEngine(ctx context.Context, cancel context.CancelFunc, setupKey strin wgPort := 33100 + i conf := &EngineConfig{ - WgIfaceName: ifaceName, - WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address), - WgPrivateKey: key, - WgPort: wgPort, - MTU: iface.DefaultMTU, + WgIfaceName: ifaceName, + WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address), + WgPrivateKey: key, + WgPort: wgPort, + IFaceBlackList: profilemanager.DefaultInterfaceBlacklist, + MTU: iface.DefaultMTU, } relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU) diff --git a/client/internal/engine_stdnet.go b/client/internal/engine_stdnet.go index 1ebb5779c..86f6d297a 100644 --- a/client/internal/engine_stdnet.go +++ b/client/internal/engine_stdnet.go @@ -6,6 +6,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func (e *Engine) newStdNet() (*stdnet.Net, error) { +func (e *Engine) newStdNet() *stdnet.Net { return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList) } diff --git a/client/internal/engine_stdnet_android.go b/client/internal/engine_stdnet_android.go index de3c80bcf..b14deeadf 100644 --- a/client/internal/engine_stdnet_android.go +++ b/client/internal/engine_stdnet_android.go @@ -2,6 +2,6 @@ package internal import "github.com/netbirdio/netbird/client/internal/stdnet" -func (e *Engine) newStdNet() (*stdnet.Net, error) { +func (e *Engine) newStdNet() *stdnet.Net { return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList) } diff --git a/client/internal/engine_test.go b/client/internal/engine_test.go index 2a7ecd652..3856cae22 100644 --- a/client/internal/engine_test.go +++ b/client/internal/engine_test.go @@ -161,7 +161,6 @@ func (m *MockWGIface) GetProxy() wgproxy.Proxy { return m.GetProxyFunc() } - func (m *MockWGIface) GetNet() *netstack.Net { return m.GetNetFunc() } @@ -689,10 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) { StatusRecorder: peer.NewRecorder("https://mgm"), }, MobileDependency{}) engine.ctx = ctx - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) opts := iface.WGIFaceOpts{ IFaceName: wgIfaceName, @@ -897,10 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) { }, MobileDependency{}) engine.ctx = ctx - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) opts := iface.WGIFaceOpts{ IFaceName: wgIfaceName, Address: wgaddr.MustParseWGAddress(wgAddr), diff --git a/client/internal/peer/ice/agent.go b/client/internal/peer/ice/agent.go index c74b46d10..6cd8c48de 100644 --- a/client/internal/peer/ice/agent.go +++ b/client/internal/peer/ice/agent.go @@ -39,10 +39,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c iceFailedTimeout := iceFailedTimeout() iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait() - transportNet, err := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList) - if err != nil { - log.Errorf("failed to create pion's stdnet: %s", err) - } + transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList) fac := logging.NewDefaultLoggerFactory() diff --git a/client/internal/peer/ice/stdnet.go b/client/internal/peer/ice/stdnet.go index 685ed0363..0c819ff66 100644 --- a/client/internal/peer/ice/stdnet.go +++ b/client/internal/peer/ice/stdnet.go @@ -8,6 +8,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) { +func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net { return stdnet.NewNet(ctx, ifaceBlacklist) } diff --git a/client/internal/peer/ice/stdnet_android.go b/client/internal/peer/ice/stdnet_android.go index 5033ec1b9..2962ecf66 100644 --- a/client/internal/peer/ice/stdnet_android.go +++ b/client/internal/peer/ice/stdnet_android.go @@ -6,6 +6,6 @@ import ( "github.com/netbirdio/netbird/client/internal/stdnet" ) -func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) { +func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net { return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist) } diff --git a/client/internal/relay/relay.go b/client/internal/relay/relay.go index 051717608..f0c65301e 100644 --- a/client/internal/relay/relay.go +++ b/client/internal/relay/relay.go @@ -201,11 +201,7 @@ func (p *StunTurnProbe) probeSTUN(ctx context.Context, uri *stun.URI) (addr stri } }() - net, err := stdnet.NewNet(ctx, nil) - if err != nil { - probeErr = fmt.Errorf("new net: %w", err) - return - } + net := stdnet.NewNet(ctx, nil) client, err := stun.DialURI(uri, &stun.DialConfig{ Net: net, @@ -290,11 +286,7 @@ func (p *StunTurnProbe) probeTURN(ctx context.Context, uri *stun.URI) (addr stri } }() - net, err := stdnet.NewNet(ctx, nil) - if err != nil { - probeErr = fmt.Errorf("new net: %w", err) - return - } + net := stdnet.NewNet(ctx, nil) cfg := &turn.ClientConfig{ STUNServerAddr: turnServerAddr, TURNServerAddr: turnServerAddr, diff --git a/client/internal/routemanager/manager_test.go b/client/internal/routemanager/manager_test.go index 18b44820a..a1624cf46 100644 --- a/client/internal/routemanager/manager_test.go +++ b/client/internal/routemanager/manager_test.go @@ -8,6 +8,7 @@ import ( "net/netip" "testing" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/stdnet" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" @@ -406,10 +407,7 @@ func TestManagerUpdateRoutes(t *testing.T) { for n, testCase := range testCases { t.Run(testCase.name, func(t *testing.T) { peerPrivateKey, _ := wgtypes.GeneratePrivateKey() - newNet, err := stdnet.NewNet(context.Background(), nil) - if err != nil { - t.Fatal(err) - } + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) opts := iface.WGIFaceOpts{ IFaceName: fmt.Sprintf("utun43%d", n), Address: wgaddr.MustParseWGAddress("100.65.65.2/24"), diff --git a/client/internal/routemanager/systemops/systemops_generic_test.go b/client/internal/routemanager/systemops/systemops_generic_test.go index c4f739c30..5b569ebd6 100644 --- a/client/internal/routemanager/systemops/systemops_generic_test.go +++ b/client/internal/routemanager/systemops/systemops_generic_test.go @@ -15,6 +15,7 @@ import ( "syscall" "testing" + "github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/stdnet" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -436,8 +437,7 @@ func createWGInterface(t *testing.T, interfaceName, ipAddressCIDR string, listen peerPrivateKey, err := wgtypes.GeneratePrivateKey() require.NoError(t, err) - newNet, err := stdnet.NewNet(context.Background(), nil) - require.NoError(t, err) + newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist) opts := iface.WGIFaceOpts{ IFaceName: interfaceName, diff --git a/client/internal/stdnet/stdnet.go b/client/internal/stdnet/stdnet.go index 381886ac6..c3a9d3d97 100644 --- a/client/internal/stdnet/stdnet.go +++ b/client/internal/stdnet/stdnet.go @@ -45,7 +45,7 @@ type Net struct { } // NewNetWithDiscover creates a new StdNet instance. -func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) (*Net, error) { +func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) *Net { if ctx == nil { ctx = context.Background() } @@ -60,20 +60,19 @@ func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover } else { n.iFaceDiscover = newMobileIFaceDiscover(iFaceDiscover) } - return n, n.UpdateInterfaces() + return n } // NewNet creates a new StdNet instance. -func NewNet(ctx context.Context, disallowList []string) (*Net, error) { +func NewNet(ctx context.Context, disallowList []string) *Net { if ctx == nil { ctx = context.Background() } - n := &Net{ + return &Net{ iFaceDiscover: pionDiscover{}, interfaceFilter: InterfaceFilter(disallowList), ctx: ctx, } - return n, n.UpdateInterfaces() } // resolveAddr performs DNS resolution with context support and timeout. @@ -122,45 +121,18 @@ func (n *Net) resolveAddr(network, address string) (netip.AddrPort, error) { return netip.AddrPortFrom(addrs[0], uint16(port)), nil } -// UpdateInterfaces updates the internal list of network interfaces -// and associated addresses filtering them by name. -// The interfaces are discovered by an external iFaceDiscover function or by a default discoverer if the external one -// wasn't specified. -func (n *Net) UpdateInterfaces() (err error) { - n.mu.Lock() - defer n.mu.Unlock() - - return n.updateInterfaces() -} - -func (n *Net) updateInterfaces() (err error) { - allIfaces, err := n.iFaceDiscover.iFaces() - if err != nil { - return err - } - - n.interfaces = n.filterInterfaces(allIfaces) - - n.lastUpdate = time.Now() - - return nil -} - // Interfaces returns a slice of interfaces which are available on the // system func (n *Net) Interfaces() ([]*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - if time.Since(n.lastUpdate) < updateInterval { - return slices.Clone(n.interfaces), nil + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err } - if err := n.updateInterfaces(); err != nil { - return nil, fmt.Errorf("update interfaces: %w", err) - } - - return slices.Clone(n.interfaces), nil + return slices.Clone(iFaces), nil } // InterfaceByIndex returns the interface specified by index. @@ -171,7 +143,13 @@ func (n *Net) Interfaces() ([]*transport.Interface, error) { func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - for _, ifc := range n.interfaces { + + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err + } + + for _, ifc := range iFaces { if ifc.Index == index { return ifc, nil } @@ -184,7 +162,13 @@ func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) { func (n *Net) InterfaceByName(name string) (*transport.Interface, error) { n.mu.Lock() defer n.mu.Unlock() - for _, ifc := range n.interfaces { + + iFaces, err := n.freshInterfacesLocked() + if err != nil { + return nil, err + } + + for _, ifc := range iFaces { if ifc.Name == name { return ifc, nil } @@ -193,6 +177,31 @@ func (n *Net) InterfaceByName(name string) (*transport.Interface, error) { return nil, fmt.Errorf("%w: %s", transport.ErrInterfaceNotFound, name) } +func (n *Net) freshInterfacesLocked() ([]*transport.Interface, error) { + if time.Since(n.lastUpdate) < updateInterval { + return n.interfaces, nil + } + + if err := n.updateInterfacesLocked(); err != nil { + return nil, fmt.Errorf("update interfaces: %w", err) + } + + return n.interfaces, nil +} + +func (n *Net) updateInterfacesLocked() error { + allIFaces, err := n.iFaceDiscover.iFaces() + if err != nil { + return err + } + + n.interfaces = n.filterInterfaces(allIFaces) + + n.lastUpdate = time.Now() + + return nil +} + func (n *Net) filterInterfaces(interfaces []*transport.Interface) []*transport.Interface { if n.interfaceFilter == nil { return interfaces diff --git a/client/internal/stdnet/stdnet_test.go b/client/internal/stdnet/stdnet_test.go new file mode 100644 index 000000000..822972f39 --- /dev/null +++ b/client/internal/stdnet/stdnet_test.go @@ -0,0 +1,136 @@ +package stdnet + +import ( + "context" + "errors" + "net" + "testing" + + "github.com/pion/transport/v3" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type countingDiscover struct { + calls int + list []*transport.Interface + err error +} + +func (d *countingDiscover) iFaces() ([]*transport.Interface, error) { + d.calls++ + if d.err != nil { + return nil, d.err + } + return d.list, nil +} + +func newTestNet(t *testing.T, d iFaceDiscover) *Net { + t.Helper() + return &Net{ + iFaceDiscover: d, + ctx: context.Background(), + } +} + +func testIFace(index int, name string) *transport.Interface { + return transport.NewInterface(net.Interface{Index: index, Name: name}) +} + +func TestNet_InterfacesDiscoversLazilyAndCaches(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}} + n := newTestNet(t, d) + + require.Zero(t, d.calls, "construction must not discover interfaces") + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, 1, d.calls) + + _, err = n.Interfaces() + require.NoError(t, err) + assert.Equal(t, 1, d.calls) +} + +func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) { + n := NewNet(context.Background(), nil) + require.NotNil(t, n) + assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") +} + +func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) { + n := NewNetWithDiscover(context.Background(), nil, nil) + require.NotNil(t, n) + assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold") +} + +func TestNet_InterfacesRetryAfterDiscoveryFailure(t *testing.T) { + discoverErr := errors.New("discover failed") + d := &countingDiscover{err: discoverErr} + n := newTestNet(t, d) + + _, err := n.Interfaces() + require.ErrorIs(t, err, discoverErr) + + d.err = nil + d.list = []*transport.Interface{testIFace(1, "eth0")} + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, 2, d.calls) +} + +func TestNet_InterfaceByNameRefreshes(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}} + n := newTestNet(t, d) + + ifc, err := n.InterfaceByName("eth0") + require.NoError(t, err) + assert.Equal(t, "eth0", ifc.Name) + assert.Equal(t, 1, d.calls) + + _, err = n.InterfaceByName("nope") + require.ErrorIs(t, err, transport.ErrInterfaceNotFound) +} + +func TestNet_InterfaceByIndexRefreshes(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}} + n := newTestNet(t, d) + + ifc, err := n.InterfaceByIndex(3) + require.NoError(t, err) + assert.Equal(t, "eth0", ifc.Name) + assert.Equal(t, 1, d.calls) + + _, err = n.InterfaceByIndex(99) + require.ErrorIs(t, err, transport.ErrInterfaceNotFound) +} + +func TestNet_InterfaceLookupPropagatesDiscoveryError(t *testing.T) { + discoverErr := errors.New("discover failed") + n := newTestNet(t, &countingDiscover{err: discoverErr}) + + _, err := n.InterfaceByName("eth0") + require.ErrorIs(t, err, discoverErr) + + _, err = n.InterfaceByIndex(1) + require.ErrorIs(t, err, discoverErr) +} + +func TestNet_InterfacesReturnsCopy(t *testing.T) { + d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}} + n := newTestNet(t, d) + + iFaces, err := n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + + iFaces[0] = testIFace(2, "tampered") + + iFaces, err = n.Interfaces() + require.NoError(t, err) + require.Len(t, iFaces, 1) + assert.Equal(t, "eth0", iFaces[0].Name) +} From 3c4358dd3625e458fef379dd15bd356b31147bfd Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Fri, 2 Oct 2026 19:03:00 +0200 Subject: [PATCH 06/14] [management] Require a private proxy cluster for cluster and direct upstream targets (#7984) * [management] Require a private proxy cluster for cluster and direct upstream targets Cluster targets and direct upstream targets make the proxy dial the upstream from its own host network instead of through the embedded NetBird client. Only clusters running in private mode are meant to do that, but the service API accepted these targets on any cluster. Service create and update now reject such targets unless the service's proxy cluster reports the private capability. An unreported capability is treated as unsupported. * [management] Require every proxy in the cluster to be private The private capability is aggregated as any-true, so a cluster where only one proxy runs in private mode passed the check. The mapping is delivered to every proxy in the cluster, so the non-private ones would serve cluster and direct upstream targets from their host network too. Validate these targets against a unanimous aggregation instead. The existing any-true lookup stays as is for the dashboard flags and the agent network gateway. --- .../modules/reverseproxy/proxy/manager.go | 1 + .../reverseproxy/proxy/manager/manager.go | 6 + .../proxy/manager/manager_test.go | 3 + .../reverseproxy/proxy/manager_mock.go | 14 ++ .../reverseproxy/service/manager/manager.go | 50 ++++ .../service/manager/private_cluster_test.go | 218 ++++++++++++++++++ management/server/store/sql_store_proxy.go | 8 + management/server/store/store.go | 1 + management/server/store/store_mock.go | 14 ++ proxy/management_integration_test.go | 4 + 10 files changed, 319 insertions(+) create mode 100644 management/internals/modules/reverseproxy/service/manager/private_cluster_test.go diff --git a/management/internals/modules/reverseproxy/proxy/manager.go b/management/internals/modules/reverseproxy/proxy/manager.go index c0b8435ec..9350ad9b9 100644 --- a/management/internals/modules/reverseproxy/proxy/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager.go @@ -20,6 +20,7 @@ type Manager interface { ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool CleanupStale(ctx context.Context, inactivityDuration time.Duration) error GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error) diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager.go b/management/internals/modules/reverseproxy/proxy/manager/manager.go index 7ddb66eec..5a95ea94a 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager.go @@ -23,6 +23,7 @@ type store interface { GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error) @@ -149,6 +150,11 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string) return m.store.GetClusterSupportsPrivate(ctx, clusterAddr) } +// ClusterAllProxiesPrivate reports whether every active proxy claims the private capability (nil = unreported). +func (m Manager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + return m.store.GetClusterAllProxiesPrivate(ctx, clusterAddr) +} + // ClusterSupportsSessionCode reports whether all active proxies support session codes. func (m Manager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool { versions, err := m.store.GetActiveProxyVersions(ctx, clusterAddr) diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go index 66ddb95bd..56806613a 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go @@ -105,6 +105,9 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool { return nil } +func (m *mockStore) GetClusterAllProxiesPrivate(_ context.Context, _ string) *bool { + return nil +} func (m *mockStore) GetActiveProxyVersions(ctx context.Context, clusterAddress string) ([]string, error) { if m.getActiveProxyVersionsFunc != nil { return m.getActiveProxyVersionsFunc(ctx, clusterAddress) diff --git a/management/internals/modules/reverseproxy/proxy/manager_mock.go b/management/internals/modules/reverseproxy/proxy/manager_mock.go index d6f7197d7..5f3404096 100644 --- a/management/internals/modules/reverseproxy/proxy/manager_mock.go +++ b/management/internals/modules/reverseproxy/proxy/manager_mock.go @@ -56,6 +56,20 @@ func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *go return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration) } +// ClusterAllProxiesPrivate mocks base method. +func (m *MockManager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ClusterAllProxiesPrivate", ctx, clusterAddr) + ret0, _ := ret[0].(*bool) + return ret0 +} + +// ClusterAllProxiesPrivate indicates an expected call of ClusterAllProxiesPrivate. +func (mr *MockManagerMockRecorder) ClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterAllProxiesPrivate", reflect.TypeOf((*MockManager)(nil).ClusterAllProxiesPrivate), ctx, clusterAddr) +} + // ClusterRequireSubdomain mocks base method. func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool { m.ctrl.T.Helper() diff --git a/management/internals/modules/reverseproxy/service/manager/manager.go b/management/internals/modules/reverseproxy/service/manager/manager.go index 62897c9ae..900b7759f 100644 --- a/management/internals/modules/reverseproxy/service/manager/manager.go +++ b/management/internals/modules/reverseproxy/service/manager/manager.go @@ -84,6 +84,7 @@ type CapabilityProvider interface { ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool } type Manager struct { @@ -332,6 +333,10 @@ func (m *Manager) persistNewService(ctx context.Context, accountID string, svc * return err } + if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil { + return err + } + return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil { return err @@ -369,6 +374,43 @@ func (m *Manager) clusterCustomPorts(ctx context.Context, svc *service.Service) return m.capabilities.ClusterSupportsCustomPorts(ctx, svc.ProxyCluster) } +// validatePrivateClusterTargets rejects cluster and direct upstream targets unless +// every active proxy in the service's cluster reports the private capability. The +// mapping reaches all proxies in the cluster, so one non-private proxy would serve +// these targets too. An unreported capability is treated as unsupported. Must be +// called outside a transaction, like clusterCustomPorts. +func (m *Manager) validatePrivateClusterTargets(ctx context.Context, targets []*service.Target, cluster string) error { + target := firstPrivateClusterTarget(targets) + if target == nil { + return nil + } + + if private := m.capabilities.ClusterAllProxiesPrivate(ctx, cluster); private != nil && *private { + return nil + } + + if target.TargetType == service.TargetTypeCluster { + return status.Errorf(status.InvalidArgument, + "target_type %q requires a proxy cluster with private mode enabled, cluster %s does not support it", + service.TargetTypeCluster, cluster) + } + return status.Errorf(status.InvalidArgument, + "direct_upstream requires a proxy cluster with private mode enabled, cluster %s does not support it", cluster) +} + +// firstPrivateClusterTarget returns the first target that only a private cluster may serve. +func firstPrivateClusterTarget(targets []*service.Target) *service.Target { + for _, target := range targets { + if target == nil { + continue + } + if target.TargetType == service.TargetTypeCluster || target.Options.DirectUpstream { + return target + } + } + return nil +} + // ensureL4Port auto-assigns a listen port when needed and validates cluster support. // customPorts must be pre-computed via clusterCustomPorts before entering a transaction. func (m *Manager) ensureL4Port(ctx context.Context, tx store.Store, svc *service.Service, customPorts *bool, serviceUpdate bool) error { @@ -464,6 +506,10 @@ func (m *Manager) persistNewEphemeralService(ctx context.Context, accountID, pee return err } + if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil { + return err + } + return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil { return err @@ -584,6 +630,10 @@ func (m *Manager) persistServiceUpdate(ctx context.Context, accountID string, se return nil, err } + if err := m.validatePrivateClusterTargets(ctx, service.Targets, effectiveCluster); err != nil { + return nil, err + } + // Validate subdomain requirement *before* the transaction: the underlying // capability lookup talks to the main DB pool, and SQLite's single-connection // pool would self-deadlock if this ran while the tx already held the only diff --git a/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go b/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go new file mode 100644 index 000000000..1f507294e --- /dev/null +++ b/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go @@ -0,0 +1,218 @@ +package manager + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/metric/noop" + "go.uber.org/mock/gomock" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" + 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/store" + "github.com/netbirdio/netbird/shared/management/status" +) + +// setupPrivateClusterTest wires the real proxy manager as the capability +// provider and connects one proxy to testCluster reporting the given private +// capability. A nil private connects no proxy, so the capability is unreported. +func setupPrivateClusterTest(t *testing.T, private *bool) (*Manager, store.Store) { + t.Helper() + + mgr, testStore := setupIntegrationTest(t) + + proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter("")) + require.NoError(t, err) + mgr.capabilities = proxyMgr + + if private != nil { + connectTestProxy(t, proxyMgr, "proxy-1", &proxy.Capabilities{Private: private}) + } + + return mgr, testStore +} + +func connectTestProxy(t *testing.T, proxyMgr *proxymanager.Manager, proxyID string, caps *proxy.Capabilities) { + t.Helper() + _, err := proxyMgr.Connect(context.Background(), proxyID, "session-"+proxyID, testCluster, "127.0.0.1", "", nil, caps) + require.NoError(t, err) +} + +func clusterTarget() *rpservice.Target { + return &rpservice.Target{ + TargetId: testCluster, + TargetType: rpservice.TargetTypeCluster, + Host: "backend.lan", + Port: 8080, + Protocol: "http", + Enabled: true, + Options: rpservice.TargetOptions{DirectUpstream: true}, + } +} + +func directUpstreamPeerTarget() *rpservice.Target { + return &rpservice.Target{ + TargetId: testPeerID, + TargetType: rpservice.TargetTypePeer, + Host: "backend.lan", + Port: 8080, + Protocol: "http", + Enabled: true, + Options: rpservice.TargetOptions{DirectUpstream: true}, + } +} + +func TestCreateService_PrivateClusterTargets(t *testing.T) { + tests := []struct { + name string + private *bool + target *rpservice.Target + wantErr string + }{ + {name: "cluster target on private cluster", private: boolPtr(true), target: clusterTarget()}, + {name: "direct upstream on private cluster", private: boolPtr(true), target: directUpstreamPeerTarget()}, + {name: "cluster target on non-private cluster", private: boolPtr(false), target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "direct upstream on non-private cluster", private: boolPtr(false), target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + {name: "cluster target with unreported capability", private: nil, target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "direct upstream with unreported capability", private: nil, target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, tc.private) + + svc := newTestService("app.test.netbird.io") + svc.Targets = []*rpservice.Target{tc.target} + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, svc) + + services, listErr := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, listErr) + + if tc.wantErr == "" { + require.NoError(t, err) + assert.Len(t, services, 1, "the service should be persisted") + return + } + + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantErr) + sErr, ok := status.FromError(err) + require.True(t, ok, "the caller must receive a typed error") + assert.Equal(t, status.InvalidArgument, sErr.Type(), "the rejection should be an invalid argument") + assert.Empty(t, services, "a rejected service must not be persisted") + }) + } +} + +// A cluster where only some proxies run in private mode must not accept these +// targets: the mapping is delivered to every proxy in the cluster, so the +// non-private ones would serve the target from their host network as well. +func TestCreateService_MixedClusterRejectsPrivateTargets(t *testing.T) { + tests := []struct { + name string + secondCaps *proxy.Capabilities + }{ + {name: "second proxy reports not private", secondCaps: &proxy.Capabilities{Private: boolPtr(false)}}, + {name: "second proxy predates capability reporting", secondCaps: nil}, + } + + for _, tc := range tests { + for _, target := range []*rpservice.Target{clusterTarget(), directUpstreamPeerTarget()} { + t.Run(tc.name+"/"+string(target.TargetType), func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(true)) + connectTestProxy(t, mgr.capabilities.(*proxymanager.Manager), "proxy-2", tc.secondCaps) + + svc := newTestService("app.test.netbird.io") + svc.Targets = []*rpservice.Target{target} + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, svc) + require.Error(t, err, "a cluster with a non-private proxy must not accept the target") + assert.Contains(t, err.Error(), "requires a proxy cluster with private mode enabled") + + services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID) + require.NoError(t, err) + assert.Empty(t, services, "a rejected service must not be persisted") + }) + } + } +} + +func TestCreateService_RegularTargetIgnoresPrivateCapability(t *testing.T) { + ctx := context.Background() + mgr, _ := setupPrivateClusterTest(t, boolPtr(false)) + + _, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err, "a peer target without direct upstream must not need a private cluster") +} + +func TestUpdateService_PrivateClusterTargets(t *testing.T) { + tests := []struct { + name string + target *rpservice.Target + wantErr string + }{ + {name: "switch to cluster target", target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`}, + {name: "enable direct upstream", target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(false)) + + created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err) + + updated := newTestService("app.test.netbird.io") + updated.ID = created.ID + updated.AccountID = testAccountID + updated.Targets = []*rpservice.Target{tc.target} + + _, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantErr) + + stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID) + require.NoError(t, err) + require.Len(t, stored.Targets, 1) + assert.Equal(t, rpservice.TargetTypePeer, stored.Targets[0].TargetType, "the stored target must be unchanged") + assert.False(t, stored.Targets[0].Options.DirectUpstream, "the stored target must keep direct upstream disabled") + }) + } +} + +func TestUpdateService_PrivateClusterAllowsClusterTarget(t *testing.T) { + ctx := context.Background() + mgr, testStore := setupPrivateClusterTest(t, boolPtr(true)) + + created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io")) + require.NoError(t, err) + + updated := newTestService("app.test.netbird.io") + updated.ID = created.ID + updated.AccountID = testAccountID + updated.Targets = []*rpservice.Target{clusterTarget()} + + _, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated) + require.NoError(t, err) + + stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID) + require.NoError(t, err) + require.Len(t, stored.Targets, 1) + assert.Equal(t, rpservice.TargetTypeCluster, stored.Targets[0].TargetType, "the cluster target should be stored") +} + +func TestValidatePrivateClusterTargets_NoLookupWithoutPrivateTargets(t *testing.T) { + ctrl := gomock.NewController(t) + // No ClusterAllProxiesPrivate expectation: a lookup would fail the test. + mgr := &Manager{capabilities: proxy.NewMockManager(ctrl)} + + targets := []*rpservice.Target{{TargetId: testPeerID, TargetType: rpservice.TargetTypePeer}} + require.NoError(t, mgr.validatePrivateClusterTargets(context.Background(), targets, testCluster)) +} diff --git a/management/server/store/sql_store_proxy.go b/management/server/store/sql_store_proxy.go index 58fa86468..bdccd282c 100644 --- a/management/server/store/sql_store_proxy.go +++ b/management/server/store/sql_store_proxy.go @@ -358,6 +358,14 @@ func (s *SqlStore) GetClusterSupportsPrivate(ctx context.Context, clusterAddr st return s.getClusterCapability(ctx, clusterAddr, "private") } +// GetClusterAllProxiesPrivate reports whether every active proxy in the cluster +// has the private capability. Returns nil when no proxy reported the capability. +// Use it where any proxy in the cluster may serve the result, since a single +// non-private proxy would serve it without the private guarantees. +func (s *SqlStore) GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + return s.getClusterUnanimousCapability(ctx, clusterAddr, "private") +} + // GetClusterSupportsCrowdSec returns whether all active proxies in the cluster // have CrowdSec configured. Returns nil when no proxy reported the capability. // Unlike other capabilities that use ANY-true (for rolling upgrades), CrowdSec diff --git a/management/server/store/store.go b/management/server/store/store.go index 177c2a47c..01aaf4892 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -336,6 +336,7 @@ type Store interface { GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool + GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error) diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 53e35b866..956cac4b8 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -1870,6 +1870,20 @@ func (mr *MockStoreMockRecorder) GetAnyAccountID(ctx any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAnyAccountID", reflect.TypeOf((*MockStore)(nil).GetAnyAccountID), ctx) } +// GetClusterAllProxiesPrivate mocks base method. +func (m *MockStore) GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetClusterAllProxiesPrivate", ctx, clusterAddr) + ret0, _ := ret[0].(*bool) + return ret0 +} + +// GetClusterAllProxiesPrivate indicates an expected call of GetClusterAllProxiesPrivate. +func (mr *MockStoreMockRecorder) GetClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusterAllProxiesPrivate", reflect.TypeOf((*MockStore)(nil).GetClusterAllProxiesPrivate), ctx, clusterAddr) +} + // GetClusterRequireSubdomain mocks base method. func (m *MockStore) GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool { m.ctrl.T.Helper() diff --git a/proxy/management_integration_test.go b/proxy/management_integration_test.go index 000d8ce72..03a9855de 100644 --- a/proxy/management_integration_test.go +++ b/proxy/management_integration_test.go @@ -246,6 +246,10 @@ func (m *testProxyManager) ClusterSupportsPrivate(_ context.Context, _ string) * return nil } +func (m *testProxyManager) ClusterAllProxiesPrivate(_ context.Context, _ string) *bool { + return nil +} + func (m *testProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool { return m.supportsSessionCode } From 88b26bb74f50f12f25ad7b5db8eddb171d91a1e5 Mon Sep 17 00:00:00 2001 From: Nicolas Frati Date: Fri, 2 Oct 2026 22:53:36 +0200 Subject: [PATCH 07/14] [infrastructure] Create the preflight artifacts directory before submitting (#7947) * [infrastructure] Create the preflight artifacts directory before submitting * [infrastructure] Add a workflow to certify UBI images on demand Move the Red Hat certification job into redhat-certify.yml so it can be run by hand for any released version and component, or for all of them. release.yml calls it with component "all" on stable tags. Component IDs now come from REDHAT_CERT_ID_ repository variables. With "all", components without a variable are skipped. * [infrastructure] Fail Red Hat certification on missing IDs or timeout Fail before certifying when a selected component's REDHAT_CERT_ID_* variable is missing, listing every missing variable. Filter the Pyxis poll by tag so older versions are found past the first page, and fail the job when both architectures are not certified within 10 minutes. --- .github/workflows/redhat-certify.yml | 199 +++++++++++++++++++++++++++ .github/workflows/release.yml | 126 ++--------------- 2 files changed, 208 insertions(+), 117 deletions(-) create mode 100644 .github/workflows/redhat-certify.yml diff --git a/.github/workflows/redhat-certify.yml b/.github/workflows/redhat-certify.yml new file mode 100644 index 000000000..e592dabc2 --- /dev/null +++ b/.github/workflows/redhat-certify.yml @@ -0,0 +1,199 @@ +name: Red Hat Certification + +# Certify published UBI images in the Red Hat Ecosystem Catalog. Called by +# release.yml on stable tags, or run by hand to (re)certify any released +# version. preflight submits every architecture of an image's manifest list +# to Pyxis; auto-publish on the component makes it public once certified. +# +# Each component's Partner Connect ID comes from the REDHAT_CERT_ID_ +# repository variable, e.g. REDHAT_CERT_ID_CLIENT_ROOTLESS. The run fails +# before certifying anything if a selected component's variable is not set. + +on: + workflow_call: + inputs: + component: + type: string + required: true + version: + type: string + required: true + secrets: + PYXIS_API_TOKEN: + required: true + workflow_dispatch: + inputs: + component: + description: "Component to certify" + type: choice + required: true + default: all + options: + - all + - client-rootless + - reverse-proxy + version: + description: "Released version, e.g. v0.80.0" + type: string + required: true + +permissions: + contents: read + +jobs: + resolve: + name: Resolve components + runs-on: ubuntu-24.04 + outputs: + version: ${{ steps.resolve.outputs.version }} + matrix: ${{ steps.resolve.outputs.matrix }} + steps: + - name: Resolve components and images + id: resolve + env: + COMPONENT: ${{ inputs.component }} + INPUT_VERSION: ${{ inputs.version }} + REPO_VARS: ${{ toJSON(vars) }} + run: | + set -euo pipefail + version="${INPUT_VERSION#v}" + if [[ ! "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then + echo "::error::Only stable x.y.z versions are certified, got '${INPUT_VERSION}'" + exit 1 + fi + # name, image repository, tag suffix (must match .goreleaser.yaml). + # Keep the names in sync with the workflow_dispatch options above. + components=( + "client-rootless ghcr.io/netbirdio/netbird -rootless-ubi" + "reverse-proxy ghcr.io/netbirdio/reverse-proxy -ubi" + ) + matrix="[]" + missing=() + for c in "${components[@]}"; do + read -r name repo suffix <<< "$c" + [[ "$COMPONENT" == all || "$COMPONENT" == "$name" ]] || continue + var="REDHAT_CERT_ID_${name^^}"; var="${var//-/_}" + id="$(jq -r --arg v "$var" '.[$v] // empty' <<< "$REPO_VARS")" + if [[ -z "$id" ]]; then + missing+=("$var") + continue + fi + matrix="$(jq -c --arg n "$name" --arg t "${version}${suffix}" --arg r "${repo}:${version}${suffix}" --arg i "$id" \ + '. + [{component: $n, tag: $t, ref: $r, component_id: $i}]' <<< "$matrix")" + done + if (( ${#missing[@]} )); then + echo "::error::Set these repository variables to the Partner Connect component IDs: ${missing[*]}" + exit 1 + fi + if [[ "$matrix" == "[]" ]]; then + echo "::error::No component to certify for '${COMPONENT}'" + exit 1 + fi + echo "Components to certify: ${matrix}" + echo "version=${version}" >> "$GITHUB_OUTPUT" + echo "matrix=${matrix}" >> "$GITHUB_OUTPUT" + + certify: + name: "Certify ${{ matrix.component }} UBI image" + needs: resolve + runs-on: ubuntu-24.04 + strategy: + fail-fast: false + matrix: + include: ${{ fromJSON(needs.resolve.outputs.matrix) }} + env: + PREFLIGHT_VERSION: "1.21.0" + # sha256 of preflight-linux-amd64 from the 1.21.0 GitHub release. + # Red Hat publishes no checksum file, so the value is pinned here. + PREFLIGHT_SHA256: "5e653135503c72f8702bbe31d7643197d12937c68086879133dd6b9650a9a449" + steps: + - name: Verify the multi-arch image is on ghcr.io + env: + IMAGE_REF: ${{ matrix.ref }} + run: | + set -euo pipefail + docker buildx imagetools inspect "$IMAGE_REF" --raw > manifest.json + for arch in amd64 arm64; do + if ! jq -e --arg a "$arch" '.manifests[] | select(.platform.architecture == $a)' manifest.json > /dev/null; then + echo "::error::${IMAGE_REF} has no ${arch} manifest" + exit 1 + fi + done + echo "Manifest list for ${IMAGE_REF}:" + jq -r '.manifests[] | "\(.platform.os)/\(.platform.architecture) \(.digest)"' manifest.json + + - name: Install preflight + run: | + set -euo pipefail + curl -fsSL --proto '=https' --proto-redir '=https' -o preflight \ + "https://github.com/redhat-openshift-ecosystem/openshift-preflight/releases/download/${PREFLIGHT_VERSION}/preflight-linux-amd64" + echo "${PREFLIGHT_SHA256} preflight" | sha256sum -c - + chmod +x preflight + ./preflight --version + + - name: Run preflight checks and submit to Red Hat + env: + IMAGE_REF: ${{ matrix.ref }} + PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }} + PFLT_CERTIFICATION_COMPONENT_ID: ${{ matrix.component_id }} + PFLT_ARTIFACTS: artifacts + PFLT_LOGFILE: artifacts/preflight.log + PFLT_LOGLEVEL: info + PFLT_JUNIT: "true" + run: | + set -euo pipefail + # No --platform: preflight walks the manifest list and submits every + # architecture in one run, grouped under one manifest-list digest. + # preflight does not create the PFLT_LOGFILE directory, and --submit + # fails if the log file is missing. + mkdir -p artifacts + ./preflight check container "$IMAGE_REF" --submit + + - name: Fail if any check did not pass + run: | + set -euo pipefail + shopt -s nullglob + results=(artifacts/results.json artifacts/*/results.json) + if [[ ${#results[@]} -eq 0 ]]; then + echo "::error::preflight produced no results.json" + exit 1 + fi + status=0 + for f in "${results[@]}"; do + arch="$(basename "$(dirname "$f")")" + passed="$(jq -r '.passed' "$f")" + failed="$(jq -r '[.results.failed[]?.name] | join(", ")' "$f")" + echo "${arch}: passed=${passed} ${failed:+failed checks: ${failed}}" + [[ "$passed" == "true" ]] || status=1 + done + exit $status + + - name: Upload preflight artifacts + if: always() + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 + with: + name: redhat-preflight-${{ matrix.component }}-${{ needs.resolve.outputs.version }} + path: artifacts/ + retention-days: 30 + + - name: Wait for Pyxis to mark both architectures certified + env: + TAG: ${{ matrix.tag }} + COMPONENT_ID: ${{ matrix.component_id }} + PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }} + run: | + set -euo pipefail + # Filter on the tag server-side so older versions are found past the first page. + url="https://catalog.redhat.com/api/containers/v1/projects/certification/id/${COMPONENT_ID}/images?filter=repositories.tags.name==${TAG}&page_size=100" + for attempt in $(seq 1 20); do + certified="$(curl -fsS --proto '=https' --proto-redir '=https' -H "X-API-KEY: ${PFLT_PYXIS_API_TOKEN}" "$url" \ + | jq -r --arg t "$TAG" '[.data[] | select(.repositories[]?.tags[]?.name == $t) | select(.certified == true) | .architecture] | unique | join(",")')" + echo "attempt ${attempt}: certified architectures for ${TAG}: ${certified:-none}" + if [[ "$certified" == "amd64,arm64" ]]; then + echo "Both architectures certified. Auto-publish is enabled on the component, so the catalog updates on its own." + exit 0 + fi + sleep 30 + done + echo "::error::Pyxis has not marked both architectures certified after 10 minutes. Check https://connect.redhat.com/component/view/${COMPONENT_ID}/images" + exit 1 diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 9c9af5f17..dee4d398d 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -380,131 +380,23 @@ jobs: path: dist/netbird_darwin** retention-days: 7 - # Certify and publish the rootless UBI client image in the Red Hat Ecosystem - # Catalog. Stable tags only: goreleaser pushes -rootless-ubi to - # ghcr.io in the release job above, and preflight submits every architecture - # of that manifest list to Pyxis. Auto-publish on the component makes the new - # version public once certification passes. + # Certify the UBI images in the Red Hat Ecosystem Catalog on stable tags. + # See redhat-certify.yml, which can also be run by hand for any released version. redhat_certification: - name: "Red Hat / Certify rootless UBI image" + name: "Red Hat" needs: release if: | github.repository == 'netbirdio/netbird' && startsWith(github.ref, 'refs/tags/v') && !contains(github.ref_name, '-') - runs-on: ubuntu-24.04 permissions: contents: read - env: - PREFLIGHT_VERSION: "1.21.0" - # sha256 of preflight-linux-amd64 from the 1.21.0 GitHub release. - # Red Hat publishes no checksum file, so the value is pinned here. - PREFLIGHT_SHA256: "5e653135503c72f8702bbe31d7643197d12937c68086879133dd6b9650a9a449" - IMAGE_REPOSITORY: "ghcr.io/netbirdio/netbird" - # Component "NetBird Client Container Image (rootless)" in Partner Connect. - # Override with the REDHAT_CERT_COMPONENT_ID repository variable if it changes. - DEFAULT_COMPONENT_ID: "6aa3ca4b4676aefdf07aaa97" - steps: - - name: Resolve image reference - id: image - env: - INPUT_VERSION: ${{ github.ref_name }} - run: | - set -euo pipefail - version="${INPUT_VERSION#v}" - if [[ ! "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then - echo "::error::Only stable x.y.z versions are certified, got '${INPUT_VERSION}'" - exit 1 - fi - echo "version=${version}" >> "$GITHUB_OUTPUT" - echo "ref=${IMAGE_REPOSITORY}:${version}-rootless-ubi" >> "$GITHUB_OUTPUT" - - - name: Verify the multi-arch image is on ghcr.io - env: - IMAGE_REF: ${{ steps.image.outputs.ref }} - run: | - set -euo pipefail - docker buildx imagetools inspect "$IMAGE_REF" --raw > manifest.json - for arch in amd64 arm64; do - if ! jq -e --arg a "$arch" '.manifests[] | select(.platform.architecture == $a)' manifest.json > /dev/null; then - echo "::error::${IMAGE_REF} has no ${arch} manifest" - exit 1 - fi - done - echo "Manifest list for ${IMAGE_REF}:" - jq -r '.manifests[] | "\(.platform.os)/\(.platform.architecture) \(.digest)"' manifest.json - - - name: Install preflight - run: | - set -euo pipefail - curl -fsSL --proto '=https' --proto-redir '=https' -o preflight \ - "https://github.com/redhat-openshift-ecosystem/openshift-preflight/releases/download/${PREFLIGHT_VERSION}/preflight-linux-amd64" - echo "${PREFLIGHT_SHA256} preflight" | sha256sum -c - - chmod +x preflight - ./preflight --version - - - name: Run preflight checks and submit to Red Hat - env: - IMAGE_REF: ${{ steps.image.outputs.ref }} - PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }} - PFLT_CERTIFICATION_COMPONENT_ID: ${{ vars.REDHAT_CERT_COMPONENT_ID || env.DEFAULT_COMPONENT_ID }} - PFLT_ARTIFACTS: artifacts - PFLT_LOGFILE: artifacts/preflight.log - PFLT_LOGLEVEL: info - PFLT_JUNIT: "true" - run: | - set -euo pipefail - # No --platform: preflight walks the manifest list and submits every - # architecture in one run, grouped under one manifest-list digest. - ./preflight check container "$IMAGE_REF" --submit - - - name: Fail if any check did not pass - run: | - set -euo pipefail - shopt -s nullglob - results=(artifacts/results.json artifacts/*/results.json) - if [[ ${#results[@]} -eq 0 ]]; then - echo "::error::preflight produced no results.json" - exit 1 - fi - status=0 - for f in "${results[@]}"; do - arch="$(basename "$(dirname "$f")")" - passed="$(jq -r '.passed' "$f")" - failed="$(jq -r '[.results.failed[]?.name] | join(", ")' "$f")" - echo "${arch}: passed=${passed} ${failed:+failed checks: ${failed}}" - [[ "$passed" == "true" ]] || status=1 - done - exit $status - - - name: Upload preflight artifacts - if: always() - uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 - with: - name: redhat-preflight-${{ steps.image.outputs.version }} - path: artifacts/ - retention-days: 30 - - - name: Wait for Pyxis to mark both architectures certified - env: - VERSION: ${{ steps.image.outputs.version }} - PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }} - COMPONENT_ID: ${{ vars.REDHAT_CERT_COMPONENT_ID || env.DEFAULT_COMPONENT_ID }} - run: | - set -euo pipefail - tag="${VERSION}-rootless-ubi" - url="https://catalog.redhat.com/api/containers/v1/projects/certification/id/${COMPONENT_ID}/images?page_size=100" - for attempt in $(seq 1 20); do - certified="$(curl -fsS --proto '=https' --proto-redir '=https' -H "X-API-KEY: ${PFLT_PYXIS_API_TOKEN}" "$url" \ - | jq -r --arg t "$tag" '[.data[] | select(.repositories[]?.tags[]?.name == $t) | select(.certified == true) | .architecture] | unique | join(",")')" - echo "attempt ${attempt}: certified architectures for ${tag}: ${certified:-none}" - if [[ "$certified" == "amd64,arm64" ]]; then - echo "Both architectures certified. Auto-publish is enabled on the component, so the catalog updates on its own." - exit 0 - fi - sleep 30 - done - echo "::warning::Pyxis has not marked both architectures certified after 10 minutes. Check https://connect.redhat.com/component/view/${COMPONENT_ID}/images" + uses: ./.github/workflows/redhat-certify.yml + with: + component: all + version: ${{ github.ref_name }} + secrets: + PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }} release_ui: runs-on: ubuntu-latest From 7b8fa29031add85380adab1c504186eae4a90247 Mon Sep 17 00:00:00 2001 From: PizzaLovingNerd Date: Mon, 5 Oct 2026 02:23:47 -0700 Subject: [PATCH 08/14] [self-hosted] Replace "which" dependency by "command" from configure.sh script (#8007) --- infrastructure_files/configure.sh | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/infrastructure_files/configure.sh b/infrastructure_files/configure.sh index 92252d0b3..ce1a041e6 100755 --- a/infrastructure_files/configure.sh +++ b/infrastructure_files/configure.sh @@ -1,14 +1,14 @@ #!/bin/bash set -e -if ! which curl >/dev/null 2>&1; then +if ! command -v curl >/dev/null 2>&1; then echo "This script uses curl fetch OpenID configuration from IDP." echo "Please install curl and re-run the script https://curl.se/" echo "" exit 1 fi -if ! which jq >/dev/null 2>&1; then +if ! command -v jq >/dev/null 2>&1; then echo "This script uses jq to load OpenID configuration from IDP." echo "Please install jq and re-run the script https://stedolan.github.io/jq/" echo "" @@ -18,13 +18,13 @@ fi source setup.env source base.setup.env -if ! which envsubst >/dev/null 2>&1; then +if ! command -v envsubst >/dev/null 2>&1; then echo "envsubst is needed to run this script" if [[ $(uname) == "Darwin" ]]; then echo "you can install it with homebrew (https://brew.sh):" echo "brew install gettext" else - if which apt-get >/dev/null 2>&1; then + if command -v apt-get >/dev/null 2>&1; then echo "you can install it by running" echo "apt-get update && apt-get install gettext-base" else From 0cd27ca14b52ba01fd3e57bf3f18328d10047fc3 Mon Sep 17 00:00:00 2001 From: Redouan El Rhazouani <81578195+redouan-rhazouani@users.noreply.github.com> Date: Mon, 5 Oct 2026 12:31:22 +0200 Subject: [PATCH 09/14] [management] Improve Base62 encoding/decoding performance and robustness (#3391) --- base62/base62.go | 85 ++++++++++++++++++++++++++----------------- base62/base62_test.go | 64 +++++++++++++++++++++++++------- 2 files changed, 102 insertions(+), 47 deletions(-) diff --git a/base62/base62.go b/base62/base62.go index efafbc768..1a02e98e2 100644 --- a/base62/base62.go +++ b/base62/base62.go @@ -3,56 +3,75 @@ package base62 import ( "fmt" "math" - "strings" ) const ( - alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" - base = uint32(len(alphabet)) + alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + base = uint32(len(alphabet)) + maxBase62Digits = 6 // max number of digits required to encode MaxUint32 + ) +var ( + ErrEmptyString = fmt.Errorf("empty string") + ErrInvalidChar = fmt.Errorf("invalid character") + ErrOverflow = fmt.Errorf("integer overflow") +) + +// Fixed-size arrays have better performance and lower memory overhead compared to maps for small static sets of data +var charToIndex [123]int8 // Assuming ASCII from '\0' - 'z' + +func init() { + for i := range charToIndex { + charToIndex[i] = -1 + } + for i, c := range alphabet { + charToIndex[c] = int8(i) + } +} + // Encode encodes a uint32 value to a base62 string. -func Encode(num uint32) string { - if num == 0 { - return string(alphabet[0]) +// The returned string will be between 1-6 characters long. +func Encode(n uint32) string { + if n < base { + return string(alphabet[n]) + } + // avoid dynamic memory usage for small, fixed size data + buf := [maxBase62Digits]byte{} + idx := len(buf) + + for n > 0 { + idx-- + buf[idx] = alphabet[n%base] + n /= base } - var encoded strings.Builder - - for num > 0 { - remainder := num % base - encoded.WriteByte(alphabet[remainder]) - num /= base - } - - // Reverse the encoded string - encodedString := encoded.String() - reversed := reverse(encodedString) - return reversed + return string(buf[idx:]) } // Decode decodes a base62 string to a uint32 value. +// Returns an error if the input string is empty, contains invalid characters, +// or would result in integer overflow. func Decode(encoded string) (uint32, error) { + if len(encoded) == 0 { + return 0, ErrEmptyString + } var decoded uint32 - strLen := len(encoded) - - for i, char := range encoded { - index := strings.IndexRune(alphabet, char) + for _, char := range encoded { + index := int8(-1) + if int(char) < len(charToIndex) { + index = charToIndex[char] + } if index < 0 { - return 0, fmt.Errorf("invalid character: %c", char) + return 0, fmt.Errorf("%w: %c", ErrInvalidChar, char) + } + // Add overflow check when calculating the decoded value to prevent silent overflow of uint32 + if decoded > (math.MaxUint32-uint32(index))/base { + return 0, fmt.Errorf("%w: %s", ErrOverflow, encoded) } - decoded += uint32(index) * uint32(math.Pow(float64(base), float64(strLen-i-1))) + decoded = decoded*base + uint32(index) } return decoded, nil } - -// Reverse a string. -func reverse(s string) string { - runes := []rune(s) - for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 { - runes[i], runes[j] = runes[j], runes[i] - } - return string(runes) -} diff --git a/base62/base62_test.go b/base62/base62_test.go index 00da2124a..f2ad06d6f 100644 --- a/base62/base62_test.go +++ b/base62/base62_test.go @@ -1,31 +1,67 @@ package base62 import ( + "errors" + "math" "testing" ) func TestEncodeDecode(t *testing.T) { - tests := []struct { - num uint32 + testCases := []struct { + input uint32 + expected string }{ - {0}, - {1}, - {42}, - {12345}, - {99999}, - {123456789}, + {0, "0"}, + {1, "1"}, + {5, "5"}, + {9, "9"}, + {10, "A"}, + {42, "g"}, + {61, "z"}, + {62, "10"}, + {'0', "m"}, + {'9', "v"}, + {'A', "13"}, + {'Z', "1S"}, + {'a', "1Z"}, + {'z', "1y"}, + {99999, "Q0t"}, + {12345, "3D7"}, + {123456789, "8M0kX"}, + {math.MaxUint32, "4gfFC3"}, } - for _, tt := range tests { - encoded := Encode(tt.num) + for _, tc := range testCases { + encoded := Encode(tc.input) + if encoded != tc.expected { + t.Errorf("Encode(%d) = %s; want %s", tc.input, encoded, tc.expected) + } decoded, err := Decode(encoded) - if err != nil { - t.Errorf("Decode error: %v", err) + t.Errorf("Expected error nil, got %v", err) } - if decoded != tt.num { - t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tt.num) + if decoded != tc.input { + t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tc.input) } } } + +// Decode handles empty string input with appropriate error +func TestDecodeEmptyString(t *testing.T) { + if _, err := Decode(""); !errors.Is(err, ErrEmptyString) { + t.Errorf("Expected error %v, got %v", ErrEmptyString, err) + } +} + +func TestDecodeOverflow(t *testing.T) { + if _, err := Decode("4gfFC4"); !errors.Is(err, ErrOverflow) { + t.Errorf("Expected error %v, got %v", ErrOverflow, err) + } +} + +func TestDecodeInvalid(t *testing.T) { + if _, err := Decode("/"); !errors.Is(err, ErrInvalidChar) { + t.Errorf("Expected error %v, got %v", ErrInvalidChar, err) + } +} From 19c54b82264965bb20ca18886324fb41b55f7fe8 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Mon, 5 Oct 2026 12:48:53 +0200 Subject: [PATCH 10/14] [management] Refresh only affected peers on DNS zone and record changes (#8050) Zone and record changes refreshed every peer in the account, and zone create/update passed the request context to the update goroutine, so it could be cancelled when the handler returned. They now compute affected peers from the zone's distribution groups inside the transaction and dispatch through ExpandAndUpdateAffected, which detaches the context. The resolver did not know about zones, so changing a group referenced only by a zone never pushed the zone to its added or removed members. It now folds the distribution groups of shipped zones on whole-group changes. --- .../modules/zones/manager/manager.go | 84 +++++++--- .../modules/zones/records/manager/manager.go | 29 +++- management/server/affected_peers_zone_test.go | 145 ++++++++++++++++++ management/server/affectedpeers/resolver.go | 29 +++- 4 files changed, 258 insertions(+), 29 deletions(-) create mode 100644 management/server/affected_peers_zone_test.go diff --git a/management/internals/modules/zones/manager/manager.go b/management/internals/modules/zones/manager/manager.go index d5348d3d0..6f6ba6c40 100644 --- a/management/internals/modules/zones/manager/manager.go +++ b/management/internals/modules/zones/manager/manager.go @@ -3,15 +3,16 @@ package manager import ( "context" "fmt" + "slices" "github.com/netbirdio/netbird/management/internals/modules/zones" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/status" ) @@ -69,6 +70,9 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, } zone = zones.NewZone(accountID, zone.Name, zone.Domain, zone.Enabled, zone.EnableSearchDomain, zone.DistributionGroups) + var snap *affectedpeers.Snapshot + change := affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { existingZone, err := transaction.GetZoneByDomain(ctx, accountID, zone.Domain) if err != nil { @@ -88,7 +92,15 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, } if err = transaction.CreateZone(ctx, zone); err != nil { - return fmt.Errorf("failed to create zone: %w", err) + return fmt.Errorf("create zone: %w", err) + } + + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + + if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil { + return fmt.Errorf("increment network serial: %w", err) } return nil @@ -99,6 +111,8 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string, m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneCreated, zone.EventMeta()) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) + return zone, nil } @@ -111,21 +125,26 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, return nil, status.NewPermissionDeniedError() } - zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID) - if err != nil { - return nil, fmt.Errorf("failed to get zone: %w", err) - } - - if zone.Domain != updatedZone.Domain { - return nil, status.Errorf(status.InvalidArgument, "zone domain cannot be updated") - } - - zone.Name = updatedZone.Name - zone.Enabled = updatedZone.Enabled - zone.EnableSearchDomain = updatedZone.EnableSearchDomain - zone.DistributionGroups = updatedZone.DistributionGroups + var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID) + if err != nil { + return fmt.Errorf("get zone: %w", err) + } + + if zone.Domain != updatedZone.Domain { + return status.Errorf(status.InvalidArgument, "zone domain cannot be updated") + } + + oldGroups := zone.DistributionGroups + zone.Name = updatedZone.Name + zone.Enabled = updatedZone.Enabled + zone.EnableSearchDomain = updatedZone.EnableSearchDomain + zone.DistributionGroups = updatedZone.DistributionGroups + for _, groupID := range zone.DistributionGroups { _, err = transaction.GetGroupByID(ctx, store.LockingStrengthNone, accountID, groupID) if err != nil { @@ -134,7 +153,16 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, } if err = transaction.UpdateZone(ctx, zone); err != nil { - return fmt.Errorf("failed to update zone: %w", err) + return fmt.Errorf("update zone: %w", err) + } + + change = affectedpeers.Change{DistributionGroupIDs: slices.Concat(zone.DistributionGroups, oldGroups)} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + + if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil { + return fmt.Errorf("increment network serial: %w", err) } return nil @@ -145,7 +173,7 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string, m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneUpdated, zone.EventMeta()) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationUpdate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return zone, nil } @@ -159,13 +187,23 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID return status.NewPermissionDeniedError() } - zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) - if err != nil { - return fmt.Errorf("failed to get zone: %w", err) - } - + var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change var eventsToStore []func() + err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { + zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) + if err != nil { + return fmt.Errorf("get zone: %w", err) + } + + // Load before delete: the post-delete state no longer references the groups. + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + records, err := transaction.GetZoneDNSRecords(ctx, store.LockingStrengthNone, accountID, zoneID) if err != nil { return fmt.Errorf("failed to get records: %w", err) @@ -207,7 +245,7 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID event() } - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationDelete}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return nil } diff --git a/management/internals/modules/zones/records/manager/manager.go b/management/internals/modules/zones/records/manager/manager.go index b041aca30..16839c1b4 100644 --- a/management/internals/modules/zones/records/manager/manager.go +++ b/management/internals/modules/zones/records/manager/manager.go @@ -9,11 +9,11 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/zones/records" "github.com/netbirdio/netbird/management/server/account" "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/affectedpeers" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions/modules" "github.com/netbirdio/netbird/management/server/permissions/operations" "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/status" ) @@ -65,6 +65,8 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI } var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change record = records.NewRecord(accountID, zoneID, record.Name, record.Type, record.Content, record.TTL) err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { @@ -82,6 +84,11 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to create dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -96,7 +103,7 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordCreated, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationCreate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return record, nil } @@ -112,6 +119,8 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI var zone *zones.Zone var record *records.Record + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) @@ -141,6 +150,11 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to update dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -155,7 +169,7 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordUpdated, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationUpdate}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return record, nil } @@ -171,6 +185,8 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI var record *records.Record var zone *zones.Zone + var snap *affectedpeers.Snapshot + var change affectedpeers.Change err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID) @@ -188,6 +204,11 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI return fmt.Errorf("failed to delete dns record: %w", err) } + change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups} + if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + err = transaction.IncrementNetworkSerial(ctx, accountID) if err != nil { return fmt.Errorf("failed to increment network serial: %w", err) @@ -202,7 +223,7 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI meta := record.EventMeta(zone.ID, zone.Name) m.accountManager.StoreEvent(ctx, userID, recordID, accountID, activity.DNSRecordDeleted, meta) - go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationDelete}) + m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change) return nil } diff --git a/management/server/affected_peers_zone_test.go b/management/server/affected_peers_zone_test.go new file mode 100644 index 000000000..4d622325c --- /dev/null +++ b/management/server/affected_peers_zone_test.go @@ -0,0 +1,145 @@ +package server + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/controllers/network_map" + "github.com/netbirdio/netbird/management/internals/modules/zones" + "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/management/server/affectedpeers" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" +) + +const affectedZoneDomain = "zone.test" + +// createAffectedZone stores a zone distributed to the given groups, optionally with +// one A record so the network map actually ships it. +func createAffectedZone(t *testing.T, s store.Store, accountID, domain string, enabled, withRecord bool, groups []string) *zones.Zone { + t.Helper() + ctx := context.Background() + + zone := zones.NewZone(accountID, domain, domain, enabled, false, groups) + require.NoError(t, s.CreateZone(ctx, zone)) + + if withRecord { + record := records.NewRecord(accountID, zone.ID, "host."+domain, records.RecordTypeA, "10.0.0.1", 300) + require.NoError(t, s.CreateDNSRecord(ctx, record)) + } + + return zone +} + +func TestCollectGroupChange_ZoneLinked(t *testing.T) { + _, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]}) + + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]}) + assert.Contains(t, groups, groupIDs[0], "group distributed a zone should be affected by its own change") + + groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]}) + assert.Empty(t, groups, "group not referenced by any zone should not be affected") +} + +func TestCollectGroupChange_UnshippedZoneNotLinked(t *testing.T) { + _, s, accountID, _, groupIDs := setupAffectedPeersTest(t) + ctx := context.Background() + + // Disabled zone and zone without records are never shipped by the network map. + createAffectedZone(t, s, accountID, "disabled."+affectedZoneDomain, false, true, []string{groupIDs[0]}) + createAffectedZone(t, s, accountID, "empty."+affectedZoneDomain, true, false, []string{groupIDs[1]}) + + groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0], groupIDs[1]}) + assert.Empty(t, groups, "groups referenced only by unshipped zones should not be affected") +} + +func TestResolveAffectedPeers_ZoneGroupMembershipChange(t *testing.T) { + _, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + + createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]}) + + // Same change shape UpdateGroup builds: the group changed as a whole and peer1 + // left it, so peer1 must refresh to drop the zone. + change := affectedpeers.Change{ + ChangedGroupIDs: []string{groupIDs[0]}, + RemovedPeersByGroup: map[string][]string{groupIDs[0]: {peerIDs[1]}}, + } + + result := resolveAffected(t, s, accountID, change) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result, "current and removed members of the zone group should be affected") +} + +func TestResolveAffectedPeers_ZoneDistributionChange(t *testing.T) { + _, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t) + + // Zone create/update/delete passes old and new distribution groups. + change := affectedpeers.Change{DistributionGroupIDs: []string{groupIDs[0], groupIDs[2]}} + + result := resolveAffected(t, s, accountID, change) + assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2]}, result, "only members of the distribution groups should be affected") +} + +// TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone verifies that adding a peer +// to a group referenced only by a zone pushes the zone to the new member and leaves +// unrelated peers alone. +func TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone(t *testing.T) { + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID)) + } + + zoneGroup := &types.Group{ID: "zone-grp", Name: "ZoneGroup", Peers: []string{peer1.ID}} + require.NoError(t, manager.CreateGroup(ctx, accountID, userID, zoneGroup)) + + createAffectedZone(t, manager.Store, accountID, affectedZoneDomain, true, true, []string{zoneGroup.ID}) + + updMsg1 := updateManager.CreateChannel(ctx, peer1.ID) + updMsg2 := updateManager.CreateChannel(ctx, peer2.ID) + updMsg3 := updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + zoneGroup.Peers = []string{peer1.ID, peer2.ID} + require.NoError(t, manager.UpdateGroup(ctx, accountID, userID, zoneGroup)) + + peerShouldReceiveUpdate(t, updMsg1) + msg := receivePeerUpdate(t, updMsg2) + assert.True(t, syncHasCustomZone(msg, affectedZoneDomain+"."), "new zone group member should receive the zone") + peerShouldNotReceiveUpdate(t, updMsg3) +} + +func receivePeerUpdate(t *testing.T, ch <-chan *network_map.UpdateMessage) *network_map.UpdateMessage { + t.Helper() + select { + case msg := <-ch: + require.NotNil(t, msg, "update message should not be nil") + return msg + case <-time.After(peerUpdateTimeout): + require.FailNow(t, "timed out waiting for update message") + return nil + } +} + +func syncHasCustomZone(msg *network_map.UpdateMessage, domain string) bool { + for _, zone := range msg.Update.GetNetworkMap().GetDNSConfig().GetCustomZones() { + if zone.GetDomain() == domain { + return true + } + } + return false +} diff --git a/management/server/affectedpeers/resolver.go b/management/server/affectedpeers/resolver.go index cb2063ac9..895e4fd36 100644 --- a/management/server/affectedpeers/resolver.go +++ b/management/server/affectedpeers/resolver.go @@ -22,6 +22,7 @@ import ( nbdns "github.com/netbirdio/netbird/dns" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/internals/modules/zones" resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" networkTypes "github.com/netbirdio/netbird/management/server/networks/types" @@ -50,6 +51,7 @@ type Snapshot struct { policies []*types.Policy routes []*route.Route nsGroups []*nbdns.NameServerGroup + zones []*zones.Zone dnsSettings *types.DNSSettings routers []*routerTypes.NetworkRouter resources []*resourceTypes.NetworkResource @@ -127,12 +129,15 @@ func (snap *Snapshot) loadRoutesAndProxy(ctx context.Context, s store.Store, acc return snap.loadProxyServices(ctx, s, accountID) } -// loadDNS loads the nameserver groups and account DNS settings. +// loadDNS loads the nameserver groups, custom DNS zones and account DNS settings. func (snap *Snapshot) loadDNS(ctx context.Context, s store.Store, accountID string) error { var err error if snap.nsGroups, err = s.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID); err != nil { return err } + if snap.zones, err = s.GetAccountZones(ctx, store.LockingStrengthNone, accountID); err != nil { + return err + } snap.dnsSettings, err = s.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID) return err } @@ -357,7 +362,7 @@ func (s policySide) opposite() policySide { // - a changed router/resource/network sits on a NETWORK -> fold the SOURCE side of // the policies whose destination reaches it (and the routers it implies). // -// Routes, nameserver groups, DNS and embedded-proxy services distribute to their own +// Routes, nameserver groups, DNS zones, DNS and embedded-proxy services distribute to their own // member peers, outside the policy graph, and are folded here too. func (r *resolver) walk() { for _, policy := range r.bothSidesPolicies() { @@ -369,6 +374,7 @@ func (r *resolver) walk() { r.collectFromPolicies() r.collectFromRoutes() r.collectFromNameServers() + r.collectFromZones() r.collectFromDNSSettings() r.collectFromNetworkRouters() r.collectFromProxyServices() @@ -829,6 +835,25 @@ func (r *resolver) collectFromNameServers() { } } +// collectFromZones folds the distribution groups of the custom DNS zones that +// reference a linked group. Like nameserver groups, a zone has no opposite side, so +// only a whole-group change folds its groups. Zones the network map does not ship +// (disabled or without records) are skipped. +func (r *resolver) collectFromZones() { + if len(r.linkGroups) == 0 { + return + } + for _, zone := range r.snap.zones { + if !zone.Enabled || len(zone.Records) == 0 { + continue + } + if anyInSet(zone.DistributionGroups, r.linkGroups) { + log.WithContext(r.ctx).Tracef("collectFromZones: zone %s references a linked group -> folding its groups %v (outputGroups only)", zone.ID, zone.DistributionGroups) + r.foldOutputGroups(zone.DistributionGroups) + } + } +} + // collectFromSSHAuthorizedGroups folds the destinations of the enabled SSH rules that // authorize a group whose user membership changed. Those destination peers carry the // group -> user mapping for the groups they authorize, so they refresh even when no From 1c7d87d5fc618babc6f524dfbe885760ed208bdc Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 5 Oct 2026 14:29:28 +0200 Subject: [PATCH 11/14] [client,android] Generate debug bundle to file (#7528) * [client] Add a debug bundle file export to the Android bridge The Android app can only upload a debug bundle and hand the user a key. Users who want to inspect what leaves their device before sharing it have no way to get the zip itself. Add DebugBundleFile, which generates the bundle into the cache directory and returns its path instead of uploading; the app copies it wherever the user chose and removes it. DebugBundle keeps its behavior. Both entry points share the unexported debugBundle with an upload switch, so the body stays where it was and merges cleanly with the MDM overlay change on main. Because the file variant leaves the zip to the caller and the upload variant only removes it after the upload finishes, a process killed in between leaves a zip behind in the cache. Remove stale bundles before generating a new one: RemoveStaleBundles deletes zips matching the generator's pattern that are older than an hour. Remote debug jobs write to the same directory, so younger files are treated as still in use. * Preserve network map for debug bundle on Android * [client] Keep exported Android debug bundles out of the stale cleanup DebugBundleFile hands the zip to the caller, but the file kept the netbird.debug.*.zip name that RemoveStaleBundles matches, so a later debug run could delete it once it was older than an hour. Rename the exported bundle to netbird.debug-file.*.zip after generation so the cleanup only ever touches bundles no caller owns. * [client] Warn when a stale debug bundle cannot be removed A failed removal means bundles pile up in the cache directory, so log it at Warn instead of Debug. A file that is already gone was removed by a concurrent cleanup and is skipped silently. * [client] Drop the outdated debugBundle comment The comment still said the file variant leaves the zip in place, but it is renamed by debug.ExportBundle since the stale-cleanup change. * [client] Test that the network map reaches the debug bundle Cover both halves of the path Android now relies on: the engine keeps the latest sync response once persistence is enabled, and the bundle generator writes it to network_map.json (anonymized or not) and omits the file when there is no sync response. * [client] Remove abandoned exported debug bundles after a day An exported bundle is owned by the caller, but if the app is killed before it copies and deletes the file, nothing ever removes it from the cache directory. Let RemoveStaleBundles also match exported bundles, with a 24 hour max age instead of the caller-provided one, so a bundle that is still being saved survives while an abandoned one goes. --- client/android/client.go | 23 ++++++++++++ client/internal/debug/debug.go | 54 ++++++++++++++++++++++++++++- client/internal/debug/debug_test.go | 50 ++++++++++++++++++++++++++ client/internal/engine_test.go | 18 ++++++++++ 4 files changed, 144 insertions(+), 1 deletion(-) diff --git a/client/android/client.go b/client/android/client.go index e47a1c13d..b870337d1 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -213,6 +213,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, internal.WithNetEvents(c.netMgr)) c.setState(cfg, cacheDir, cfgFile, connectClient) + connectClient.SetSyncResponsePersistence(true) // This path runs the interactive SSO flow, so reaching here means the peer // is authenticated again — release the latch Status() reports from. Clear // only once the fresh connect client is installed: until then Status() @@ -256,6 +257,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR connectClient := internal.NewConnectClient(ctx, cfg, c.recorder, internal.WithNetEvents(c.netMgr)) c.setState(cfg, cacheDir, cfgFile, connectClient) + connectClient.SetSyncResponsePersistence(true) return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) } @@ -327,6 +329,19 @@ func (c *Client) NotifyNetworkChange() { // or "strict"; strict also anonymizes internal IP ranges, peer names, and // WireGuard public keys, and implies anonymize. func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) { + return c.debugBundle(platformFiles, anonymize, anonymizeLevel, true) +} + +// DebugBundleFile generates a debug bundle and returns the path of the zip in +// the cache directory instead of uploading it, so the app can hand the file to +// the user for inspection. The caller owns the file and removes it once done; +// the stale-bundle cleanup of later runs removes it only after a day. +// anonymize and anonymizeLevel behave as in DebugBundle. +func (c *Client) DebugBundleFile(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) { + return c.debugBundle(platformFiles, anonymize, anonymizeLevel, false) +} + +func (c *Client) debugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string, upload bool) (string, error) { cfg, cacheDir, cc := c.stateSnapshot() // If the engine hasn't been started, load config from disk @@ -342,6 +357,11 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym cacheDir = platformFiles.CacheDir() } + // Clear what an interrupted earlier run may have left in the cache before + // adding to it. Remote debug jobs write to the same directory, so anything + // younger than an hour is treated as possibly still in use. + debug.RemoveStaleBundles(cacheDir, time.Hour) + deps := debug.GeneratorDependencies{ InternalConfig: cfg, StatusRecorder: c.recorder, @@ -379,6 +399,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym if err != nil { return "", fmt.Errorf("generate debug bundle: %w", err) } + if !upload { + return debug.ExportBundle(path) + } defer func() { if err := os.Remove(path); err != nil { log.Errorf("failed to remove debug bundle file: %v", err) diff --git a/client/internal/debug/debug.go b/client/internal/debug/debug.go index b362ae293..f4d1c3598 100644 --- a/client/internal/debug/debug.go +++ b/client/internal/debug/debug.go @@ -379,9 +379,38 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen } } +// bundleFilePattern names the bundle zips Generate creates in tempDir; the +// asterisk is filled in by os.CreateTemp. +const bundleFilePattern = "netbird.debug.*.zip" + +const exportedBundlePrefix = "netbird.debug-file." + +const exportedBundleMaxAge = 24 * time.Hour + +// RemoveStaleBundles deletes bundle zips that an interrupted generation or +// upload left behind in dir. Only files older than maxAge go, so a bundle that +// another caller is still writing or uploading in the same directory survives. +// Exported bundles are kept for exportedBundleMaxAge instead. +func RemoveStaleBundles(dir string, maxAge time.Duration) { + removeStaleFiles(dir, bundleFilePattern, maxAge) + removeStaleFiles(dir, exportedBundlePrefix+"*.zip", exportedBundleMaxAge) +} + +// ExportBundle renames a generated bundle out of the RemoveStaleBundles pattern +// and returns the new path. The caller owns the file from then on; an export +// abandoned for longer than exportedBundleMaxAge is removed by RemoveStaleBundles. +func ExportBundle(path string) (string, error) { + base := strings.TrimPrefix(filepath.Base(path), strings.SplitN(bundleFilePattern, "*", 2)[0]) + exported := filepath.Join(filepath.Dir(path), exportedBundlePrefix+base) + if err := os.Rename(path, exported); err != nil { + return "", fmt.Errorf("export debug bundle: %w", err) + } + return exported, nil +} + // Generate creates a debug bundle and returns the location. func (g *BundleGenerator) Generate() (resp string, err error) { - bundlePath, err := os.CreateTemp(g.tempDir, "netbird.debug.*.zip") + bundlePath, err := os.CreateTemp(g.tempDir, bundleFilePattern) if err != nil { return "", fmt.Errorf("create zip file: %w", err) } @@ -1725,3 +1754,26 @@ func anonymizeSlice(v []any, anonymizer *anonymize.Anonymizer) []any { } return v } + +func removeStaleFiles(dir, pattern string, maxAge time.Duration) { + matches, err := filepath.Glob(filepath.Join(dir, pattern)) + if err != nil { + log.Debugf("glob stale debug bundles in %s: %v", dir, err) + return + } + + cutoff := time.Now().Add(-maxAge) + for _, path := range matches { + info, err := os.Stat(path) + if err != nil || info.ModTime().After(cutoff) { + continue + } + if err := os.Remove(path); err != nil { + if !errors.Is(err, fs.ErrNotExist) { + log.Warnf("remove stale debug bundle %s: %v", path, err) + } + continue + } + log.Infof("removed stale debug bundle %s", path) + } +} diff --git a/client/internal/debug/debug_test.go b/client/internal/debug/debug_test.go index 17d520358..6a810bccc 100644 --- a/client/internal/debug/debug_test.go +++ b/client/internal/debug/debug_test.go @@ -4,6 +4,7 @@ import ( "archive/zip" "bytes" "encoding/json" + "fmt" "net" "net/netip" "net/url" @@ -969,3 +970,52 @@ func renderAddConfigSpecific(g *BundleGenerator) string { func newAnonymizerForTest() *anonymize.Anonymizer { return anonymize.NewAnonymizer(anonymize.DefaultAddresses()) } + +func TestRemoveStaleBundles(t *testing.T) { + dir := t.TempDir() + stale := filepath.Join(dir, "netbird.debug.111.zip") + fresh := filepath.Join(dir, "netbird.debug.222.zip") + other := filepath.Join(dir, "netbird.debug.333.txt") + owned := filepath.Join(dir, "netbird.debug.444.zip") + abandoned := filepath.Join(dir, "netbird.debug.555.zip") + for _, p := range []string{stale, fresh, other, owned, abandoned} { + require.NoError(t, os.WriteFile(p, []byte("x"), 0o600)) + } + exported, err := ExportBundle(owned) + require.NoError(t, err) + exportedAbandoned, err := ExportBundle(abandoned) + require.NoError(t, err) + old := time.Now().Add(-2 * time.Hour) + for _, p := range []string{stale, other, exported} { + require.NoError(t, os.Chtimes(p, old, old)) + } + ancient := time.Now().Add(-exportedBundleMaxAge - time.Hour) + require.NoError(t, os.Chtimes(exportedAbandoned, ancient, ancient)) + + RemoveStaleBundles(dir, time.Hour) + + assert.NoFileExists(t, stale, "bundle older than maxAge should be removed") + assert.FileExists(t, fresh, "bundle younger than maxAge must survive, it may still be uploading") + assert.FileExists(t, other, "files outside the bundle pattern must not be touched") + assert.NoFileExists(t, owned) + assert.FileExists(t, exported, "exported bundle is caller-owned and must survive maxAge") + assert.NoFileExists(t, exportedAbandoned, "exported bundle older than exportedBundleMaxAge is abandoned") +} + +func TestBundleIncludesNetworkMap(t *testing.T) { + for _, anonymize := range []bool{false, true} { + t.Run(fmt.Sprintf("anonymize=%t", anonymize), func(t *testing.T) { + g := NewBundleGenerator(GeneratorDependencies{ + SyncResponse: &mgmProto.SyncResponse{NetworkMap: &mgmProto.NetworkMap{Serial: 1}}, + }, BundleConfig{Anonymize: anonymize}) + + require.Contains(t, bundleEntries(t, g), "network_map.json") + }) + } +} + +func TestBundleOmitsNetworkMapWithoutSyncResponse(t *testing.T) { + g := NewBundleGenerator(GeneratorDependencies{}, BundleConfig{}) + + require.NotContains(t, bundleEntries(t, g), "network_map.json") +} diff --git a/client/internal/engine_test.go b/client/internal/engine_test.go index 3856cae22..14076c051 100644 --- a/client/internal/engine_test.go +++ b/client/internal/engine_test.go @@ -1493,3 +1493,21 @@ func TestOverlayAddrsFromAllowedIPs(t *testing.T) { }) } } + +func TestEngine_SyncResponsePersistence(t *testing.T) { + e := &Engine{} + + _, err := e.GetLatestSyncResponse() + require.Error(t, err, "persistence is disabled by default") + + e.SetSyncResponsePersistence(true) + e.persistSyncResponse(&mgmtProto.SyncResponse{NetworkMap: &mgmtProto.NetworkMap{Serial: 7}}) + + got, err := e.GetLatestSyncResponse() + require.NoError(t, err) + assert.Equal(t, uint64(7), got.GetNetworkMap().GetSerial()) + + e.SetSyncResponsePersistence(false) + _, err = e.GetLatestSyncResponse() + require.Error(t, err) +} From 6b3cfbabd2d6efe270029e5a52769970214f4af9 Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 5 Oct 2026 15:11:08 +0200 Subject: [PATCH 12/14] [client] Resolve the Android network route peer by HA unique ID instead of scanning the full status (#7705) * [client] Build the Android network list from a single peer snapshot Networks() called GetFullStatus() once per network to find the peer that serves it, copying every peer state and taking eight recorder locks each time. With 100+ peers and the UI calling Networks() from every peer list change, this queued hundreds of callers on the status recorder lock. Take one snapshot per call and index it by route. * [client] Track the active route peer by HA unique ID in the status recorder The Android network list resolved the route owner by scanning peer state route keys. Those keys are the handler string: a prefix for static routes and the domain pattern for dynamic ones, so the prefix-based lookup never matched dynamic routes, and two networks sharing a prefix resolved to the same owner. The route watcher now records the chosen route peer under the route's HA unique ID in the status recorder, and the Android binding looks the owner up by that ID. The key is unique per network and independent of the handler string format, so both anomalies are gone. The prefix-based routeOwners helper is removed. * [client] Record the active route peer before notifying listeners AddPeerStateRoute and RemovePeerStateRoute fire the peer list change callback and wake the status subscribers. The active route peer mapping was written after those calls, so a Networks() call landing in between found no mapping for the network and fell back to the first connected peer, or kept showing the previous peer on removal. Nothing re-notified after the mapping write, so the wrong peer stayed until the next peer list change. Write and delete the mapping before the notifying calls so a listener reacting to the notification always reads the current owner. --- client/android/client.go | 16 ++++++------- client/internal/peer/status.go | 20 ++++++++++++++++ client/internal/peer/status_test.go | 23 +++++++++++++++++++ client/internal/routemanager/client/client.go | 2 ++ 4 files changed, 52 insertions(+), 9 deletions(-) diff --git a/client/android/client.go b/client/android/client.go index b870337d1..6f5eaacf3 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -498,6 +498,7 @@ func (c *Client) Networks() *NetworkArray { routesMap := routeManager.GetClientRoutesWithNetID() v6Merged := route.V6ExitMergeSet(routesMap) resolvedDomains := c.recorder.GetResolvedDomainsStates() + activeRoutePeers := c.recorder.GetActiveRoutePeers() networkArray := &NetworkArray{ items: make([]Network, 0), @@ -511,7 +512,7 @@ func (c *Client) Networks() *NetworkArray { continue } - network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged) + network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged, activeRoutePeers) if network == nil { continue } @@ -520,14 +521,14 @@ func (c *Client) Networks() *NetworkArray { return networkArray } -func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}) *Network { +func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}, activeRoutePeers map[route.HAUniqueID]string) *Network { r := routes[0] netStr := r.Network.String() if r.IsDynamic() { netStr = r.Domains.SafeString() } - routePeer, err := c.findBestRoutePeer(routes) + routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers) if err != nil { log.Errorf("could not get peer info for route %s: %v", id, err) return nil @@ -551,12 +552,9 @@ func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bo // findBestRoutePeer returns the peer actively routing traffic for the given // HA route group. Falls back to the first connected peer, then the first peer. -func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) { - netStr := routes[0].Network.String() - - fullStatus := c.recorder.GetFullStatus() - for _, p := range fullStatus.Peers { - if _, ok := p.GetRoutes()[netStr]; ok { +func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) { + if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok { + if p, err := c.recorder.GetPeer(peerKey); err == nil { return p, nil } } diff --git a/client/internal/peer/status.go b/client/internal/peer/status.go index d753ee43e..826bf6fe0 100644 --- a/client/internal/peer/status.go +++ b/client/internal/peer/status.go @@ -196,6 +196,7 @@ type Status struct { muxRelays sync.RWMutex peers map[string]State ipToKey map[string]string + activeRoutePeers map[route.HAUniqueID]string changeNotify map[string]map[string]*StatusChangeSubscription // map[peerID]map[subscriptionID]*StatusChangeSubscription signalState bool signalError error @@ -257,6 +258,7 @@ func NewRecorder(mgmAddress string) *Status { return &Status{ peers: make(map[string]State), ipToKey: make(map[string]string), + activeRoutePeers: make(map[route.HAUniqueID]string), changeNotify: make(map[string]map[string]*StatusChangeSubscription), eventStreams: make(map[string]chan *proto.SystemEvent), eventQueue: NewEventQueue(eventQueueSize), @@ -481,6 +483,24 @@ func (d *Status) RemovePeerStateRoute(peer string, route string) error { return nil } +func (d *Status) AddActiveRoutePeer(haID route.HAUniqueID, peer string) { + d.mux.Lock() + defer d.mux.Unlock() + d.activeRoutePeers[haID] = peer +} + +func (d *Status) RemoveActiveRoutePeer(haID route.HAUniqueID) { + d.mux.Lock() + defer d.mux.Unlock() + delete(d.activeRoutePeers, haID) +} + +func (d *Status) GetActiveRoutePeers() map[route.HAUniqueID]string { + d.mux.RLock() + defer d.mux.RUnlock() + return maps.Clone(d.activeRoutePeers) +} + // CheckRoutes checks if the source and destination addresses are within the same route // and returns the resource ID of the route that contains the addresses func (d *Status) CheckRoutes(ip netip.Addr) ([]byte, bool) { diff --git a/client/internal/peer/status_test.go b/client/internal/peer/status_test.go index 82dff0d6f..b3f01b217 100644 --- a/client/internal/peer/status_test.go +++ b/client/internal/peer/status_test.go @@ -9,6 +9,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/route" ) func TestAddPeer(t *testing.T) { @@ -372,3 +374,24 @@ func TestMarkServerStateDoesNotNotifyWhenUnchanged(t *testing.T) { status.MarkManagementDisconnected(err) assert.False(t, notified(ch), "redundant disconnect should not notify") } + +func TestActiveRoutePeers(t *testing.T) { + status := NewRecorder("https://mgm") + netA := route.HAUniqueID("net-a-10.0.0.0/24") + netB := route.HAUniqueID("net-b-10.0.0.0/24") + + status.AddActiveRoutePeer(netA, "peerA") + status.AddActiveRoutePeer(netB, "peerB") + + active := status.GetActiveRoutePeers() + assert.Equal(t, "peerA", active[netA]) + assert.Equal(t, "peerB", active[netB]) + + status.RemoveActiveRoutePeer(netA) + delete(active, netB) + + active = status.GetActiveRoutePeers() + _, ok := active[netA] + assert.False(t, ok) + assert.Equal(t, "peerB", active[netB]) +} diff --git a/client/internal/routemanager/client/client.go b/client/internal/routemanager/client/client.go index c691c54f8..973cf1ab8 100644 --- a/client/internal/routemanager/client/client.go +++ b/client/internal/routemanager/client/client.go @@ -294,6 +294,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error { return fmt.Errorf("add allowed IPs for peer %s: %w", route.Peer, err) } + w.statusRecorder.AddActiveRoutePeer(route.GetHAUniqueID(), route.Peer) if err := w.statusRecorder.AddPeerStateRoute(route.Peer, w.handler.String(), route.GetResourceID()); err != nil { log.Warnf("Failed to update peer state: %v", err) } @@ -303,6 +304,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error { } func (w *Watcher) removeAllowedIPs(route *route.Route, rsn reason) error { + w.statusRecorder.RemoveActiveRoutePeer(route.GetHAUniqueID()) if err := w.statusRecorder.RemovePeerStateRoute(route.Peer, w.handler.String()); err != nil { log.Warnf("Failed to update peer state: %v", err) } From bc44cdc37a3b41298187ab76799d327c8a1152cc Mon Sep 17 00:00:00 2001 From: Zoltan Papp Date: Mon, 5 Oct 2026 15:32:48 +0200 Subject: [PATCH 13/14] [client] Fix browser login popup show from go (#7408) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * [client] Show the SSO login popup and open the browser from Go The browser-login popup was created hidden and relied on its own webview to size and show itself and to launch the external browser. On macOS a hidden WKWebView gets throttled or suspended (App Nap / hidden-window throttling), so on the first-use path nothing appeared and the browser never opened, leaving the session-expiration dialog disabled until the PKCE flow timed out. Reproduced by freezing the popup's WebContent process: the old code showed nothing, the new code shows the popup and opens the browser within 30 ms regardless of the webview state. Show and focus the popup from Go right after creation and launch the browser from Go on both the create and reuse paths. The popup's frontend no longer shows or focuses itself, so the browser keeps the foreground once it activates. This also fixes the reuse path, where a fragment-only SetURL kept the mounted React tree and the once-only guard skipped opening the browser for the new URI. Browser launch failures surface in the error dialog instead of being swallowed. * [client] Show every dialog window from Go once its frontend has painted Dialog windows (browser-login, session-expiration, install-progress, welcome, error) were created hidden and made visible only by their own webview's Show call after sizing. A hidden WKWebView on macOS can be throttled or suspended before that code runs, which left the window hidden forever. The main and settings windows already avoided this with the painted event plus a fallback timer, but that timer was armed on WindowRuntimeReady, which a frozen webview never reaches either. Route all dialogs through the same mechanism: the auto-size hook emits the painted event instead of showing the window, Go shows and focuses it on that event, and a fallback timer armed at creation shows it after 3 s regardless. The browser-login popup opens the browser in an after-show callback so the browser still lands in front of the popup, also on the fallback path. * [client] Tie install-progress hidden-window restore to the current popup CloseInstallProgress nils s.installProgress before calling w.Close(), so a replacement popup can open before the old window's WindowClosing event runs. The old callback then restored the windows the replacement had just hidden, because the restore sat outside the identity check. Guard the restore with the same check the state reset uses, and restore from CloseInstallProgress itself so the programmatic close path still re-shows the hidden windows — mirroring how CloseBrowserLogin already handles it. Co-Authored-By: Claude Opus 5 (1M context) * [client] Correlate painted reports with the window generation that sent them A painted report carried only the window name, so a late report from a popup that was already closed and replaced marked its replacement ready. The replacement was then shown before its own frontend had rendered, which is the blank-dialog case this flow exists to prevent. Each dialog start URL now carries a monotonic generation token, echoed back by ReadySignal, and a report whose token no longer matches the live window is dropped. * [client] Separate a window being painted from its frontend being mounted One flag gated both showing a window and emitting to it, so the fallback timer set it for a frontend that had not subscribed yet: the queued events were flushed into a window that could not hear them, losing the login trigger and the settings tab selection. Showing is now gated on painted and emitting on mounted, and only a real frontend report sets mounted. The fallback timer also moved to its own helper so the runtime-ready hook can rearm it, giving the frontend a full budget to mount rather than sharing one with webview boot. * [client] Tag hidden windows with the popup that hid them Windows hidden while a popup owned the screen went into one untagged list, so whichever popup closed first restored all of them and emptied the list. An install started during SSO login re-showed the main window the login popup had deliberately hidden, and left the login popup with nothing to restore. Each entry now records the popup that hid it, and a restore releases only that popup's own entries. This also subsumes the manual filtering CloseRenewFlow did to keep its own session-expiration window from being re-shown. * [client] Cover the hidden-window bookkeeping with tests application.Window carries unexported methods, so the hide/restore paths could not be faked and the earlier tests could only assert which entries survived a restore, never which windows were actually shown. The bookkeeping now goes through hideableWindow, the four methods it needs, with the window enumeration and the main-window raise behind seams that are nil in production. That makes the case the owner tag exists for testable end to end: an install started during SSO login restores only the login popup it hid, and leaves the main window hidden until the login popup itself closes. * [client] Report the first paint from unstamped windows too The main and settings windows carry no generation token, so ReadySignal saw an empty generation that already matched the ref's initial value and never emitted the painted event. Those windows only became visible through the fallback timer, and their frontend was never marked mounted, so the login trigger and the requested settings tab stayed queued. Start the ref from null so the first report goes out regardless of the generation value. * [client] Hand covered windows over when a popup closes under another Closing the browser-login popup while the install-progress popup was still up restored the main window the login had hidden, even though the install popup was meant to own the screen until it finished. The owner tag on each hidden entry only stops a popup from restoring another's windows; it says nothing about what to do with its own when a second popup still covers them. Track which popups currently own the screen and, on restore, re-tag the entries another live popup covers to that popup instead of showing them. A popup is never handed its own window, so closing the popup on top still brings the one below back. --------- Co-authored-by: Claude Opus 5 (1M context) --- .../frontend/src/components/ReadySignal.tsx | 13 +- .../frontend/src/hooks/useAutoSizeWindow.ts | 24 +- .../login/LoginWaitingForBrowserDialog.tsx | 10 +- client/ui/services/connection.go | 30 +- client/ui/services/windowmanager.go | 484 +++++++++++++----- client/ui/services/windowmanager_test.go | 293 ++++++++++- 6 files changed, 695 insertions(+), 159 deletions(-) diff --git a/client/ui/frontend/src/components/ReadySignal.tsx b/client/ui/frontend/src/components/ReadySignal.tsx index 0d040cabc..6a98b30ea 100644 --- a/client/ui/frontend/src/components/ReadySignal.tsx +++ b/client/ui/frontend/src/components/ReadySignal.tsx @@ -1,4 +1,5 @@ import { useEffect, useRef } from "react"; +import { useSearchParams } from "react-router-dom"; import { Events } from "@wailsio/runtime"; import { useStatus } from "@/contexts/StatusContext.tsx"; @@ -6,13 +7,15 @@ const EVENT_WINDOW_PAINTED = "netbird:window-painted"; export const ReadySignal = () => { const { isReady } = useStatus(); - const sent = useRef(false); + const [params] = useSearchParams(); + const generation = params.get("gen") ?? ""; + const sent = useRef(null); useEffect(() => { - if (!isReady || sent.current) return; - sent.current = true; - void Events.Emit(EVENT_WINDOW_PAINTED); - }, [isReady]); + if (!isReady || sent.current === generation) return; + sent.current = generation; + void Events.Emit(EVENT_WINDOW_PAINTED, generation); + }, [isReady, generation]); return null; }; diff --git a/client/ui/frontend/src/hooks/useAutoSizeWindow.ts b/client/ui/frontend/src/hooks/useAutoSizeWindow.ts index d4f4d80b2..6623e72c2 100644 --- a/client/ui/frontend/src/hooks/useAutoSizeWindow.ts +++ b/client/ui/frontend/src/hooks/useAutoSizeWindow.ts @@ -1,23 +1,27 @@ import { useLayoutEffect, useRef } from "react"; -import { Window } from "@wailsio/runtime"; +import { useSearchParams } from "react-router-dom"; +import { Events, Window } from "@wailsio/runtime"; import i18next from "@/lib/i18n"; import { isLinux } from "@/lib/platform"; +const EVENT_WINDOW_PAINTED = "netbird:window-painted"; + // Sizes the current Wails window to the measured content height (keeping `width`), -// then shows it. Re-applies on content resize and language change. +// then reports it as painted so Go shows it. Re-applies on content resize and language change. export function useAutoSizeWindow(width: number, ready: boolean = true) { const ref = useRef(null); + const [params] = useSearchParams(); + const generation = params.get("gen") ?? ""; useLayoutEffect(() => { const el = ref.current; if (!el) return; - let shown = false; + let painted = false; let raf1 = 0; let raf2 = 0; - const showOnce = () => { - if (shown) return; - shown = true; - Window.Show().catch(() => {}); - Window.Focus().catch(() => {}); + const paintedOnce = () => { + if (painted) return; + painted = true; + Events.Emit(EVENT_WINDOW_PAINTED, generation).catch(() => {}); }; const apply = async () => { if (!ready) return; @@ -33,7 +37,7 @@ export function useAutoSizeWindow(width: number, ready: b await Window.SetMaxSize(width, targetH); } await Window.SetSize(width, targetH); - showOnce(); + paintedOnce(); } catch { // window gone / not ready — ignore } @@ -55,6 +59,6 @@ export function useAutoSizeWindow(width: number, ready: b cancelAnimationFrame(raf2); i18next.off("languageChanged", scheduleApply); }; - }, [width, ready]); + }, [width, ready, generation]); return ref; } diff --git a/client/ui/frontend/src/modules/login/LoginWaitingForBrowserDialog.tsx b/client/ui/frontend/src/modules/login/LoginWaitingForBrowserDialog.tsx index efbd1ee84..f03751c4d 100644 --- a/client/ui/frontend/src/modules/login/LoginWaitingForBrowserDialog.tsx +++ b/client/ui/frontend/src/modules/login/LoginWaitingForBrowserDialog.tsx @@ -1,4 +1,4 @@ -import { useCallback, useEffect, useRef } from "react"; +import { useCallback } from "react"; import { useTranslation } from "react-i18next"; import { useSearchParams } from "react-router-dom"; import { Events } from "@wailsio/runtime"; @@ -21,7 +21,6 @@ export default function LoginWaitingForBrowserDialog() { const [params] = useSearchParams(); const uri = params.get("uri") ?? ""; const contentRef = useAutoSizeWindow(WINDOW_WIDTH); - const openedRef = useRef(false); const reportOpenFailure = useCallback( (e: unknown) => { @@ -33,13 +32,6 @@ export default function LoginWaitingForBrowserDialog() { [t], ); - // Open the browser only after mount, or it lands on top of the still-hidden popup. - useEffect(() => { - if (!uri || openedRef.current) return; - openedRef.current = true; - Connection.OpenURL(uri).catch(reportOpenFailure); - }, [uri, reportOpenFailure]); - const tryAgain = useCallback(() => { if (!uri) return; Connection.OpenURL(uri).catch(reportOpenFailure); diff --git a/client/ui/services/connection.go b/client/ui/services/connection.go index f78ce4c0f..f6a8eca72 100644 --- a/client/ui/services/connection.go +++ b/client/ui/services/connection.go @@ -205,19 +205,7 @@ func (s *Connection) Down(ctx context.Context) error { // window.open, so the SSO verification page can't pop inline. Honors $BROWSER // before the platform default. func (s *Connection) OpenURL(url string) error { - if browser := os.Getenv("BROWSER"); browser != "" { - return exec.Command(browser, url).Start() - } - switch runtime.GOOS { - case "windows": - return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start() - case "darwin": - return exec.Command("open", url).Start() - case "linux": - return exec.Command("xdg-open", url).Start() - default: - return fmt.Errorf("unsupported platform") - } + return openURL(url) } func (s *Connection) Logout(ctx context.Context, p LogoutParams) error { @@ -288,3 +276,19 @@ func (s *Connection) waitSSOLogin(ctx context.Context, p WaitSSOParams) (string, func (s *Connection) classifyDaemonError(err error) *ClientError { return s.classifier.classify(err) } + +func openURL(url string) error { + if browser := os.Getenv("BROWSER"); browser != "" { + return exec.Command(browser, url).Start() + } + switch runtime.GOOS { + case "windows": + return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start() + case "darwin": + return exec.Command("open", url).Start() + case "linux": + return exec.Command("xdg-open", url).Start() + default: + return fmt.Errorf("unsupported platform") + } +} diff --git a/client/ui/services/windowmanager.go b/client/ui/services/windowmanager.go index 24319dae0..af6d726a3 100644 --- a/client/ui/services/windowmanager.go +++ b/client/ui/services/windowmanager.go @@ -5,6 +5,7 @@ package services import ( "net/url" "strconv" + "strings" "sync" "sync/atomic" "time" @@ -26,6 +27,16 @@ type windowOp func(w *application.WebviewWindow, created bool) type windowCloser func(w *application.WebviewWindow) +// hideableWindow is the slice of application.Window the hide/restore bookkeeping needs. +// Narrow enough to fake in tests, which application.Window itself is not: it carries +// unexported methods. +type hideableWindow interface { + Show() application.Window + Hide() application.Window + IsVisible() bool + Name() string +} + // EventTriggerLogin asks the frontend's startLogin() to begin an SSO flow. const EventTriggerLogin = "trigger-login" @@ -37,7 +48,10 @@ const EventSettingsOpen = "netbird:settings:open" const EventWindowPainted = "netbird:window-painted" -const paintedFallback = 2 * time.Second +// generationParam carries the painted-report token in each dialog's start URL. +const generationParam = "gen" + +const paintedFallback = 3 * time.Second const headlessTeardownDelay = 2 * time.Second @@ -201,6 +215,12 @@ func DialogWindowOptions(name, title, url string, linuxIcon []byte) application. } } +// hiddenWindow records a window hidden by owner, the name of the popup that hid it. +type hiddenWindow struct { + win hideableWindow + owner string +} + type WindowManager struct { app *application.App mainWindow *application.WebviewWindow @@ -213,19 +233,35 @@ type WindowManager struct { installProgress *application.WebviewWindow welcome *application.WebviewWindow errorDialog *application.WebviewWindow - // hiddenForLogin holds windows hidden while the BrowserLogin popup is open, restored on close. - hiddenForLogin []application.Window - mu sync.Mutex - newMain func(startURL string) *application.WebviewWindow - creating map[string]bool - pendingOps map[string][]windowOp - pendingClose map[string]windowCloser - restoreGen uint64 - ready map[uint]bool + // hiddenWindows holds windows hidden while a popup owns the screen, each tagged with + // the popup that hid it so closing one popup cannot restore what another still hides. + hiddenWindows []hiddenWindow + hiding map[string]bool + // allWindows and raiseMain are the seams the hide/restore tests replace; both are nil + // in production, where the Wails app and the platform helper are used directly. + allWindows func() []hideableWindow + raiseMain func() + mu sync.Mutex + newMain func(startURL string) *application.WebviewWindow + creating map[string]bool + pendingOps map[string][]windowOp + pendingClose map[string]windowCloser + restoreGen map[string]uint64 + // painted gates showing a window: set by the frontend's first render, or by the + // fallback timer so a webview that never wakes up still becomes visible. + painted map[uint]bool + // mounted gates emitting to a window: set only by a real frontend report, since an + // event emitted to a frontend that has not subscribed yet is dropped, not queued. + mounted map[uint]bool showPending map[uint]bool pendingTab map[uint]string pendingEmits map[uint][]string fallbackTimers map[uint]*time.Timer + afterShow map[uint]func() + // generation maps a window name to the token stamped into its current start URL, so a + // painted report from a replaced window can be told apart from the live one's. + generation map[string]uint64 + lastGeneration uint64 headlessMain bool headlessTimer *time.Timer // recenterOnShow is set only on the minimal-WM/XEmbed path, where the WM neither centers nor @@ -243,11 +279,16 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo creating: map[string]bool{}, pendingOps: map[string][]windowOp{}, pendingClose: map[string]windowCloser{}, - ready: map[uint]bool{}, + restoreGen: map[string]uint64{}, + hiding: map[string]bool{}, + painted: map[uint]bool{}, + mounted: map[uint]bool{}, showPending: map[uint]bool{}, pendingTab: map[uint]string{}, pendingEmits: map[uint][]string{}, fallbackTimers: map[uint]*time.Timer{}, + afterShow: map[uint]func(){}, + generation: map[string]uint64{}, } s.watchPainted() s.watchTriggerLogin() @@ -307,13 +348,13 @@ func (s *WindowManager) OpenSettings(tab string) { s.withWindow(windowSettings, &s.settings, s.newSettingsWindow, func(w *application.WebviewWindow, _ bool) { s.mu.Lock() - ready := s.ready[w.ID()] - if !ready { + mounted := s.mounted[w.ID()] + if !mounted { s.pendingTab[w.ID()] = target } s.mu.Unlock() - if ready { + if mounted { s.app.Event.Emit(EventSettingsOpen, target) } s.showWhenReady(w) @@ -327,21 +368,37 @@ func (s *WindowManager) OpenBrowserLogin(uri string) { startURL = "/#/dialog/browser-login?uri=" + url.QueryEscape(uri) } s.withWindow(windowBrowserLogin, &s.browserLogin, func() *application.WebviewWindow { - return s.newBrowserLoginWindow(startURL) + return s.newBrowserLoginWindow(s.stampGeneration(windowBrowserLogin, startURL)) }, func(w *application.WebviewWindow, created bool) { - if created { - s.centerOnCursorScreen(w) - return - } - if uri != "" { - w.SetURL(startURL) + if !created && uri != "" { + w.SetURL(s.stampGeneration(windowBrowserLogin, startURL)) } s.centerOnCursorScreen(w) - w.Show() - w.Focus() + s.showThenOpenBrowser(w, uri) }) } +func (s *WindowManager) showThenOpenBrowser(w *application.WebviewWindow, uri string) { + if uri != "" { + s.mu.Lock() + s.afterShow[w.ID()] = func() { s.openBrowser(uri) } + s.mu.Unlock() + } + s.showWhenReady(w) +} + +func (s *WindowManager) openBrowser(uri string) { + if uri == "" { + return + } + go func() { + if err := openURL(uri); err != nil { + log.Errorf("open browser for SSO login: %v", err) + s.OpenError(s.title("browserLogin.openFailedTitle"), err.Error(), "") + } + }() +} + func (s *WindowManager) newBrowserLoginWindow(startURL string) *application.WebviewWindow { s.hideOtherWindows(windowBrowserLogin) opts := DialogWindowOptions(windowBrowserLogin, s.title("window.title.signIn"), startURL, s.linuxIcon) @@ -360,12 +417,14 @@ func (s *WindowManager) newBrowserLoginWindow(startURL string) *application.Webv if userClosed { s.browserLogin = nil } + s.forgetWindowLocked(w) s.mu.Unlock() if userClosed { - s.restoreHiddenWindows() + s.restoreHiddenWindows(windowBrowserLogin) s.app.Event.Emit(EventBrowserLoginCancel) } }) + s.armReady(w) return w } @@ -386,13 +445,11 @@ func (s *WindowManager) InstallProgressWindow() *application.WebviewWindow { } func (s *WindowManager) CloseBrowserLogin() { - // The WindowClosing hook no-ops on a programmatic close, so restore here — - // but only if a popup was actually open. The frontend calls this even when no - // popup was ever shown (e.g. resetDialog() after an early RequestExtend failure, - // or connection.ts's catch path), and hiddenForLogin is shared with - // OpenInstallProgress, so an unconditional restore could re-show windows a - // still-running install-progress is hiding. - s.closeWindow(windowBrowserLogin, &s.browserLogin, s.restoreAndClose) + // The WindowClosing hook no-ops on a programmatic close, so the closer restores. + // The frontend calls this even when no popup was ever shown (resetDialog() after an + // early RequestExtend failure, or connection.ts's catch path); closeWindow skips the + // closer then, and an owner-scoped restore cannot touch what install-progress hides. + s.closeWindow(windowBrowserLogin, &s.browserLogin, s.restoringCloser(windowBrowserLogin)) } // OpenSessionExpiration shows the countdown warning on the cursor's display; seconds seeds @@ -404,16 +461,13 @@ func (s *WindowManager) OpenSessionExpiration(seconds int, deadlineUnixMilli int startURL += "&deadline=" + strconv.FormatInt(deadlineUnixMilli, 10) } s.withWindow(windowSessionExpiration, &s.sessionExpiration, func() *application.WebviewWindow { - return s.newSessionExpirationWindow(startURL) + return s.newSessionExpirationWindow(s.stampGeneration(windowSessionExpiration, startURL)) }, func(w *application.WebviewWindow, created bool) { - if created { - s.centerOnCursorScreen(w) - return + if !created { + w.SetURL(s.stampGeneration(windowSessionExpiration, startURL)) } - w.SetURL(startURL) s.centerOnCursorScreen(w) - w.Show() - w.Focus() + s.showWhenReady(w) }) } @@ -427,8 +481,10 @@ func (s *WindowManager) newSessionExpirationWindow(startURL string) *application if s.sessionExpiration == w { s.sessionExpiration = nil } + s.forgetWindowLocked(w) s.mu.Unlock() }) + s.armReady(w) return w } @@ -440,20 +496,20 @@ func (s *WindowManager) CloseSessionExpiration() { // closes the browser-login popup and the session-expiration window together. func (s *WindowManager) CloseRenewFlow() { s.mu.Lock() - bl := s.takeWindowLocked(windowBrowserLogin, &s.browserLogin, s.restoreAndClose) + bl := s.takeWindowLocked(windowBrowserLogin, &s.browserLogin, s.restoringCloser(windowBrowserLogin)) se := s.takeWindowLocked(windowSessionExpiration, &s.sessionExpiration, closeOnly) if se != nil { - kept := s.hiddenForLogin[:0] - for _, w := range s.hiddenForLogin { - if w != se { - kept = append(kept, w) + kept := s.hiddenWindows[:0] + for _, hidden := range s.hiddenWindows { + if !sameWindow(hidden.win, se) { + kept = append(kept, hidden) } } - s.hiddenForLogin = kept + s.hiddenWindows = kept } s.mu.Unlock() - s.restoreHiddenWindows() + s.restoreHiddenWindows(windowBrowserLogin) // Close after unlock so the re-entrant handlers can take s.mu. if bl != nil { bl.Close() @@ -471,14 +527,12 @@ func (s *WindowManager) OpenInstallProgress(version string) { startURL = "/#/dialog/install-progress?version=" + url.QueryEscape(version) } s.withWindow(windowInstallProgress, &s.installProgress, func() *application.WebviewWindow { - return s.newInstallProgressWindow(startURL) + return s.newInstallProgressWindow(s.stampGeneration(windowInstallProgress, startURL)) }, func(w *application.WebviewWindow, created bool) { if !created { - w.SetURL(startURL) - w.Show() - w.Focus() + w.SetURL(s.stampGeneration(windowInstallProgress, startURL)) } - s.centerWhenReady(w) + s.showWhenReady(w) }) } @@ -489,32 +543,33 @@ func (s *WindowManager) newInstallProgressWindow(startURL string) *application.W ) w.OnWindowEvent(events.Common.WindowClosing, func(_ *application.WindowEvent) { s.mu.Lock() - if s.installProgress == w { + userClosed := s.installProgress == w + if userClosed { s.installProgress = nil } + s.forgetWindowLocked(w) s.mu.Unlock() - s.restoreHiddenWindows() + if userClosed { + s.restoreHiddenWindows(windowInstallProgress) + } }) + s.armReady(w) return w } func (s *WindowManager) CloseInstallProgress() { - s.closeWindow(windowInstallProgress, &s.installProgress, closeOnly) + s.closeWindow(windowInstallProgress, &s.installProgress, s.restoringCloser(windowInstallProgress)) } // OpenWelcome shows the first-launch onboarding window. Singleton, destroyed on close. func (s *WindowManager) OpenWelcome() { - s.withWindow(windowWelcome, &s.welcome, s.newWelcomeWindow, func(w *application.WebviewWindow, created bool) { - if !created { - w.Show() - w.Focus() - } - s.centerWhenReady(w) + s.withWindow(windowWelcome, &s.welcome, s.newWelcomeWindow, func(w *application.WebviewWindow, _ bool) { + s.showWhenReady(w) }) } func (s *WindowManager) newWelcomeWindow() *application.WebviewWindow { - opts := DialogWindowOptions(windowWelcome, s.title("window.title.welcome"), "/#/dialog/welcome", s.linuxIcon) + opts := DialogWindowOptions(windowWelcome, s.title("window.title.welcome"), s.stampGeneration(windowWelcome, "/#/dialog/welcome"), s.linuxIcon) opts.Width = 420 opts.InitialPosition = application.WindowCentered w := s.app.Window.NewWithOptions(opts) @@ -523,8 +578,10 @@ func (s *WindowManager) newWelcomeWindow() *application.WebviewWindow { if s.welcome == w { s.welcome = nil } + s.forgetWindowLocked(w) s.mu.Unlock() }) + s.armReady(w) return w } @@ -542,14 +599,12 @@ func (s *WindowManager) OpenError(title, message, command string) { } startURL := errorDialogURL(title, message, command) s.withWindow(windowError, &s.errorDialog, func() *application.WebviewWindow { - return s.newErrorWindow(startURL) + return s.newErrorWindow(s.stampGeneration(windowError, startURL)) }, func(w *application.WebviewWindow, created bool) { if !created { - w.SetURL(startURL) - w.Show() - w.Focus() + w.SetURL(s.stampGeneration(windowError, startURL)) } - s.centerWhenReady(w) + s.showWhenReady(w) }) } @@ -562,8 +617,10 @@ func (s *WindowManager) newErrorWindow(startURL string) *application.WebviewWind if s.errorDialog == w { s.errorDialog = nil } + s.forgetWindowLocked(w) s.mu.Unlock() }) + s.armReady(w) return w } @@ -589,14 +646,14 @@ func (s *WindowManager) ShowMainAndEmit(event string) { s.ensureMain("/", func(w *application.WebviewWindow, _ bool) { id := w.ID() s.mu.Lock() - ready := s.ready[id] - if !ready { + mounted := s.mounted[id] + if !mounted { s.pendingEmits[id] = append(s.pendingEmits[id], event) } s.mu.Unlock() s.showWhenReady(w) - if ready { + if mounted { s.app.Event.Emit(event) } }) @@ -741,31 +798,66 @@ func (s *WindowManager) releaseCreationLocked(name string) { delete(s.pendingClose, name) } -func (s *WindowManager) restoreAndClose(w *application.WebviewWindow) { - s.restoreHiddenWindows() - w.Close() +func (s *WindowManager) restoringCloser(owner string) windowCloser { + return func(w *application.WebviewWindow) { + s.restoreHiddenWindows(owner) + w.Close() + } } +// armReady starts the fallback that shows w even if its frontend never reports a first +// render. The timer starts at creation, because a hidden webview can be suspended before +// it reaches WindowRuntimeReady — the very case this fallback covers. That makes the first +// budget cover webview boot as well, so the runtime-ready hook rearms it to give the +// frontend its own full budget to mount and paint. func (s *WindowManager) armReady(w *application.WebviewWindow) { if w == nil { return } + s.armPaintedFallback(w) w.RegisterHook(events.Common.WindowRuntimeReady, func(_ *application.WindowEvent) { - timer := time.AfterFunc(paintedFallback, func() { - log.Warnf("window %q never reported a first render, showing it anyway", w.Name()) - s.markReady(w) - }) - s.mu.Lock() - s.fallbackTimers[w.ID()] = timer - s.mu.Unlock() + s.armPaintedFallback(w) }) } +func (s *WindowManager) armPaintedFallback(w *application.WebviewWindow) { + id := w.ID() + timer := time.AfterFunc(paintedFallback, func() { + s.mu.Lock() + painted := s.painted[id] + s.mu.Unlock() + if painted { + return + } + log.Warnf("window %q never reported a first render, showing it anyway", w.Name()) + s.markPainted(w) + }) + + s.mu.Lock() + if prev := s.fallbackTimers[id]; prev != nil { + prev.Stop() + } + if s.painted[id] { + timer.Stop() + delete(s.fallbackTimers, id) + } else { + s.fallbackTimers[id] = timer + } + s.mu.Unlock() +} + func (s *WindowManager) watchPainted() { s.app.Event.On(EventWindowPainted, func(e *application.CustomEvent) { - if w := s.windowByName(e.Sender); w != nil { - s.markReady(w) + w := s.windowByName(e.Sender) + if w == nil { + return } + if !s.matchesGeneration(e.Sender, paintedGeneration(e.Data)) { + log.Debugf("ignoring stale painted report for window %q", e.Sender) + return + } + s.markPainted(w) + s.markMounted(w) }) } @@ -777,7 +869,7 @@ func (s *WindowManager) watchTriggerLogin() { s.headlessTimer = nil } w := s.mainWindow - ready := w != nil && s.ready[w.ID()] + ready := w != nil && s.mounted[w.ID()] s.mu.Unlock() if ready { return @@ -788,7 +880,7 @@ func (s *WindowManager) watchTriggerLogin() { if created { s.headlessMain = true } - pending := !s.ready[w.ID()] + pending := !s.mounted[w.ID()] if pending { s.pendingEmits[w.ID()] = append(s.pendingEmits[w.ID()], EventTriggerLogin) } @@ -850,18 +942,67 @@ func (s *WindowManager) forgetWindowLocked(w *application.WebviewWindow) { timer.Stop() } delete(s.fallbackTimers, id) - delete(s.ready, id) + delete(s.painted, id) + delete(s.mounted, id) delete(s.showPending, id) delete(s.pendingTab, id) delete(s.pendingEmits, id) + delete(s.afterShow, id) - kept := s.hiddenForLogin[:0] - for _, hidden := range s.hiddenForLogin { - if hidden != application.Window(w) { + kept := s.hiddenWindows[:0] + for _, hidden := range s.hiddenWindows { + if !sameWindow(hidden.win, w) { kept = append(kept, hidden) } } - s.hiddenForLogin = kept + s.hiddenWindows = kept +} + +func (s *WindowManager) stampGeneration(name, startURL string) string { + s.mu.Lock() + defer s.mu.Unlock() + s.lastGeneration++ + s.generation[name] = s.lastGeneration + return appendGeneration(startURL, s.lastGeneration) +} + +func (s *WindowManager) matchesGeneration(name string, gen uint64) bool { + s.mu.Lock() + defer s.mu.Unlock() + want, tracked := s.generation[name] + if !tracked { + return true + } + return want == gen +} + +func (s *WindowManager) hideableWindows() []hideableWindow { + if s.allWindows != nil { + return s.allWindows() + } + all := s.app.Window.GetAll() + windows := make([]hideableWindow, 0, len(all)) + for _, w := range all { + windows = append(windows, w) + } + return windows +} + +func (s *WindowManager) isMainWindow(w hideableWindow, mainWindow *application.WebviewWindow) bool { + if s.allWindows != nil { + return w != nil && w.Name() == windowMain + } + return sameWindow(w, mainWindow) +} + +func (s *WindowManager) raiseMainWindow(mainWindow *application.WebviewWindow) { + if s.raiseMain != nil { + s.raiseMain() + return + } + if mainWindow != nil { + raiseToForeground(mainWindow) + } } func (s *WindowManager) windowByName(name string) *application.WebviewWindow { @@ -872,24 +1013,50 @@ func (s *WindowManager) windowByName(name string) *application.WebviewWindow { return s.mainWindow case windowSettings: return s.settings + case windowBrowserLogin: + return s.browserLogin + case windowSessionExpiration: + return s.sessionExpiration + case windowInstallProgress: + return s.installProgress + case windowWelcome: + return s.welcome + case windowError: + return s.errorDialog default: return nil } } -func (s *WindowManager) markReady(w *application.WebviewWindow) { +func (s *WindowManager) markPainted(w *application.WebviewWindow) { id := w.ID() s.mu.Lock() - already := s.ready[id] - s.ready[id] = true + already := s.painted[id] + s.painted[id] = true wanted := s.showPending[id] - tab, hasTab := s.pendingTab[id] - emits := s.pendingEmits[id] + delete(s.showPending, id) if timer := s.fallbackTimers[id]; timer != nil { timer.Stop() delete(s.fallbackTimers, id) } - delete(s.showPending, id) + s.mu.Unlock() + + if already || !wanted { + return + } + s.showNow(w) +} + +// markMounted records that the window's frontend is subscribed, and flushes the events +// held back for it. The fallback timer never calls this: showing a blank window is +// recoverable, emitting into a frontend that cannot hear it is not. +func (s *WindowManager) markMounted(w *application.WebviewWindow) { + id := w.ID() + s.mu.Lock() + already := s.mounted[id] + s.mounted[id] = true + tab, hasTab := s.pendingTab[id] + emits := s.pendingEmits[id] delete(s.pendingTab, id) delete(s.pendingEmits, id) s.mu.Unlock() @@ -902,10 +1069,6 @@ func (s *WindowManager) markReady(w *application.WebviewWindow) { s.app.Event.Emit(EventSettingsOpen, tab) } - if wanted { - s.showNow(w) - } - for _, event := range emits { s.app.Event.Emit(event) } @@ -918,18 +1081,19 @@ func (s *WindowManager) showWhenReady(w *application.WebviewWindow) { id := w.ID() s.mu.Lock() - ready := s.ready[id] - if !ready { + painted := s.painted[id] + if !painted { s.showPending[id] = true } s.mu.Unlock() - if ready { + if painted { s.showNow(w) } } func (s *WindowManager) showNow(w *application.WebviewWindow) { + id := w.ID() s.mu.Lock() if w == s.mainWindow { s.headlessMain = false @@ -938,10 +1102,15 @@ func (s *WindowManager) showNow(w *application.WebviewWindow) { s.headlessTimer = nil } } + after := s.afterShow[id] + delete(s.afterShow, id) s.mu.Unlock() w.Show() w.Focus() s.centerWhenReady(w) + if after != nil { + after() + } } func (s *WindowManager) ShowMainAt(url string) { @@ -1070,13 +1239,19 @@ func (s *WindowManager) retitleAll() { } } +// hideOtherWindows hides every visible window except keepName, recording them against +// keepName so only its own restore brings them back. A window already hidden by an +// earlier popup is skipped, leaving it tagged to the popup that actually hid it. The +// per-owner generation catches a restore for keepName that ran between the snapshot and +// the record, in which case the windows are re-shown rather than stranded. func (s *WindowManager) hideOtherWindows(keepName string) { s.mu.Lock() - gen := s.restoreGen + s.hiding[keepName] = true + gen := s.restoreGen[keepName] s.mu.Unlock() - var hidden []application.Window - for _, w := range s.app.Window.GetAll() { + var hidden []hideableWindow + for _, w := range s.hideableWindows() { if w == nil || w.Name() == keepName || !w.IsVisible() { continue } @@ -1088,9 +1263,11 @@ func (s *WindowManager) hideOtherWindows(keepName string) { } s.mu.Lock() - restored := s.restoreGen != gen + restored := s.restoreGen[keepName] != gen if !restored { - s.hiddenForLogin = append(s.hiddenForLogin, hidden...) + for _, w := range hidden { + s.hiddenWindows = append(s.hiddenWindows, hiddenWindow{win: w, owner: keepName}) + } } s.mu.Unlock() if !restored { @@ -1101,33 +1278,58 @@ func (s *WindowManager) hideOtherWindows(keepName string) { } } -// restoreHiddenWindows re-shows windows hidden by hideOtherWindows. If the main -// window was among them, raiseToForeground lifts it above the SSO browser, which -// still owns the foreground — a plain Show/Focus would be demoted to a taskbar -// flash and leave it stranded behind. -func (s *WindowManager) restoreHiddenWindows() { +// restoreHiddenWindows re-shows the windows owner hid, unless another popup still covers +// them, in which case they are handed to that popup. If the main window was among them, +// raiseToForeground lifts it above the SSO browser, which still owns the foreground — a +// plain Show/Focus would be demoted to a taskbar flash and leave it stranded behind. +func (s *WindowManager) restoreHiddenWindows(owner string) { s.mu.Lock() - hidden := s.hiddenForLogin - s.hiddenForLogin = nil - s.restoreGen++ mainWindow := s.mainWindow + delete(s.hiding, owner) + var restore []hideableWindow + kept := s.hiddenWindows[:0] + for _, hidden := range s.hiddenWindows { + if hidden.owner != owner { + kept = append(kept, hidden) + continue + } + if coverer, covered := s.coveringPopupLocked(hidden.win); covered { + hidden.owner = coverer + kept = append(kept, hidden) + continue + } + if hidden.win != nil { + restore = append(restore, hidden.win) + } + } + s.hiddenWindows = kept + s.restoreGen[owner]++ s.mu.Unlock() mainRestored := false - for _, w := range hidden { - if w == nil { - continue - } + for _, w := range restore { w.Show() - if w == mainWindow { + if s.isMainWindow(w, mainWindow) { mainRestored = true } } - if mainRestored && mainWindow != nil { - raiseToForeground(mainWindow) + if mainRestored { + s.raiseMainWindow(mainWindow) } } +func (s *WindowManager) coveringPopupLocked(w hideableWindow) (string, bool) { + if w == nil { + return "", false + } + for name := range s.hiding { + if name != w.Name() { + return name, true + } + } + return "", false +} + // getScreenBasedOnCursorPosition returns the cursor's display, falling back to the // main-window screen, then nil (OS-default placement). func (s *WindowManager) getScreenBasedOnCursorPosition() *application.Screen { @@ -1169,6 +1371,48 @@ func errorDialogURL(title, message, command string) string { return startURL } +// appendGeneration adds the painted-report token to a dialog start URL, keeping any +// existing query params intact across the "/#/path?params" hash-router form. +func appendGeneration(startURL string, gen uint64) string { + sep := "?" + if strings.Contains(startURL, "?") { + sep = "&" + } + return startURL + sep + generationParam + "=" + strconv.FormatUint(gen, 10) +} + +// paintedGeneration reads the token a painted report carries back, returning 0 when the +// frontend sent none (an older bundle, or the main window, which is never stamped). +func paintedGeneration(data any) uint64 { + switch v := data.(type) { + case string: + gen, err := strconv.ParseUint(v, 10, 64) + if err != nil { + return 0 + } + return gen + case float64: + return uint64(v) + case []any: + if len(v) == 0 { + return 0 + } + return paintedGeneration(v[0]) + default: + return 0 + } +} + +// sameWindow reports whether a hidden entry refers to w, comparing through the interface +// so a nil entry never matches a live window. +func sameWindow(hidden hideableWindow, w *application.WebviewWindow) bool { + if hidden == nil || w == nil { + return false + } + other, ok := hidden.(*application.WebviewWindow) + return ok && other == w +} + // u32ptr returns a pointer to v, for the optional *uint32 Wails theme fields. func u32ptr(v uint32) *uint32 { return &v } diff --git a/client/ui/services/windowmanager_test.go b/client/ui/services/windowmanager_test.go index 13c8548ab..890fba24f 100644 --- a/client/ui/services/windowmanager_test.go +++ b/client/ui/services/windowmanager_test.go @@ -17,9 +17,65 @@ func newTestWindowManager() *WindowManager { creating: map[string]bool{}, pendingOps: map[string][]windowOp{}, pendingClose: map[string]windowCloser{}, + restoreGen: map[string]uint64{}, + hiding: map[string]bool{}, + generation: map[string]uint64{}, } } +type fakeWindow struct { + name string + visible bool + shown int + hidden int +} + +func newFakeWindow(name string) *fakeWindow { + return &fakeWindow{name: name, visible: true} +} + +func (f *fakeWindow) Show() application.Window { + f.visible = true + f.shown++ + return nil +} + +func (f *fakeWindow) Hide() application.Window { + f.visible = false + f.hidden++ + return nil +} + +func (f *fakeWindow) IsVisible() bool { return f.visible } + +func (f *fakeWindow) Name() string { return f.name } + +type fakeDesktop struct { + windows []*fakeWindow + raised int +} + +func newFakeDesktop(s *WindowManager, windows ...*fakeWindow) *fakeDesktop { + d := &fakeDesktop{windows: windows} + s.allWindows = func() []hideableWindow { + all := make([]hideableWindow, 0, len(d.windows)) + for _, w := range d.windows { + all = append(all, w) + } + return all + } + s.raiseMain = func() { d.raised++ } + return d +} + +func ownersOf(hidden []hiddenWindow) []string { + owners := make([]string, 0, len(hidden)) + for _, h := range hidden { + owners = append(owners, h.owner) + } + return owners +} + func waitDone(t *testing.T, done <-chan struct{}, msg string) { t.Helper() select { @@ -339,12 +395,245 @@ func TestCloseRenewFlowDuringBrowserLoginCreationRestoresHiddenWindows(t *testin // Seeded after the call so the deferred closer, not CloseRenewFlow's own // immediate restore, is what has to drain it. A nil entry is skipped by // restoreHiddenWindows, so no Wails window is needed. - s.hiddenForLogin = []application.Window{nil} + s.hiddenWindows = []hiddenWindow{{owner: windowBrowserLogin}} return &application.WebviewWindow{} }, func(*application.WebviewWindow, bool) {}) require.Nil(t, s.browserLogin) - require.Empty(t, s.hiddenForLogin) + require.Empty(t, s.hiddenWindows) require.Empty(t, s.creating) require.Empty(t, s.pendingClose) } + +func TestHideOtherWindowsSkipsKeepNameAndInvisible(t *testing.T) { + main := newFakeWindow(windowMain) + settings := newFakeWindow(windowSettings) + settings.visible = false + popup := newFakeWindow(windowBrowserLogin) + s := newTestWindowManager() + newFakeDesktop(s, main, settings, popup) + + s.hideOtherWindows(windowBrowserLogin) + + require.False(t, main.visible) + require.Equal(t, 1, main.hidden) + require.Equal(t, 0, settings.hidden, "an already hidden window must not be recorded") + require.Equal(t, 0, popup.hidden, "the popup itself must stay visible") + require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows)) +} + +func TestInstallDuringLoginKeepsMainHiddenUntilLoginCloses(t *testing.T) { + main := newFakeWindow(windowMain) + login := newFakeWindow(windowBrowserLogin) + install := newFakeWindow(windowInstallProgress) + install.visible = false + s := newTestWindowManager() + d := newFakeDesktop(s, main, login, install) + + s.hideOtherWindows(windowBrowserLogin) + require.False(t, main.visible) + + install.visible = true + s.hideOtherWindows(windowInstallProgress) + require.False(t, login.visible, "the install popup hides the login popup") + + s.restoreHiddenWindows(windowInstallProgress) + require.True(t, login.visible, "the install popup restores the login popup it hid") + require.False(t, main.visible, "the main window stays hidden for the login popup") + require.Equal(t, 0, d.raised) + require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows)) + + s.restoreHiddenWindows(windowBrowserLogin) + require.True(t, main.visible) + require.Equal(t, 1, d.raised, "restoring the main window raises it above the SSO browser") + require.Empty(t, s.hiddenWindows) +} + +func TestLoginClosingUnderInstallHandsMainToInstall(t *testing.T) { + main := newFakeWindow(windowMain) + login := newFakeWindow(windowBrowserLogin) + install := newFakeWindow(windowInstallProgress) + install.visible = false + s := newTestWindowManager() + d := newFakeDesktop(s, main, login, install) + + s.hideOtherWindows(windowBrowserLogin) + install.visible = true + s.hideOtherWindows(windowInstallProgress) + require.False(t, main.visible) + require.False(t, login.visible, "the install popup hides the login popup") + + // The login popup closes while the install popup is still up: the main window it + // hid must not resurface under the install popup, it is handed over instead. + s.restoreHiddenWindows(windowBrowserLogin) + require.False(t, main.visible, "the install popup still covers the main window") + require.Equal(t, 0, d.raised) + require.Equal(t, []string{windowInstallProgress, windowInstallProgress}, ownersOf(s.hiddenWindows)) + + s.restoreHiddenWindows(windowInstallProgress) + require.True(t, main.visible, "the install popup restores the handed-over main window") + require.Equal(t, 1, d.raised) + require.Empty(t, s.hiddenWindows) +} + +func TestInstallClosingUnderLoginHandsMainToLogin(t *testing.T) { + main := newFakeWindow(windowMain) + install := newFakeWindow(windowInstallProgress) + login := newFakeWindow(windowBrowserLogin) + login.visible = false + s := newTestWindowManager() + d := newFakeDesktop(s, main, install, login) + + s.hideOtherWindows(windowInstallProgress) + login.visible = true + s.hideOtherWindows(windowBrowserLogin) + require.False(t, install.visible, "the login popup hides the install popup") + + s.restoreHiddenWindows(windowInstallProgress) + require.False(t, main.visible, "the login popup still covers the main window") + require.Equal(t, 0, d.raised) + require.Equal(t, []string{windowBrowserLogin, windowBrowserLogin}, ownersOf(s.hiddenWindows)) + + s.restoreHiddenWindows(windowBrowserLogin) + require.True(t, main.visible) + require.Equal(t, 1, d.raised) + require.Empty(t, s.hiddenWindows) +} + +func TestPopupClosingReshowsTheCoveringPopupItself(t *testing.T) { + main := newFakeWindow(windowMain) + install := newFakeWindow(windowInstallProgress) + login := newFakeWindow(windowBrowserLogin) + login.visible = false + s := newTestWindowManager() + d := newFakeDesktop(s, main, install, login) + + s.hideOtherWindows(windowInstallProgress) + login.visible = true + s.hideOtherWindows(windowBrowserLogin) + + // The login popup hid the install popup itself; closing the login popup must bring + // the install popup back rather than hand it over to its own owner. + s.restoreHiddenWindows(windowBrowserLogin) + require.True(t, install.visible, "a popup is never handed over to itself") + require.False(t, main.visible, "the main window stays with the install popup") + require.Equal(t, 0, d.raised) + require.Equal(t, []string{windowInstallProgress}, ownersOf(s.hiddenWindows)) +} + +func TestRestoreHiddenWindowsUnknownOwnerKeepsEverything(t *testing.T) { + main := newFakeWindow(windowMain) + s := newTestWindowManager() + d := newFakeDesktop(s, main) + s.hideOtherWindows(windowBrowserLogin) + + s.restoreHiddenWindows(windowWelcome) + + require.False(t, main.visible) + require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows)) + require.Equal(t, 0, d.raised) +} + +func TestRestoreHiddenWindowsWithoutMainDoesNotRaise(t *testing.T) { + settings := newFakeWindow(windowSettings) + s := newTestWindowManager() + d := newFakeDesktop(s, settings) + s.hideOtherWindows(windowBrowserLogin) + + s.restoreHiddenWindows(windowBrowserLogin) + + require.True(t, settings.visible) + require.Equal(t, 0, d.raised) +} + +func TestRestoreHiddenWindowsEmptyIsNoop(t *testing.T) { + s := newTestWindowManager() + require.NotPanics(t, func() { s.restoreHiddenWindows(windowBrowserLogin) }) + require.Empty(t, s.hiddenWindows) +} + +func TestHideOtherWindowsRacingOwnRestoreReshowsWhatItHid(t *testing.T) { + main := newFakeWindow(windowMain) + s := newTestWindowManager() + d := newFakeDesktop(s, main) + enumerate := s.allWindows + // A restore for the same owner lands between the generation snapshot and the record. + s.allWindows = func() []hideableWindow { + s.restoreHiddenWindows(windowBrowserLogin) + return enumerate() + } + + s.hideOtherWindows(windowBrowserLogin) + + require.True(t, main.visible) + require.Equal(t, 1, main.hidden) + require.Empty(t, s.hiddenWindows) + require.Equal(t, 0, d.raised) +} + +func TestHideOtherWindowsIgnoresRestoreOfAnotherOwner(t *testing.T) { + main := newFakeWindow(windowMain) + s := newTestWindowManager() + newFakeDesktop(s, main) + enumerate := s.allWindows + s.allWindows = func() []hideableWindow { + s.restoreHiddenWindows(windowInstallProgress) + return enumerate() + } + + s.hideOtherWindows(windowBrowserLogin) + + require.False(t, main.visible) + require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows)) +} + +func TestRestoringCloserRestoresOnlyItsOwner(t *testing.T) { + main := newFakeWindow(windowMain) + s := newTestWindowManager() + newFakeDesktop(s, main) + s.hideOtherWindows(windowBrowserLogin) + s.hiddenWindows = append(s.hiddenWindows, hiddenWindow{owner: windowInstallProgress}) + + s.restoringCloser(windowBrowserLogin)(&application.WebviewWindow{}) + + require.True(t, main.visible) + require.Equal(t, []string{windowInstallProgress}, ownersOf(s.hiddenWindows)) +} + +func TestStampGenerationTracksLatestPerWindow(t *testing.T) { + s := newTestWindowManager() + + first := s.stampGeneration(windowBrowserLogin, "/#/dialog/browser-login") + require.Equal(t, "/#/dialog/browser-login?gen=1", first) + require.True(t, s.matchesGeneration(windowBrowserLogin, 1)) + + second := s.stampGeneration(windowBrowserLogin, "/#/dialog/browser-login?uri=x") + require.Equal(t, "/#/dialog/browser-login?uri=x&gen=2", second) + require.False(t, s.matchesGeneration(windowBrowserLogin, 1)) + require.True(t, s.matchesGeneration(windowBrowserLogin, 2)) +} + +func TestMatchesGenerationUntrackedWindowAccepts(t *testing.T) { + s := newTestWindowManager() + require.True(t, s.matchesGeneration(windowMain, 0)) +} + +func TestPaintedGeneration(t *testing.T) { + tests := []struct { + name string + data any + want uint64 + }{ + {"string", "7", 7}, + {"float", float64(7), 7}, + {"slice", []any{"7"}, 7}, + {"empty slice", []any{}, 0}, + {"unparsable", "abc", 0}, + {"nil", nil, 0}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + require.Equal(t, tc.want, paintedGeneration(tc.data)) + }) + } +} From ad03081e1fb8a65f052a6795d5812a5476a7c16c Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Mon, 5 Oct 2026 15:37:06 +0200 Subject: [PATCH 14/14] [management] Refresh only affected peers on IPv6 settings changes (#8051) Account settings changes refreshed every peer, even an IPv6 group toggle that re-addresses a few, and the refresh goroutine got the request context, so it could be cancelled when the handler returned. Group paths that reconcile IPv6 addresses only walked the changed group, so peers reaching a re-addressed peer through its other groups missed the new address. The IPv6 reconcile now returns the peers whose address changed and callers pass them as changed peers. An IPv6-only settings change dispatches affected peers, adding every IPv6 holder on a range change since the interface prefix comes from the range. IPv4 range and account-wide changes keep the full refresh with a detached context. --- management/server/account.go | 115 +++++++-- management/server/affected_peers_ipv6_test.go | 243 ++++++++++++++++++ management/server/affected_peers_user_test.go | 10 +- management/server/group.go | 46 +++- management/server/user.go | 4 +- 5 files changed, 372 insertions(+), 46 deletions(-) create mode 100644 management/server/affected_peers_ipv6_test.go diff --git a/management/server/account.go b/management/server/account.go index 1c09c8252..038c5d8db 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -334,6 +334,9 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco var groupChangesAffectPeers bool var reloadReverseProxy bool var effectiveOldNetworkRange netip.Prefix + var ipv6Changed bool + var ipv6Snap *affectedpeers.Snapshot + var ipv6Change affectedpeers.Change err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { var groupsUpdated bool @@ -379,10 +382,10 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco } if ipv6SettingsChanged(oldSettings, newSettings) { - if err = am.updatePeerIPv6Addresses(ctx, transaction, accountID, newSettings); err != nil { + if ipv6Change, err = am.applyIPv6SettingsChange(ctx, transaction, accountID, oldSettings, newSettings); err != nil { return err } - updateAccountPeers = true + ipv6Changed = true } if oldSettings.RoutingPeerDNSResolutionEnabled != newSettings.RoutingPeerDNSResolutionEnabled || @@ -419,12 +422,20 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco return err } - if updateAccountPeers || groupsUpdated { + if updateAccountPeers || groupsUpdated || ipv6Changed { if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil { return err } } + // A full account refresh already covers the IPv6 change, so the affected-peers + // snapshot is only needed when nothing account-wide changed. + if ipv6Changed && !updateAccountPeers && !groupChangesAffectPeers { + if ipv6Snap, err = affectedpeers.Load(ctx, transaction, accountID, ipv6Change); err != nil { + return fmt.Errorf("load affected peers: %w", err) + } + } + return nil }) if err != nil { @@ -486,13 +497,34 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco } } - if updateAccountPeers || extraSettingsChanged || groupChangesAffectPeers { - go am.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceAccountSettings, Operation: types.UpdateOperationUpdate}) + switch { + case updateAccountPeers || extraSettingsChanged || groupChangesAffectPeers: + go am.UpdateAccountPeers(context.WithoutCancel(ctx), accountID, types.UpdateReason{Resource: types.UpdateResourceAccountSettings, Operation: types.UpdateOperationUpdate}) + case ipv6Snap != nil: + am.ExpandAndUpdateAffected(ctx, accountID, ipv6Snap, ipv6Change) } return newSettings, nil } +// applyIPv6SettingsChange reconciles peer IPv6 addresses for new IPv6 settings and +// returns the affected-peers change: peers whose address changed refresh together +// with every peer that reaches them. On a range change every peer holding an address +// also refreshes itself, since its interface prefix comes from the account range even +// when its address stays inside the new one. +func (am *DefaultAccountManager) applyIPv6SettingsChange(ctx context.Context, transaction store.Store, accountID string, oldSettings, newSettings *types.Settings) (affectedpeers.Change, error) { + result, err := am.updatePeerIPv6Addresses(ctx, transaction, accountID, newSettings) + if err != nil { + return affectedpeers.Change{}, err + } + + change := affectedpeers.Change{ChangedPeerIDs: result.changed} + if oldSettings.NetworkRangeV6 != newSettings.NetworkRangeV6 { + change.OutputPeerIDs = result.withIPv6 + } + return change, nil +} + func ipv6SettingsChanged(old, updated *types.Settings) bool { if old.NetworkRangeV6 != updated.NetworkRangeV6 { return true @@ -1742,9 +1774,11 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth change.LinkGroups = allGroupChanges - if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges) + if err != nil { return fmt.Errorf("reconcile IPv6 for group changes: %w", err) } + change.ChangedPeerIDs = append(change.ChangedPeerIDs, ipv6Changed...) if err = transaction.IncrementNetworkSerial(ctx, userAuth.AccountId); err != nil { return fmt.Errorf("error incrementing network serial: %w", err) @@ -2334,7 +2368,8 @@ func (am *DefaultAccountManager) propagateUserGroupMemberships(ctx context.Conte return false, false, err } - if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, updatedGroups); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, updatedGroups) + if err != nil { return false, false, fmt.Errorf("reconcile IPv6 for group changes: %w", err) } @@ -2343,7 +2378,7 @@ func (am *DefaultAccountManager) propagateUserGroupMemberships(ctx context.Conte return false, false, fmt.Errorf("error checking if group changes affect peers: %w", err) } - return len(updatedGroups) > 0, peersAffected, nil + return len(updatedGroups) > 0, peersAffected || len(ipv6Changed) > 0, nil } // propagateAutoGroupsForUsers adds each user's peers to their AutoGroups where not already present. @@ -2440,56 +2475,78 @@ func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transac return nil } -func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) error { +// ipv6Reassignment reports the outcome of an IPv6 address reconciliation. +type ipv6Reassignment struct { + // changed are the peers whose IPv6 address was assigned, removed or reallocated. + changed []string + // withIPv6 are all peers holding an IPv6 address after the reconciliation. + withIPv6 []string +} + +func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) (ipv6Reassignment, error) { peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "") if err != nil { - return fmt.Errorf("get peers: %w", err) + return ipv6Reassignment{}, fmt.Errorf("get peers: %w", err) } network, err := transaction.GetAccountNetwork(ctx, store.LockingStrengthUpdate, accountID) if err != nil { - return fmt.Errorf("get network: %w", err) + return ipv6Reassignment{}, fmt.Errorf("get network: %w", err) } if err := am.ensureIPv6Subnet(ctx, transaction, accountID, settings, network); err != nil { - return err + return ipv6Reassignment{}, err } allowedPeers, err := am.buildIPv6AllowedPeers(ctx, transaction, accountID, settings) if err != nil { - return err + return ipv6Reassignment{}, err } v6Prefix, err := netip.ParsePrefix(network.NetV6.String()) if err != nil { - return fmt.Errorf("parse IPv6 prefix: %w", err) + return ipv6Reassignment{}, fmt.Errorf("parse IPv6 prefix: %w", err) } - if err := am.assignPeerIPv6Addresses(ctx, transaction, accountID, peers, network, allowedPeers, v6Prefix); err != nil { - return err + changed, err := am.assignPeerIPv6Addresses(ctx, transaction, accountID, peers, network, allowedPeers, v6Prefix) + if err != nil { + return ipv6Reassignment{}, err } - log.WithContext(ctx).Infof("updated IPv6 addresses for %d peers in account %s (groups=%d)", - len(peers), accountID, len(settings.IPv6EnabledGroups)) + result := ipv6Reassignment{changed: changed} + for _, peer := range peers { + if peer.IPv6.IsValid() { + result.withIPv6 = append(result.withIPv6, peer.ID) + } + } - return nil + log.WithContext(ctx).Infof("updated IPv6 addresses for %d of %d peers in account %s (groups=%d)", + len(changed), len(peers), accountID, len(settings.IPv6EnabledGroups)) + + return result, nil } // reconcileIPv6ForGroupChanges checks whether the given group IDs overlap with // the account's IPv6EnabledGroups. If they do, it runs a full IPv6 address // reconciliation so that peers gaining or losing membership in an IPv6-enabled -// group get their addresses assigned or removed. -func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) error { +// group get their addresses assigned or removed. It returns the peers whose IPv6 +// address changed, which callers pass as changed peers so every peer that can +// reach them refreshes. +func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) ([]string, error) { settings, err := transaction.GetAccountSettings(ctx, store.LockingStrengthNone, accountID) if err != nil { - return fmt.Errorf("get account settings: %w", err) + return nil, fmt.Errorf("get account settings: %w", err) } if !ipv6ReconcileNeeded(settings, groupIDs) { - return nil + return nil, nil } - return am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings) + result, err := am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings) + if err != nil { + return nil, err + } + return result.changed, nil } // ipv6ReconcileNeeded reports whether changes to the given groups trigger an IPv6 @@ -2528,7 +2585,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses( ctx context.Context, transaction store.Store, accountID string, peers []*nbpeer.Peer, network *types.Network, allowedPeers map[string]struct{}, v6Prefix netip.Prefix, -) error { +) ([]string, error) { takenV6 := make(map[netip.Addr]struct{}) for _, peer := range peers { if _, ok := allowedPeers[peer.ID]; ok && peer.IPv6.IsValid() && network.NetV6.Contains(peer.IPv6.AsSlice()) { @@ -2536,6 +2593,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses( } } + var changed []string for _, peer := range peers { _, allowed := allowedPeers[peer.ID] oldIPv6 := peer.IPv6 @@ -2545,7 +2603,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses( } else if !peer.IPv6.IsValid() || !network.NetV6.Contains(peer.IPv6.AsSlice()) { newIP, err := allocateIPv6WithRetry(v6Prefix, takenV6, peer.ID) if err != nil { - return err + return nil, err } peer.IPv6 = newIP } @@ -2555,10 +2613,11 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses( } if err := transaction.SavePeer(ctx, accountID, peer); err != nil { - return fmt.Errorf("save peer %s: %w", peer.ID, err) + return nil, fmt.Errorf("save peer %s: %w", peer.ID, err) } + changed = append(changed, peer.ID) } - return nil + return changed, nil } func allocateIPv6WithRetry(prefix netip.Prefix, taken map[netip.Addr]struct{}, peerID string) (netip.Addr, error) { diff --git a/management/server/affected_peers_ipv6_test.go b/management/server/affected_peers_ipv6_test.go new file mode 100644 index 000000000..c64360016 --- /dev/null +++ b/management/server/affected_peers_ipv6_test.go @@ -0,0 +1,243 @@ +package server + +import ( + "context" + "net/netip" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/controllers/network_map" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" +) + +const ( + ipv6GroupA = "ipv6-grp-a" + ipv6GroupB = "ipv6-grp-b" + ipv6GroupC = "ipv6-grp-c" + ipv6GroupD = "ipv6-grp-d" +) + +// ipv6AffectedTest holds three peers: peer1 in group A, peer2 in group B, peer3 in +// group C, with a single A<->B policy. peer3 is unrelated to peer1 and peer2. Group D +// is empty and referenced by nothing. +type ipv6AffectedTest struct { + manager *DefaultAccountManager + accountID string + peer1, peer2, peer3 *nbpeer.Peer + updMsg1, updMsg2, updMsg3 <-chan *network_map.UpdateMessage +} + +func setupIPv6AffectedTest(t *testing.T, ipv6Groups []string) *ipv6AffectedTest { + t.Helper() + + manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t) + ctx := context.Background() + accountID := account.Id + + policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + for _, p := range policies { + require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID)) + } + + for _, g := range []*types.Group{ + {ID: ipv6GroupA, Name: "IPv6-A", Peers: []string{peer1.ID}}, + {ID: ipv6GroupB, Name: "IPv6-B", Peers: []string{peer2.ID}}, + {ID: ipv6GroupC, Name: "IPv6-C", Peers: []string{peer3.ID}}, + {ID: ipv6GroupD, Name: "IPv6-D"}, + } { + require.NoError(t, manager.CreateGroup(ctx, accountID, userID, g)) + } + + _, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{ + Enabled: true, + Rules: []*types.PolicyRule{{ + Enabled: true, + Sources: []string{ipv6GroupA}, + Destinations: []string{ipv6GroupB}, + Bidirectional: true, + Action: types.PolicyTrafficActionAccept, + }}, + }, true) + require.NoError(t, err) + + // New accounts enable IPv6 for the All group; start from the requested groups. + updateIPv6TestSettings(t, manager, accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = ipv6Groups + }) + + tc := &ipv6AffectedTest{ + manager: manager, + accountID: accountID, + peer1: peer1, + peer2: peer2, + peer3: peer3, + } + tc.updMsg1 = updateManager.CreateChannel(ctx, peer1.ID) + tc.updMsg2 = updateManager.CreateChannel(ctx, peer2.ID) + tc.updMsg3 = updateManager.CreateChannel(ctx, peer3.ID) + t.Cleanup(func() { + updateManager.CloseChannel(ctx, peer1.ID) + updateManager.CloseChannel(ctx, peer2.ID) + updateManager.CloseChannel(ctx, peer3.ID) + }) + + // The setup changes above dispatch asynchronously and can land after the + // channels open, so drop them before the test acts. + drainPeerUpdates(tc.updMsg1) + drainPeerUpdates(tc.updMsg2) + drainPeerUpdates(tc.updMsg3) + + return tc +} + +// updateIPv6TestSettings applies mutate to a copy of the current settings, so only +// the mutated fields differ from what is stored. +func updateIPv6TestSettings(t *testing.T, manager *DefaultAccountManager, accountID string, mutate func(*types.Settings)) { + t.Helper() + ctx := context.Background() + + current, err := manager.Store.GetAccountSettings(ctx, store.LockingStrengthNone, accountID) + require.NoError(t, err) + + updated := current.Copy() + mutate(updated) + + _, err = manager.UpdateAccountSettings(ctx, accountID, userID, updated) + require.NoError(t, err) +} + +func (tc *ipv6AffectedTest) peerIPv6(t *testing.T, peerID string) netip.Addr { + t.Helper() + peer, err := tc.manager.Store.GetPeerByID(context.Background(), store.LockingStrengthNone, tc.accountID, peerID) + require.NoError(t, err) + return peer.IPv6 +} + +func TestAffectedPeers_IPv6GroupEnabled_RefreshesOnlyReachablePeers(t *testing.T) { + tc := setupIPv6AffectedTest(t, nil) + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = []string{ipv6GroupA} + }) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +func TestAffectedPeers_IPv6GroupDisabled_RefreshesOnlyReachablePeers(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupA}) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should start with an IPv6 address") + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = []string{} + }) + require.False(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should lose its IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +// Widening the IPv6 range keeps peer addresses, but each holder's interface prefix +// comes from the range, so holders refresh while peers that only reach them do not. +func TestAffectedPeers_IPv6RangeWidened_RefreshesAddressHolders(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupA}) + oldIPv6 := tc.peerIPv6(t, tc.peer1.ID) + require.True(t, oldIPv6.IsValid(), "peer1 should start with an IPv6 address") + + // The range is allocated on the account network; settings may leave it empty. + network, err := tc.manager.Store.GetAccountNetwork(context.Background(), store.LockingStrengthNone, tc.accountID) + require.NoError(t, err) + current := prefixFromIPNet(network.NetV6) + require.True(t, current.IsValid(), "account should have an IPv6 range") + widened := netip.PrefixFrom(current.Addr(), current.Bits()-8).Masked() + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.NetworkRangeV6 = widened + }) + require.Equal(t, oldIPv6, tc.peerIPv6(t, tc.peer1.ID), "peer1 should keep its address inside the widened range") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldNotReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +func TestAffectedPeers_IPv4RangeChange_RefreshesWholeAccount(t *testing.T) { + tc := setupIPv6AffectedTest(t, nil) + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.NetworkRange = netip.MustParsePrefix("100.70.0.0/16") + }) + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldReceiveUpdate(t, tc.updMsg3) +} + +func TestAffectedPeers_IPv6WithAccountWideChange_RefreshesWholeAccount(t *testing.T) { + tc := setupIPv6AffectedTest(t, nil) + + updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = []string{ipv6GroupA} + s.LazyConnectionEnabled = !s.LazyConnectionEnabled + }) + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldReceiveUpdate(t, tc.updMsg3) +} + +// Joining an IPv6-enabled group that no policy references gives peer1 an address. +// peer2 reaches peer1 through group A, not through the joined group, and must still +// learn the new address. +func TestAffectedPeers_GroupAddPeerIPv6_RefreshesPeersReachingThroughOtherGroups(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupD}) + + require.NoError(t, tc.manager.GroupAddPeer(context.Background(), tc.accountID, ipv6GroupD, tc.peer1.ID)) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +func TestAffectedPeers_UpdateGroupIPv6_RefreshesPeersReachingThroughOtherGroups(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupD}) + + require.NoError(t, tc.manager.UpdateGroup(context.Background(), tc.accountID, userID, &types.Group{ + ID: ipv6GroupD, + Name: "IPv6-D", + Peers: []string{tc.peer1.ID}, + })) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} + +// Deleting an IPv6-enabled group removes its members' addresses after the +// pre-delete snapshot was taken. +func TestAffectedPeers_DeleteIPv6Group_RefreshesFormerMembersAndReachablePeers(t *testing.T) { + tc := setupIPv6AffectedTest(t, []string{ipv6GroupD}) + ctx := context.Background() + + require.NoError(t, tc.manager.GroupAddPeer(ctx, tc.accountID, ipv6GroupD, tc.peer1.ID)) + require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address") + drainPeerUpdates(tc.updMsg1) + drainPeerUpdates(tc.updMsg2) + drainPeerUpdates(tc.updMsg3) + + require.NoError(t, tc.manager.DeleteGroup(ctx, tc.accountID, userID, ipv6GroupD)) + require.False(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should lose its IPv6 address") + + peerShouldReceiveUpdate(t, tc.updMsg1) + peerShouldReceiveUpdate(t, tc.updMsg2) + peerShouldNotReceiveUpdate(t, tc.updMsg3) +} diff --git a/management/server/affected_peers_user_test.go b/management/server/affected_peers_user_test.go index c0dbbb84f..3d73bbed0 100644 --- a/management/server/affected_peers_user_test.go +++ b/management/server/affected_peers_user_test.go @@ -108,11 +108,13 @@ func TestAffectedPeers_SaveUser_OnlyAffectedPeersUpdated(t *testing.T) { }) t.Run("auto group change reassigning IPv6 refreshes the changed peers and their observers", func(t *testing.T) { - account, err := manager.Store.GetAccount(ctx, accountID) - require.NoError(t, err) - account.Settings.IPv6EnabledGroups = []string{"ug-v6"} - require.NoError(t, manager.Store.SaveAccount(ctx, account)) require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "ug-v6", Name: "ug-v6"})) + // Apply through the settings API so the reconciliation that strips the other + // peers' addresses happens here, leaving the target as the only peer the + // user update reassigns. + updateIPv6TestSettings(t, manager, accountID, func(s *types.Settings) { + s.IPv6EnabledGroups = []string{"ug-v6"} + }) drainPeerUpdates(updTarget) drainPeerUpdates(upd2) diff --git a/management/server/group.go b/management/server/group.go index 88295e2f6..8d91df3ab 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -166,9 +166,11 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use return err } - if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}) + if err != nil { return err } + change.ChangedPeerIDs = ipv6Changed // A membership change does not alter which entities reference the group, so // the dependency walk runs once against the post-change snapshot. The new @@ -321,7 +323,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us var globalErr error for _, newGroup := range groups { change := affectedpeers.Change{ChangedGroupIDs: []string{newGroup.ID}} - events, snap, err := am.updateSingleGroup(ctx, accountID, userID, newGroup, change) + events, snap, change, err := am.updateSingleGroup(ctx, accountID, userID, newGroup, change) if err != nil { log.WithContext(ctx).Errorf("failed to update group %s: %v", newGroup.ID, err) if len(groups) == 1 { @@ -344,7 +346,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us return globalErr } -func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountID, userID string, newGroup *types.Group, change affectedpeers.Change) ([]func(), *affectedpeers.Snapshot, error) { +func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountID, userID string, newGroup *types.Group, change affectedpeers.Change) ([]func(), *affectedpeers.Snapshot, affectedpeers.Change, error) { var events []func() var snap *affectedpeers.Snapshot err := am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error { @@ -364,9 +366,11 @@ func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountI return err } - if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}) + if err != nil { return err } + change.ChangedPeerIDs = ipv6Changed if err := transaction.IncrementNetworkSerial(ctx, accountID); err != nil { return err @@ -377,7 +381,7 @@ func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountI snap, err = affectedpeers.Load(ctx, transaction, accountID, change) return err }) - return events, snap, err + return events, snap, change, err } // prepareGroupEvents prepares a list of event functions to be stored. @@ -480,8 +484,8 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us var allErrors error var groupIDsToDelete []string var deletedGroups []*types.Group - var snap *affectedpeers.Snapshot - var change affectedpeers.Change + var snap, ipv6Snap *affectedpeers.Snapshot + var change, ipv6Change affectedpeers.Change extraSettings, err := am.settingsManager.GetExtraSettings(ctx, accountID) if err != nil { @@ -510,10 +514,20 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us return err } - if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, groupIDsToDelete); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, groupIDsToDelete) + if err != nil { return err } + // Members of a deleted IPv6-enabled group lose their address, which the + // pre-delete snapshot cannot see, so they are resolved post-delete. + if len(ipv6Changed) > 0 { + ipv6Change = affectedpeers.Change{ChangedPeerIDs: ipv6Changed} + if ipv6Snap, err = affectedpeers.Load(ctx, transaction, accountID, ipv6Change); err != nil { + return err + } + } + return transaction.IncrementNetworkSerial(ctx, accountID) }) if err != nil { @@ -524,7 +538,7 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us am.StoreEvent(ctx, userID, group.ID, accountID, activity.GroupDeleted, group.EventMeta()) } - am.ExpandAndUpdateAffected(ctx, accountID, snap, change) + go am.dispatchAffected(ctx, accountID, []*affectedpeers.Snapshot{snap, ipv6Snap}, []affectedpeers.Change{change, ipv6Change}) return allErrors } @@ -564,11 +578,14 @@ func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, gr return err } - if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}) + if err != nil { return err } + // A peer whose IPv6 address changed is visible to every peer that reaches it + // through any of its groups, not only through this one. + change.ChangedPeerIDs = ipv6Changed - var err error if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { return err } @@ -634,11 +651,14 @@ func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID, return err } - if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}) + if err != nil { return err } + // A peer whose IPv6 address changed is visible to every peer that reaches it + // through any of its groups, not only through this one. + change.ChangedPeerIDs = ipv6Changed - var err error if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil { return err } diff --git a/management/server/user.go b/management/server/user.go index 3510a624b..5f29f4df7 100644 --- a/management/server/user.go +++ b/management/server/user.go @@ -861,9 +861,11 @@ func (am *DefaultAccountManager) processUserUpdate(ctx context.Context, transact allGroupChanges := slices.Concat(removedGroups, addedGroups) change.LinkGroups = allGroupChanges - if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, allGroupChanges); err != nil { + ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, allGroupChanges) + if err != nil { return change, nil, nil, nil, fmt.Errorf("reconcile IPv6 for group changes: %w", err) } + change.ChangedPeerIDs = append(change.ChangedPeerIDs, ipv6Changed...) } userEventsToAdd := am.prepareUserUpdateEvents(ctx, updatedUser.AccountID, initiatorUserId, oldUser, updatedUser, transferredOwnerRole, isNewUser, removedGroups, addedGroups, transaction)