diff --git a/.github/workflows/docs-ack.yml b/.github/workflows/docs-ack.yml index 7e34e2f8a..edd1eebff 100644 --- a/.github/workflows/docs-ack.yml +++ b/.github/workflows/docs-ack.yml @@ -12,6 +12,8 @@ jobs: docs-ack: name: Require docs PR URL or explicit "not needed" runs-on: ubuntu-latest + # Crowdin's translation-sync service PRs are auto-generated without the PR template. + if: github.event.pull_request.user.login != 'netbirddev' steps: - name: Read PR body diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index ff36a0854..843177b88 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -30,7 +30,7 @@ jobs: # segment by codespell and behave the same across versions; the # recursive "**" form did not take effect with the codespell shipped # by this action. - skip: go.mod,go.sum,*/proxy/web/*,*pnpm-lock.yaml,*package-lock.json,*/locales/de/*,*/locales/es/*,*/locales/fr/*,*/locales/hu/*,*/locales/it/*,*/locales/pt/*,*/locales/ru/*,*/locales/zh-CN/*,*/i18n/TRANSLATING.md + skip: go.mod,go.sum,*/proxy/web/*,*pnpm-lock.yaml,*package-lock.json,*/locales/de/*,*/locales/es/*,*/locales/fr/*,*/locales/gl/*,*/locales/hu/*,*/locales/it/*,*/locales/pt/*,*/locales/ru/*,*/locales/zh-CN/*,*/i18n/TRANSLATING.md golangci: strategy: fail-fast: false diff --git a/.github/workflows/pr-title-check.yml b/.github/workflows/pr-title-check.yml index 24d81b50f..5b769e2f7 100644 --- a/.github/workflows/pr-title-check.yml +++ b/.github/workflows/pr-title-check.yml @@ -7,6 +7,8 @@ on: jobs: check-title: runs-on: ubuntu-latest + # Crowdin's translation-sync service PRs are auto-generated with a fixed title. + if: github.event.pull_request.user.login != 'netbirddev' steps: - name: Validate PR title prefix uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9.0.0 diff --git a/.github/workflows/redhat-certify.yml b/.github/workflows/redhat-certify.yml index e592dabc2..fa0a87c64 100644 --- a/.github/workflows/redhat-certify.yml +++ b/.github/workflows/redhat-certify.yml @@ -32,6 +32,7 @@ on: - all - client-rootless - reverse-proxy + - netbird-server version: description: "Released version, e.g. v0.80.0" type: string @@ -66,6 +67,7 @@ jobs: components=( "client-rootless ghcr.io/netbirdio/netbird -rootless-ubi" "reverse-proxy ghcr.io/netbirdio/reverse-proxy -ubi" + "netbird-server ghcr.io/netbirdio/netbird-server -ubi" ) matrix="[]" missing=() diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index dee4d398d..a79357505 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -199,7 +199,7 @@ jobs: with: node-version: '22' - name: Install proxy web dependencies for license collection - # proxy/collect-licenses.sh reads the UI's license terms from node_modules. + # release_files/collect-licenses.sh -w reads the proxy UI's license terms from node_modules. working-directory: proxy/web run: npm ci --ignore-scripts - name: Set up QEMU diff --git a/.github/workflows/ui-translations.yml b/.github/workflows/ui-translations.yml index 24b7c9de2..9ac524495 100644 --- a/.github/workflows/ui-translations.yml +++ b/.github/workflows/ui-translations.yml @@ -36,7 +36,8 @@ jobs: with: node-version: "22" - # English (en) is the source of truth for translation keys; every other - # locale declared in _index.json must carry the exact same key set. + # English (en) is the source of truth for translation keys. Locales declared + # in _index.json fail on orphaned keys or placeholder mismatches; missing + # keys only warn, since they fall back to English at runtime. - name: Check translation key parity run: node client/ui/i18n/check-translations.mjs diff --git a/.goreleaser.yaml b/.goreleaser.yaml index 275c1cd7b..59d8274e9 100644 --- a/.goreleaser.yaml +++ b/.goreleaser.yaml @@ -385,7 +385,7 @@ dockers_v2: RELEASE: "{{ .Timestamp }}" hooks: pre: - - cmd: 'sh client/collect-licenses.sh "{{ .ContextDir }}/licenses" amd64 arm64' + - cmd: 'sh release_files/collect-licenses.sh -t load_wgnt_from_rsrc "{{ .ContextDir }}/licenses" ./client amd64 arm64' env: - GOOS=linux - CGO_ENABLED=0 @@ -511,6 +511,41 @@ dockers_v2: "org.opencontainers.image.revision": "{{.FullCommit}}" "org.opencontainers.image.source": "{{.GitURL}}" "maintainer": "dev@netbird.io" + - id: netbird-server-ubi + disable: "{{ .Env.SKIP_DOCKER_PUSH }}" + ids: + - netbird-server + images: + - netbirdio/netbird-server + - ghcr.io/netbirdio/netbird-server + tags: + - "{{ .Version }}-ubi" + - "{{ if eq .Env.SKIP_PUBLISH \"false\" }}ubi-latest{{ end }}" + dockerfile: combined/Dockerfile.ubi + platforms: + - linux/amd64 + - linux/arm64 + build_args: + VERSION: "{{ .Version }}" + RELEASE: "{{ .Timestamp }}" + hooks: + pre: + - cmd: 'sh release_files/collect-licenses.sh -l combined/LICENSE "{{ .ContextDir }}/licenses" ./combined amd64 arm64' + env: + - GOOS=linux + - CGO_ENABLED=1 + labels: + "org.opencontainers.image.created": "{{.Date}}" + "org.opencontainers.image.version": "{{.Version}}" + "org.opencontainers.image.revision": "{{.FullCommit}}" + "org.opencontainers.image.source": "{{.GitURL}}" + annotations: + "org.opencontainers.image.created": "{{.Date}}" + "org.opencontainers.image.title": "{{.ProjectName}}" + "org.opencontainers.image.version": "{{.Version}}" + "org.opencontainers.image.revision": "{{.FullCommit}}" + "org.opencontainers.image.source": "{{.GitURL}}" + "maintainer": "dev@netbird.io" - id: netbird-proxy disable: "{{ .Env.SKIP_DOCKER_PUSH }}" ids: @@ -552,7 +587,7 @@ dockers_v2: RELEASE: "{{ .Timestamp }}" hooks: pre: - - cmd: 'sh proxy/collect-licenses.sh "{{ .ContextDir }}/licenses" amd64 arm64' + - cmd: 'sh release_files/collect-licenses.sh -l proxy/LICENSE -w "{{ .ContextDir }}/licenses" ./proxy/cmd/proxy amd64 arm64' env: - GOOS=linux - CGO_ENABLED=0 diff --git a/LICENSE b/LICENSE index d922f155a..cea6f8f0b 100644 --- a/LICENSE +++ b/LICENSE @@ -1,4 +1,4 @@ -This BSD‑3‑Clause license applies to all parts of the repository except for the directories management/, signal/, relay/ and combined/. +This BSD-3-Clause license applies to all parts of the repository except for the directories management/, signal/, relay/ and combined/. Those directories are licensed under the GNU Affero General Public License version 3.0 (AGPLv3). See the respective LICENSE files inside each directory. BSD 3-Clause License diff --git a/client/android/client.go b/client/android/client.go index 6f5eaacf3..9705db8e0 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -104,8 +104,7 @@ type Client struct { stateChangeMu sync.Mutex stateChangeSubID string - eventSub *peer.EventSubscription - // Closed to stop the watch goroutines from delivering buffered items to a + // Closed to stop the watch goroutine from delivering buffered ticks to a // listener that has been removed or replaced. See stopStateChangeWatchLocked. stateChangeDone chan struct{} diff --git a/client/android/preferences.go b/client/android/preferences.go index 5ce31026c..3623de23f 100644 --- a/client/android/preferences.go +++ b/client/android/preferences.go @@ -46,7 +46,7 @@ func (p *Preferences) GetManagementURL() (string, error) { return p.configInput.ManagementURL, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return "", err } @@ -64,7 +64,7 @@ func (p *Preferences) GetAdminURL() (string, error) { return p.configInput.AdminURL, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return "", err } @@ -86,7 +86,7 @@ func (p *Preferences) HasPreSharedKey() (bool, error) { return *p.configInput.PreSharedKey != "", nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -112,7 +112,7 @@ func (p *Preferences) GetRosenpassEnabled() (bool, error) { return *p.configInput.RosenpassEnabled, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -133,7 +133,7 @@ func (p *Preferences) GetRosenpassPermissive() (bool, error) { return *p.configInput.RosenpassPermissive, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -149,7 +149,7 @@ func (p *Preferences) GetDisableClientRoutes() (bool, error) { return *p.configInput.DisableClientRoutes, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -170,7 +170,7 @@ func (p *Preferences) GetDisableServerRoutes() (bool, error) { return *p.configInput.DisableServerRoutes, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -188,7 +188,7 @@ func (p *Preferences) GetDisableDNS() (bool, error) { return *p.configInput.DisableDNS, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -206,7 +206,7 @@ func (p *Preferences) GetDisableFirewall() (bool, error) { return *p.configInput.DisableFirewall, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -227,7 +227,7 @@ func (p *Preferences) GetServerSSHAllowed() (bool, error) { return *p.configInput.ServerSSHAllowed, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -249,7 +249,7 @@ func (p *Preferences) GetEnableSSHRoot() (bool, error) { return *p.configInput.EnableSSHRoot, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -271,7 +271,7 @@ func (p *Preferences) GetEnableSSHSFTP() (bool, error) { return *p.configInput.EnableSSHSFTP, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -293,7 +293,7 @@ func (p *Preferences) GetEnableSSHLocalPortForwarding() (bool, error) { return *p.configInput.EnableSSHLocalPortForwarding, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -315,7 +315,7 @@ func (p *Preferences) GetEnableSSHRemotePortForwarding() (bool, error) { return *p.configInput.EnableSSHRemotePortForwarding, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -340,7 +340,7 @@ func (p *Preferences) GetBlockInbound() (bool, error) { return *p.configInput.BlockInbound, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -358,7 +358,7 @@ func (p *Preferences) GetDisableIPv6() (bool, error) { return *p.configInput.DisableIPv6, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -377,7 +377,7 @@ func (p *Preferences) GetRemoteJobsAllowed() (bool, error) { return *p.configInput.RemoteJobsAllowed, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } diff --git a/client/android/session.go b/client/android/session.go index a543db482..b2de8dadb 100644 --- a/client/android/session.go +++ b/client/android/session.go @@ -6,13 +6,8 @@ import ( "context" "fmt" - log "github.com/sirupsen/logrus" - "github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal/auth" - "github.com/netbirdio/netbird/client/internal/auth/sessionwatch" - "github.com/netbirdio/netbird/client/internal/peer" - cProto "github.com/netbirdio/netbird/client/proto" ) // StateChangeListener receives client state notifications. @@ -21,16 +16,11 @@ import ( // changed: connection state, the run-loop status label (e.g. NeedsLogin) or // the session deadline. It mirrors the daemon's SubscribeStatus stream // trigger — on each signal the consumer pulls the fresh values via -// Status() / SessionExpiresAtUnix(). -// -// OnSessionExpiring forwards the engine's session-expiry warnings, fired at -// sessionwatch.WarningLead before the deadline and again at FinalWarningLead -// (finalWarning true). The second one is suppressed when the user dismissed -// the first via DismissSessionWarning. The daemon turns the same events into -// its tray notification. +// Status() / SessionExpiresAtUnix(). The engine arms no expiry-warning +// timers on Android; the app schedules the warnings from the deadline it +// reads here. type StateChangeListener interface { OnStateChanged() - OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool) } // Status returns the connect run-loop's status label — the same value the @@ -110,11 +100,11 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) { return } - // Both subscriptions are buffered (one pending tick, ten pending events), - // so unsubscribing is not enough to stop callbacks: the loops would drain - // what is already queued and deliver it to a listener the caller has - // already removed or replaced. Gate every callback on this registration's - // own signal, which is closed before unsubscribing. + // The subscription is buffered (one pending tick), so unsubscribing is + // not enough to stop callbacks: the loop would drain what is already + // queued and deliver it to a listener the caller has already removed or + // replaced. Gate every callback on this registration's own signal, which + // is closed before unsubscribing. done := make(chan struct{}) c.stateChangeDone = done @@ -133,9 +123,6 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) { listener.OnStateChanged() } }() - - c.eventSub = c.recorder.SubscribeToEvents() - go watchSessionWarnings(c.eventSub, listener, done) } // RemoveStateChangeListener unregisters the state notification listener. @@ -145,21 +132,6 @@ func (c *Client) RemoveStateChangeListener() { c.stopStateChangeWatchLocked() } -// DismissSessionWarning records the user's "Dismiss" on the first expiry -// warning and suppresses the final one for the current deadline. A refreshed -// deadline re-arms both. No-op while the engine is not running. -func (c *Client) DismissSessionWarning() { - cc := c.getConnectClient() - if cc == nil { - return - } - engine := cc.Engine() - if engine == nil { - return - } - engine.DismissSessionWarning() -} - // ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and // asks the management server to extend the session deadline. The tunnel is // untouched: no resync, no reconnect. Async; the result arrives on the @@ -201,8 +173,8 @@ func (c *Client) CancelExtendAuthSession() { } func (c *Client) stopStateChangeWatchLocked() { - // Signal first, unsubscribe second: closing the channels only stops new - // items, and the loops would still hand whatever is buffered to a listener + // Signal first, unsubscribe second: closing the channel only stops new + // items, and the loop would still hand whatever is buffered to a listener // that is no longer registered. if c.stateChangeDone != nil { close(c.stateChangeDone) @@ -212,49 +184,6 @@ func (c *Client) stopStateChangeWatchLocked() { c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID) c.stateChangeSubID = "" } - if c.eventSub != nil { - // Closes the channel, which ends watchSessionWarnings. - c.recorder.UnsubscribeFromEvents(c.eventSub) - c.eventSub = nil - } -} - -// watchSessionWarnings forwards the engine's session-expiry warnings to the -// listener. The event stream also carries unrelated traffic — network-map -// updates on every sync, DNS and route errors — so everything but an -// AUTHENTICATION event carrying the session-warning marker is dropped. Exits -// when the subscription is closed by UnsubscribeFromEvents, or earlier when -// done is closed — the stream buffers up to ten events, and a deregistered -// listener must not receive the ones already queued. -func watchSessionWarnings(sub *peer.EventSubscription, listener StateChangeListener, done <-chan struct{}) { - for ev := range sub.Events() { - select { - case <-done: - return - default: - } - if ev.GetCategory() != cProto.SystemEvent_AUTHENTICATION { - continue - } - meta := ev.GetMetadata() - if meta[sessionwatch.MetaSessionWarning] != "true" { - // Other AUTHENTICATION events exist (e.g. a deadline rejected as - // out of range); they carry no warning marker. - continue - } - deadline, err := sessionwatch.ParseExpiresAt(meta[sessionwatch.MetaSessionExpiresAt]) - if err != nil { - log.Warnf("session warning event with unparsable deadline: %v", err) - continue - } - lead, err := sessionwatch.ParseLeadMinutes(meta[sessionwatch.MetaSessionLeadMinutes]) - if err != nil { - // Informational only — the deadline above is what drives the UI. - lead = 0 - } - listener.OnSessionExpiring(deadline.Unix(), int64(lead), - meta[sessionwatch.MetaSessionFinal] == "true") - } } func (c *Client) beginExtend() (context.Context, error) { diff --git a/client/cmd/forwarding_rules.go b/client/cmd/forwarding_rules.go deleted file mode 100644 index b3052746a..000000000 --- a/client/cmd/forwarding_rules.go +++ /dev/null @@ -1,98 +0,0 @@ -package cmd - -import ( - "fmt" - "sort" - - "github.com/spf13/cobra" - "google.golang.org/grpc/status" - - "github.com/netbirdio/netbird/client/proto" -) - -var forwardingRulesCmd = &cobra.Command{ - Use: "forwarding", - Short: "List forwarding rules", - Long: `Commands to list forwarding rules.`, -} - -var forwardingRulesListCmd = &cobra.Command{ - Use: "list", - Aliases: []string{"ls"}, - Short: "List forwarding rules", - Example: " netbird forwarding list", - Long: "Commands to list forwarding rules.", - RunE: listForwardingRules, -} - -func listForwardingRules(cmd *cobra.Command, _ []string) error { - conn, err := getClient(cmd) - if err != nil { - return err - } - defer conn.Close() - - client := proto.NewDaemonServiceClient(conn) - resp, err := client.ForwardingRules(cmd.Context(), &proto.EmptyRequest{}) - if err != nil { - return fmt.Errorf("failed to list network: %v", status.Convert(err).Message()) - } - - if len(resp.GetRules()) == 0 { - cmd.Println("No forwarding rules available.") - return nil - } - - printForwardingRules(cmd, resp.GetRules()) - return nil -} - -func printForwardingRules(cmd *cobra.Command, rules []*proto.ForwardingRule) { - cmd.Println("Available forwarding rules:") - - // Sort rules by translated address - sort.Slice(rules, func(i, j int) bool { - if rules[i].GetTranslatedAddress() != rules[j].GetTranslatedAddress() { - return rules[i].GetTranslatedAddress() < rules[j].GetTranslatedAddress() - } - if rules[i].GetProtocol() != rules[j].GetProtocol() { - return rules[i].GetProtocol() < rules[j].GetProtocol() - } - - return getFirstPort(rules[i].GetDestinationPort()) < getFirstPort(rules[j].GetDestinationPort()) - }) - - var lastIP string - for _, rule := range rules { - dPort := portToString(rule.GetDestinationPort()) - tPort := portToString(rule.GetTranslatedPort()) - if lastIP != rule.GetTranslatedAddress() { - lastIP = rule.GetTranslatedAddress() - cmd.Printf("\nTranslated peer: %s\n", rule.GetTranslatedHostname()) - } - - cmd.Printf(" Local %s/%s to %s:%s\n", rule.GetProtocol(), dPort, rule.GetTranslatedAddress(), tPort) - } -} - -func getFirstPort(portInfo *proto.PortInfo) int { - switch v := portInfo.PortSelection.(type) { - case *proto.PortInfo_Port: - return int(v.Port) - case *proto.PortInfo_Range_: - return int(v.Range.GetStart()) - default: - return 0 - } -} - -func portToString(translatedPort *proto.PortInfo) string { - switch v := translatedPort.PortSelection.(type) { - case *proto.PortInfo_Port: - return fmt.Sprintf("%d", v.Port) - case *proto.PortInfo_Range_: - return fmt.Sprintf("%d-%d", v.Range.GetStart(), v.Range.GetEnd()) - default: - return "No port specified" - } -} diff --git a/client/cmd/login.go b/client/cmd/login.go index c483ad98f..764fdd5b1 100644 --- a/client/cmd/login.go +++ b/client/cmd/login.go @@ -9,8 +9,6 @@ import ( log "github.com/sirupsen/logrus" "github.com/spf13/cobra" "golang.org/x/term" - "google.golang.org/grpc/codes" - gstatus "google.golang.org/grpc/status" "github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal/auth" @@ -145,10 +143,7 @@ func doDaemonLogin(ctx context.Context, cmd *cobra.Command, providedSetupKey str err = WithBackOff(func() error { var backOffErr error loginResp, backOffErr = client.Login(ctx, &loginRequest) - if s, ok := gstatus.FromError(backOffErr); ok && (s.Code() == codes.InvalidArgument || - s.Code() == codes.PermissionDenied || - s.Code() == codes.NotFound || - s.Code() == codes.Unimplemented) { + if terminalLoginError(backOffErr) { loginErr = backOffErr return nil } @@ -327,10 +322,28 @@ func doForegroundLogin(ctx context.Context, cmd *cobra.Command, setupKey string, } - config, err := profilemanager.ReadConfig(configFilePath) + config, err := profilemanager.ReadConfigOrDefault(configFilePath) if err != nil { return fmt.Errorf("read config file %s: %v", configFilePath, err) } + // Reading a config does not provision one: this login is about to dial + // management with the profile's identity, so mint the keys if the profile + // has none yet and put them on disk — a key that stayed in memory would + // come back different on the next run and register a second peer. + // + // Before the MDM overlay below, on purpose: the file must keep the + // profile's own values. The overlay is runtime-only and re-derived on + // every load, so persisting it would turn an enforced management URL or + // pre-shared key into one the user appears to own once the policy is + // withdrawn. + if generated, err := config.EnsureIdentity(); err != nil { + return fmt.Errorf("ensure profile identity: %v", err) + } else if generated { + if err := profilemanager.WriteOutConfig(configFilePath, config); err != nil { + return fmt.Errorf("write out config file %s: %v", configFilePath, err) + } + } + // CLI standalone login: profilemanager no longer auto-applies MDM, // so layer in the OS-native policy here. Desktop builds construct // a Loader with no fetcher — the build-tagged loadPlatform reads diff --git a/client/cmd/root.go b/client/cmd/root.go index be6479440..4525a9bd6 100644 --- a/client/cmd/root.go +++ b/client/cmd/root.go @@ -20,6 +20,8 @@ import ( "github.com/spf13/cobra" "github.com/spf13/pflag" "google.golang.org/grpc" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" "github.com/netbirdio/netbird/client/anonymize" daddr "github.com/netbirdio/netbird/client/internal/daemonaddr" @@ -175,7 +177,6 @@ func init() { rootCmd.AddCommand(versionCmd) rootCmd.AddCommand(sshCmd) rootCmd.AddCommand(networksCMD) - rootCmd.AddCommand(forwardingRulesCmd) rootCmd.AddCommand(debugCmd) rootCmd.AddCommand(profileCmd) rootCmd.AddCommand(exposeCmd) @@ -183,8 +184,6 @@ func init() { networksCMD.AddCommand(routesListCmd) networksCMD.AddCommand(routesSelectCmd, routesDeselectCmd) - forwardingRulesCmd.AddCommand(forwardingRulesListCmd) - debugCmd.AddCommand(debugBundleCmd) debugCmd.AddCommand(logCmd) logCmd.AddCommand(logLevelCmd) @@ -285,6 +284,43 @@ func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, e return grpc.DialContext(ctx, target, opts...) } +// terminalLoginError reports whether a Login failure is final, so the backoff +// cycle stops and the caller is told what the daemon said instead of "login +// backoff cycle failed" thirty seconds later. Retrying cannot change any of +// these answers: the request is malformed, the caller is not allowed, the +// target does not exist, a precondition on the daemon refuses it (the +// update-settings kill switch, an MDM-managed field), or the method is not +// implemented. +// +// Both `netbird up` and `netbird login` run Login through the backoff, and +// they each carried their own copy of this list — which is how one of them +// ended up retrying a refusal the other treated as final. +func terminalLoginError(err error) bool { + // A successful Login reaches here with a nil error, and that is not a + // terminal failure. Handled explicitly rather than left to + // gstatus.FromError, which answers (nil, true) for a nil error and leans on + // Status.Code tolerating a nil receiver to come back as codes.OK. + if err == nil { + return false + } + + s, ok := gstatus.FromError(err) + if !ok { + return false + } + + switch s.Code() { + case codes.InvalidArgument, + codes.PermissionDenied, + codes.NotFound, + codes.FailedPrecondition, + codes.Unimplemented: + return true + default: + return false + } +} + // WithBackOff execute function in backoff cycle. func WithBackOff(bf func() error) error { return backoff.RetryNotify(bf, CLIBackOffSettings, func(err error, duration time.Duration) { diff --git a/client/cmd/testutil_test.go b/client/cmd/testutil_test.go index 328a15454..46bf31837 100644 --- a/client/cmd/testutil_test.go +++ b/client/cmd/testutil_test.go @@ -6,9 +6,9 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/stretchr/testify/require" "go.opentelemetry.io/otel" + "go.uber.org/mock/gomock" "google.golang.org/grpc" "github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator" @@ -28,7 +28,6 @@ import ( mgmt "github.com/netbirdio/netbird/management/server" "github.com/netbirdio/netbird/management/server/activity" "github.com/netbirdio/netbird/management/server/groups" - "github.com/netbirdio/netbird/management/server/integrations/port_forwarding" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/settings" "github.com/netbirdio/netbird/management/server/store" @@ -124,9 +123,9 @@ func startManagement(t *testing.T, config *config.Config, testFile string) (*grp updateManager := update_channel.NewPeersUpdateManager(metrics) requestBuffer := mgmt.NewAccountRequestBuffer(ctx, store) - networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersmanager), config, nil) + networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", manager.NewEphemeralManager(store, peersmanager), config, nil) - accountManager, err := mgmt.BuildManager(ctx, config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManagerMock, false, cacheStore) + accountManager, err := mgmt.BuildManager(ctx, config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, settingsMockManager, permissionsManagerMock, false, cacheStore) if err != nil { t.Fatal(err) } diff --git a/client/cmd/up.go b/client/cmd/up.go index f5fac9749..120a25595 100644 --- a/client/cmd/up.go +++ b/client/cmd/up.go @@ -357,9 +357,17 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager // set the new config req := setupSetConfigReq(customDNSAddressConverted, cmd, activeProf.ID.String(), username.Username) if _, err := client.SetConfig(ctx, req); err != nil { - if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unavailable { - log.Warnf("setConfig method is not available in the daemon: %s", st.Message()) - } else { + switch reason, refused := refusedSettingsUpdate(err); { + case refused: + // Failing here is the point: carrying on would connect while + // silently dropping the settings the caller asked for, since + // nothing further down the line applies them. + return fmt.Errorf("the daemon refused the settings update: %s", reason) + case gstatus.Code(err) == codes.Unavailable: + // The daemon cannot serve the method at all, which is what this + // code means; an older daemon without it lands here. + log.Warnf("the daemon did not apply the settings update: %s", gstatus.Convert(err).Message()) + default: return daemonCallError("call service setConfig method", err) } } @@ -400,10 +408,7 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ err = WithBackOff(func() error { var backOffErr error loginResp, backOffErr = client.Login(ctx, loginRequest) - if s, ok := gstatus.FromError(backOffErr); ok && (s.Code() == codes.InvalidArgument || - s.Code() == codes.PermissionDenied || - s.Code() == codes.NotFound || - s.Code() == codes.Unimplemented) { + if terminalLoginError(backOffErr) { loginErr = backOffErr return nil } @@ -472,6 +477,22 @@ func setSSHSetConfigFields(req *proto.SetConfigRequest, cmd *cobra.Command) { } } +// refusedSettingsUpdate reports whether err is the daemon refusing the settings +// a request carried — the update-settings kill switch, or a field an MDM policy +// manages — and returns the reason it gave. +// +// The distinction that matters is against codes.Unavailable, which means the +// daemon cannot serve the call: that one is worth a warning, because an older +// daemon without the method lands there and the rest of `netbird up` still +// works. A refusal is not, because the settings would be silently dropped. +func refusedSettingsUpdate(err error) (string, bool) { + st, ok := gstatus.FromError(err) + if !ok || st.Code() != codes.FailedPrecondition { + return "", false + } + return st.Message(), true +} + func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, profileName, username string) *proto.SetConfigRequest { var req proto.SetConfigRequest req.ProfileName = profileName diff --git a/client/cmd/up_setconfig_refusal_test.go b/client/cmd/up_setconfig_refusal_test.go new file mode 100644 index 000000000..fdf580102 --- /dev/null +++ b/client/cmd/up_setconfig_refusal_test.go @@ -0,0 +1,85 @@ +package cmd + +import ( + "errors" + "testing" + + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + gstatus "google.golang.org/grpc/status" +) + +// A refused settings update has to fail `netbird up`, or a caller that asked +// for a setting the daemon will not apply connects as if it had been applied. +// The daemon being unable to serve the call is the case that stays a warning. +func TestRefusedSettingsUpdate(t *testing.T) { + tests := []struct { + name string + err error + wantRefused bool + }{ + { + name: "the kill switch refused the change", + err: gstatus.Errorf(codes.FailedPrecondition, "update settings are disabled, you cannot use this feature without update settings enabled"), + wantRefused: true, + }, + { + name: "an MDM policy manages the field", + err: gstatus.Errorf(codes.FailedPrecondition, "fields managed by MDM policy: managementURL"), + wantRefused: true, + }, + { + name: "the daemon cannot serve the call", + err: gstatus.Errorf(codes.Unavailable, "connection refused"), + wantRefused: false, + }, + { + name: "any other RPC failure", + err: gstatus.Errorf(codes.Internal, "boom"), + wantRefused: false, + }, + { + name: "not a status error at all", + err: errors.New("boom"), + wantRefused: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + reason, refused := refusedSettingsUpdate(tt.err) + require.Equal(t, tt.wantRefused, refused) + if tt.wantRefused { + require.Equal(t, gstatus.Convert(tt.err).Message(), reason, "the daemon's reason must reach the caller") + } + }) + } +} + +// Both `netbird up` and `netbird login` drive Login through the backoff cycle, +// and a final answer has to stop it: retrying a refusal only replaces the +// daemon's reason with "login backoff cycle failed" thirty seconds later. +func TestTerminalLoginError(t *testing.T) { + tests := []struct { + name string + err error + wantTerminal bool + }{ + {name: "settings refused by the kill switch", err: gstatus.Errorf(codes.FailedPrecondition, "update settings are disabled"), wantTerminal: true}, + {name: "field managed by MDM", err: gstatus.Errorf(codes.FailedPrecondition, "fields managed by MDM policy: managementURL"), wantTerminal: true}, + {name: "caller not allowed", err: gstatus.Errorf(codes.PermissionDenied, "nope"), wantTerminal: true}, + {name: "malformed request", err: gstatus.Errorf(codes.InvalidArgument, "nope"), wantTerminal: true}, + {name: "profile not found", err: gstatus.Errorf(codes.NotFound, "nope"), wantTerminal: true}, + {name: "method missing on an older daemon", err: gstatus.Errorf(codes.Unimplemented, "nope"), wantTerminal: true}, + {name: "daemon unreachable, worth retrying", err: gstatus.Errorf(codes.Unavailable, "connection refused"), wantTerminal: false}, + {name: "transient internal failure", err: gstatus.Errorf(codes.Internal, "boom"), wantTerminal: false}, + {name: "not a status error", err: errors.New("boom"), wantTerminal: false}, + {name: "no error at all, the login succeeded", err: nil, wantTerminal: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.wantTerminal, terminalLoginError(tt.err)) + }) + } +} diff --git a/client/collect-licenses.sh b/client/collect-licenses.sh deleted file mode 100644 index 7dfabada9..000000000 --- a/client/collect-licenses.sh +++ /dev/null @@ -1,77 +0,0 @@ -#!/bin/sh -set -eu - -if [ "$#" -lt 2 ]; then - printf '%s\n' "usage: $0 OUTPUT_DIRECTORY GOARCH..." >&2 - exit 2 -fi - -repo_root=$(CDPATH= cd -- "$(dirname "$0")/.." && pwd) -output_name=$(basename "$1") -if [ -z "$output_name" ] || [ "$output_name" = "." ] || - [ "$output_name" = ".." ] || [ "$output_name" = "/" ]; then - printf '%s\n' "OUTPUT_DIRECTORY must name a directory" >&2 - exit 2 -fi -output_parent=$(CDPATH= cd -- "$(dirname "$1")" && pwd) -output="$output_parent/$output_name" -shift -modules=$(mktemp "${TMPDIR:-/tmp}/netbird-client-licenses.modules.XXXXXX") -sorted_modules=$(mktemp "${TMPDIR:-/tmp}/netbird-client-licenses.sorted.XXXXXX") -trap 'rm -f "$modules" "$sorted_modules"' EXIT HUP INT TERM - -if [ -e "$output" ] || [ -L "$output" ]; then - printf 'output directory already exists: %s\n' "$output" >&2 - exit 1 -fi -mkdir "$output" -mkdir "$output/third_party" - -cp "$repo_root/LICENSE" "$output/BSD-3-Clause.txt" - -cd "$repo_root" -for arch in "$@"; do - GOOS=${GOOS:-linux} GOARCH="$arch" CGO_ENABLED=${CGO_ENABLED:-0} \ - go list -deps -f '{{with .Module}}{{if .Replace}}{{.Replace.Path}}{{"\t"}}{{.Replace.Version}}{{"\t"}}{{.Replace.Dir}}{{else}}{{.Path}}{{"\t"}}{{.Version}}{{"\t"}}{{.Dir}}{{end}}{{end}}' -tags load_wgnt_from_rsrc ./client >>"$modules" -done -LC_ALL=C sort -u "$modules" >"$sorted_modules" - -goroot=$(go env GOROOT) -for term in LICENSE PATENTS; do - if [ ! -f "$goroot/$term" ]; then - printf 'missing Go standard-library term: %s\n' "$goroot/$term" >&2 - exit 1 - fi - cp "$goroot/$term" "$output/Go-$term" -done - -while IFS=' ' read -r module version module_dir; do - [ -n "$module" ] || continue - [ "$module" = "github.com/netbirdio/netbird" ] && continue - - if [ -z "$version" ] || [ ! -d "$module_dir" ]; then - printf 'cannot collect terms for module %s at version %s\n' "$module" "$version" >&2 - exit 1 - fi - - destination="$output/third_party/$module/$version" - mkdir -p "$destination" - printf 'module: %s\nversion: %s\n' "$module" "$version" >"$destination/MODULE" - - found=false - for term in \ - "$module_dir"/LICENSE* "$module_dir"/License* "$module_dir"/license* \ - "$module_dir"/LICENCE* "$module_dir"/Licence* "$module_dir"/licence* \ - "$module_dir"/COPYING* "$module_dir"/Copying* "$module_dir"/copying* \ - "$module_dir"/NOTICE* "$module_dir"/Notice* "$module_dir"/notice* \ - "$module_dir"/PATENTS* "$module_dir"/Patents* "$module_dir"/patents*; do - [ -f "$term" ] || continue - cp "$term" "$destination/" - found=true - done - - if [ "$found" = false ]; then - printf 'no root license terms found for module %s at %s\n' "$module" "$module_dir" >&2 - exit 1 - fi -done <"$sorted_modules" diff --git a/client/embed/embed_test.go b/client/embed/embed_test.go index 4ff5c9978..a818af055 100644 --- a/client/embed/embed_test.go +++ b/client/embed/embed_test.go @@ -6,8 +6,8 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" "google.golang.org/grpc" "github.com/netbirdio/netbird/management/internals/controllers/network_map/controller" @@ -21,7 +21,6 @@ import ( nbcache "github.com/netbirdio/netbird/management/server/cache" "github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator" - "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" @@ -146,8 +145,8 @@ func startManagement(t *testing.T, signalAddr string) string { updateManager := update_channel.NewPeersUpdateManager(metrics) requestBuffer := mgmt.NewAccountRequestBuffer(context.Background(), testStore) - networkMapController := controller.NewController(context.Background(), testStore, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(testStore, peersManager), cfg, nil) - accountManager, err := mgmt.BuildManager(context.Background(), cfg, testStore, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore) + networkMapController := controller.NewController(context.Background(), testStore, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", manager.NewEphemeralManager(testStore, peersManager), cfg, nil) + accountManager, err := mgmt.BuildManager(context.Background(), cfg, testStore, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, settingsMockManager, permissionsManager, false, cacheStore) require.NoError(t, err) secretsManager, err := nbgrpc.NewTimeBasedAuthSecretsManager(updateManager, cfg.TURNConfig, cfg.Relay, settingsMockManager, groupsManager) diff --git a/client/firewall/iptables/dnat_linux.go b/client/firewall/iptables/dnat_linux.go index eca8386c0..f118c9dfe 100644 --- a/client/firewall/iptables/dnat_linux.go +++ b/client/firewall/iptables/dnat_linux.go @@ -8,177 +8,11 @@ import ( "strconv" "strings" - "github.com/hashicorp/go-multierror" log "github.com/sirupsen/logrus" - nberrors "github.com/netbirdio/netbird/client/errors" firewall "github.com/netbirdio/netbird/client/firewall/manager" ) -func (r *family) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { - ruleID := rule.ID() - if _, exists := r.rules[ruleID+dnatSuffix]; exists { - return rule, nil - } - - toDestination := rule.TranslatedAddress.String() - switch { - case len(rule.TranslatedPort.Values) == 0: - // no translated port, use original port - case len(rule.TranslatedPort.Values) == 1: - toDestination += fmt.Sprintf(":%d", rule.TranslatedPort.Values[0]) - case rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2: - // need the "/originalport" suffix to avoid dnat port randomization - toDestination += fmt.Sprintf(":%d-%d/%d", rule.TranslatedPort.Values[0], rule.TranslatedPort.Values[1], rule.DestinationPort.Values[0]) - default: - return nil, fmt.Errorf("invalid translated port: %v", rule.TranslatedPort) - } - - proto := strings.ToLower(string(rule.Protocol)) - - rules := make(map[firewall.RuleID]ruleInfo, 3) - - // DNAT rule - dnatRule := []string{ - "!", "-i", r.wgIface.Name(), - "-p", proto, - "-j", "DNAT", - "--to-destination", toDestination, - } - dnatRule = append(dnatRule, applyPort("--dport", &rule.DestinationPort)...) - rules[ruleID+dnatSuffix] = ruleInfo{ - table: tableNat, - chain: chainRTRdr, - rule: dnatRule, - } - - // SNAT rule - snatRule := []string{ - "-o", r.wgIface.Name(), - "-p", proto, - "-d", rule.TranslatedAddress.String(), - "-j", "MASQUERADE", - } - snatRule = append(snatRule, applyPort("--dport", &rule.TranslatedPort)...) - rules[ruleID+snatSuffix] = ruleInfo{ - table: tableNat, - chain: chainRTNAT, - rule: snatRule, - } - - // Forward filtering rule, if fwd policy is DROP - forwardRule := []string{ - "-o", r.wgIface.Name(), - "-p", proto, - "-d", rule.TranslatedAddress.String(), - "-j", "ACCEPT", - } - forwardRule = append(forwardRule, applyPort("--dport", &rule.TranslatedPort)...) - rules[ruleID+fwdSuffix] = ruleInfo{ - table: tableFilter, - chain: chainRTFwdOut, - rule: forwardRule, - } - - for key, ruleInfo := range rules { - if err := r.iptablesClient.Append(ruleInfo.table, ruleInfo.chain, ruleInfo.rule...); err != nil { - r.cleanupFailedDNATAdd(rules) - return nil, fmt.Errorf("add rule %s: %w", key, err) - } - r.rules[key] = ruleInfo.rule - } - - if err := r.ipFwdState.RequestForwarding(r.v6); err != nil { - r.cleanupFailedDNATAdd(rules) - return nil, fmt.Errorf("enable forwarding: %w", err) - } - - r.updateState() - return rule, nil -} - -// cleanupFailedDNATAdd removes the bookkeeping written by a partially applied -// AddDNATRule before rolling back the kernel rules, so no entries remain that -// never got a forwarding refcount. rollbackRules re-adds entries it failed to -// remove from the kernel. -func (r *family) cleanupFailedDNATAdd(rules map[firewall.RuleID]ruleInfo) { - for key := range rules { - delete(r.rules, key) - } - if err := r.rollbackRules(rules); err != nil { - log.Errorf("rollback failed: %v", err) - } -} - -func (r *family) rollbackRules(rules map[firewall.RuleID]ruleInfo) error { - var merr *multierror.Error - for key, ruleInfo := range rules { - if err := r.iptablesClient.DeleteIfExists(ruleInfo.table, ruleInfo.chain, ruleInfo.rule...); err != nil { - merr = multierror.Append(merr, fmt.Errorf("rollback rule %s: %w", key, err)) - // On rollback error, add to rules map for next cleanup - r.rules[key] = ruleInfo.rule - } - } - if merr != nil { - r.updateState() - } - return nberrors.FormatErrorOrNil(merr) -} - -func (r *family) DeleteDNATRule(rule firewall.Rule) error { - ruleID := rule.ID() - - _, hadDNAT := r.rules[ruleID+dnatSuffix] - _, hadSNAT := r.rules[ruleID+snatSuffix] - _, hadFWD := r.rules[ruleID+fwdSuffix] - if !hadDNAT && !hadSNAT && !hadFWD { - return nil - } - - var merr *multierror.Error - if dnatRule, exists := r.rules[ruleID+dnatSuffix]; exists { - if err := r.iptablesClient.Delete(tableNat, chainRTRdr, dnatRule...); err != nil { - merr = multierror.Append(merr, fmt.Errorf("delete DNAT rule: %w", err)) - } else { - delete(r.rules, ruleID+dnatSuffix) - } - } - - if snatRule, exists := r.rules[ruleID+snatSuffix]; exists { - if err := r.iptablesClient.Delete(tableNat, chainRTNAT, snatRule...); err != nil { - merr = multierror.Append(merr, fmt.Errorf("delete SNAT rule: %w", err)) - } else { - delete(r.rules, ruleID+snatSuffix) - } - } - - if fwdRule, exists := r.rules[ruleID+fwdSuffix]; exists { - if err := r.iptablesClient.Delete(tableFilter, chainRTFwdOut, fwdRule...); err != nil { - merr = multierror.Append(merr, fmt.Errorf("delete forward rule: %w", err)) - } else { - delete(r.rules, ruleID+fwdSuffix) - } - } - - // Release the refcount only once all rules are gone from the kernel. On - // partial failure the failed entries stay in r.rules so a retry can remove - // them and release then. - if merr == nil { - r.releaseForwarding() - } - - r.updateState() - - return nberrors.FormatErrorOrNil(merr) -} - -// releaseForwarding drops one IP forwarding reference, logging any error. -func (r *family) releaseForwarding() { - if err := r.ipFwdState.ReleaseForwarding(r.v6); err != nil { - log.Errorf("release IP forwarding: %v", err) - } -} - func (r *family) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error { ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort)) diff --git a/client/firewall/iptables/dnat_refcount_linux_test.go b/client/firewall/iptables/dnat_refcount_linux_test.go deleted file mode 100644 index 40ebc6cc3..000000000 --- a/client/firewall/iptables/dnat_refcount_linux_test.go +++ /dev/null @@ -1,240 +0,0 @@ -//go:build privileged - -package iptables - -import ( - "net/netip" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - fw "github.com/netbirdio/netbird/client/firewall/manager" - "github.com/netbirdio/netbird/client/iface" - "github.com/netbirdio/netbird/client/iface/wgaddr" -) - -func iptRefcountIfaceV4() *iFaceMock { - return &iFaceMock{ - NameFunc: func() string { return "wt-refcount" }, - AddressFunc: func() wgaddr.Address { - return wgaddr.Address{ - IP: netip.MustParseAddr("10.20.0.1"), - Network: netip.MustParsePrefix("10.20.0.0/24"), - } - }, - } -} - -func iptRefcountIfaceDual() *iFaceMock { - return &iFaceMock{ - NameFunc: func() string { return "wt-refcount" }, - AddressFunc: func() wgaddr.Address { - return wgaddr.Address{ - IP: netip.MustParseAddr("10.20.0.1"), - Network: netip.MustParsePrefix("10.20.0.0/24"), - IPv6: netip.MustParseAddr("fd00::1"), - IPv6Net: netip.MustParsePrefix("fd00::/64"), - } - }, - } -} - -func newIptRefcountManager(t *testing.T, dual bool) *Manager { - t.Helper() - var ifMock *iFaceMock - if dual { - ifMock = iptRefcountIfaceDual() - } else { - ifMock = iptRefcountIfaceV4() - } - m, err := Create(ifMock, iface.DefaultMTU) - require.NoError(t, err, "create manager") - require.NoError(t, m.Init(nil), "init manager") - t.Cleanup(func() { - require.NoError(t, m.Close(nil), "close manager") - }) - return m -} - -func iptDnatV4(port uint16) fw.ForwardRule { - return fw.ForwardRule{ - Protocol: fw.ProtocolTCP, - DestinationPort: fw.Port{Values: []uint16{port}}, - TranslatedAddress: netip.MustParseAddr("10.20.0.2"), - TranslatedPort: fw.Port{Values: []uint16{80}}, - } -} - -func iptDnatV6(port uint16) fw.ForwardRule { - return fw.ForwardRule{ - Protocol: fw.ProtocolTCP, - DestinationPort: fw.Port{Values: []uint16{port}}, - TranslatedAddress: netip.MustParseAddr("fd00::2"), - TranslatedPort: fw.Port{Values: []uint16{80}}, - } -} - -// TestIptablesRouting_RepeatedEnableSingleReference verifies that EnableRouting -// (called on every network-map update) holds at most one reference per family -// and a single DisableRouting drops both back to zero. -func TestIptablesRouting_RepeatedEnableSingleReference(t *testing.T) { - m := newIptRefcountManager(t, true) - state := m.family4.ipFwdState - - require.NoError(t, m.EnableRouting(), "first enable") - require.NoError(t, m.EnableRouting(), "second enable") - require.NoError(t, m.EnableRouting(), "third enable") - v4, v6 := state.Counts() - assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference") - assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference") - - require.NoError(t, m.DisableRouting(), "disable") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4, "single disable releases the v4 reference") - assert.Equal(t, 0, v6, "single disable releases the v6 reference") -} - -// TestIptablesRouting_DisableKeepsDNATReference verifies that an unpaired -// DisableRouting does not release references held by active DNAT rules. -func TestIptablesRouting_DisableKeepsDNATReference(t *testing.T) { - m := newIptRefcountManager(t, true) - state := m.family4.ipFwdState - - r1, err := m.AddDNATRule(iptDnatV6(9095)) - require.NoError(t, err, "add v6 dnat") - - require.NoError(t, m.DisableRouting(), "unpaired disable") - _, v6 := state.Counts() - assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting") - - require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat") - _, v6 = state.Counts() - assert.Equal(t, 0, v6, "delete releases the DNAT reference") -} - -// TestIptablesDNAT_RefcountBalancedV4 covers a Balanced Add/Delete pair on v4. -func TestIptablesDNAT_RefcountBalancedV4(t *testing.T) { - m := newIptRefcountManager(t, false) - state := m.family4.ipFwdState - - r1, err := m.AddDNATRule(iptDnatV4(7081)) - require.NoError(t, err, "add v4 dnat 1") - v4, v6 := state.Counts() - assert.Equal(t, 1, v4, "v4 refcount after first add") - assert.Equal(t, 0, v6, "v6 refcount unchanged") - - r2, err := m.AddDNATRule(iptDnatV4(7082)) - require.NoError(t, err, "add v4 dnat 2") - v4, v6 = state.Counts() - assert.Equal(t, 2, v4, "v4 refcount after second add") - assert.Equal(t, 0, v6, "v6 refcount unchanged") - - require.NoError(t, m.DeleteDNATRule(r1)) - v4, v6 = state.Counts() - assert.Equal(t, 1, v4, "v4 refcount after first delete") - assert.Equal(t, 0, v6, "v6 refcount unchanged") - - require.NoError(t, m.DeleteDNATRule(r2)) - v4, v6 = state.Counts() - assert.Equal(t, 0, v4, "v4 refcount after second delete") - assert.Equal(t, 0, v6, "v6 refcount unchanged") -} - -// TestIptablesDNAT_RefcountBalancedV6 checks the v6 path increments v6 only and -// decrements back to zero. -func TestIptablesDNAT_RefcountBalancedV6(t *testing.T) { - m := newIptRefcountManager(t, true) - require.NotNil(t, m.family6, "v6 family") - require.Same(t, m.family4.ipFwdState, m.family6.ipFwdState, "shared state") - state := m.family4.ipFwdState - - r1, err := m.AddDNATRule(iptDnatV6(9081)) - require.NoError(t, err, "add v6 dnat 1") - v4, v6 := state.Counts() - assert.Equal(t, 0, v4) - assert.Equal(t, 1, v6, "v6 refcount after first add") - - r2, err := m.AddDNATRule(iptDnatV6(9082)) - require.NoError(t, err, "add v6 dnat 2") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4, "v4 refcount unchanged") - assert.Equal(t, 2, v6, "v6 refcount after second add") - - require.NoError(t, m.DeleteDNATRule(r1)) - v4, v6 = state.Counts() - assert.Equal(t, 0, v4, "v4 refcount unchanged") - assert.Equal(t, 1, v6, "v6 refcount after first delete") - - require.NoError(t, m.DeleteDNATRule(r2)) - v4, v6 = state.Counts() - assert.Equal(t, 0, v4) - assert.Equal(t, 0, v6, "v6 refcount after second delete") -} - -// TestIptablesDNAT_DuplicateAddNoLeak verifies the duplicate-rule path returns -// without bumping the refcount. -func TestIptablesDNAT_DuplicateAddNoLeak(t *testing.T) { - m := newIptRefcountManager(t, true) - state := m.family4.ipFwdState - - rule := iptDnatV4(7083) - r1, err := m.AddDNATRule(rule) - require.NoError(t, err) - v4, _ := state.Counts() - assert.Equal(t, 1, v4) - - _, err = m.AddDNATRule(rule) - require.NoError(t, err, "duplicate add") - v4, _ = state.Counts() - assert.Equal(t, 1, v4, "duplicate add must not increment") - - require.NoError(t, m.DeleteDNATRule(r1)) - v4, _ = state.Counts() - assert.Equal(t, 0, v4, "single delete must drop to zero") -} - -// TestIptablesDNAT_DeleteMissingNoUnderflow verifies Delete on an unknown rule -// neither errors nor releases the refcount. -func TestIptablesDNAT_DeleteMissingNoUnderflow(t *testing.T) { - m := newIptRefcountManager(t, true) - state := m.family4.ipFwdState - - phantom := iptDnatV4(7099) - require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4") - v4, v6 := state.Counts() - assert.Equal(t, 0, v4) - assert.Equal(t, 0, v6) - - phantom6 := iptDnatV6(9099) - require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4) - assert.Equal(t, 0, v6) - - r1, err := m.AddDNATRule(iptDnatV4(7100)) - require.NoError(t, err) - v4, _ = state.Counts() - assert.Equal(t, 1, v4, "real add still increments after phantom delete") - require.NoError(t, m.DeleteDNATRule(r1)) -} - -// TestIptablesDNAT_DoubleDeleteNoUnderflow verifies a second Delete on the same -// rule is a no-op. -func TestIptablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) { - m := newIptRefcountManager(t, true) - state := m.family4.ipFwdState - - r1, err := m.AddDNATRule(iptDnatV6(9083)) - require.NoError(t, err) - _, v6 := state.Counts() - assert.Equal(t, 1, v6) - - require.NoError(t, m.DeleteDNATRule(r1), "first delete") - _, v6 = state.Counts() - assert.Equal(t, 0, v6) - - require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op") - _, v6 = state.Counts() - assert.Equal(t, 0, v6, "double delete must not underflow") -} diff --git a/client/firewall/iptables/family_linux.go b/client/firewall/iptables/family_linux.go index 0e1ce5440..2ac860a0a 100644 --- a/client/firewall/iptables/family_linux.go +++ b/client/firewall/iptables/family_linux.go @@ -56,10 +56,6 @@ const ( markManglePost = "mark-mangle-post" matchSet = "--match-set" - dnatSuffix firewall.RuleID = "_dnat" - snatSuffix firewall.RuleID = "_snat" - fwdSuffix firewall.RuleID = "_fwd" - // ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation. ipv4TCPHeaderSize = 40 // ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation. diff --git a/client/firewall/iptables/filter_linux.go b/client/firewall/iptables/filter_linux.go index dc606da2d..30cd81018 100644 --- a/client/firewall/iptables/filter_linux.go +++ b/client/firewall/iptables/filter_linux.go @@ -81,15 +81,6 @@ func (r *family) hasRule(id nbid.RuleID) bool { return ok } -// hasDNATRule reports whether this family owns the DNAT rule set for -// the given user id. DNAT rules live in r.rules under the well-known -// "_dnat" key; the lookup here is used by Manager.DeleteDNATRule -// to pick the right family. -func (r *family) hasDNATRule(id firewall.RuleID) bool { - _, ok := r.rules[id+dnatSuffix] - return ok -} - // DeleteFilterRule removes a previously installed filter rule. The // rule's stored chain/table identify where to delete from; source set // references are recovered from the spec via findSets and dropped diff --git a/client/firewall/iptables/manager_linux.go b/client/firewall/iptables/manager_linux.go index 0f0b0110e..a566909c8 100644 --- a/client/firewall/iptables/manager_linux.go +++ b/client/firewall/iptables/manager_linux.go @@ -323,31 +323,6 @@ func (m *Manager) DisableRouting() error { return m.family4.ipFwdState.ReleaseRouting() } -// AddDNATRule adds a DNAT rule -func (m *Manager) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { - m.mutex.Lock() - defer m.mutex.Unlock() - - if rule.TranslatedAddress.Is6() { - if !m.hasIPv6() { - return nil, fmt.Errorf("add DNAT rule: %w", firewall.ErrIPv6NotInitialized) - } - return m.family6.AddDNATRule(rule) - } - return m.family4.AddDNATRule(rule) -} - -// DeleteDNATRule deletes a DNAT rule -func (m *Manager) DeleteDNATRule(rule firewall.Rule) error { - m.mutex.Lock() - defer m.mutex.Unlock() - - if m.hasIPv6() && !m.family4.hasDNATRule(rule.ID()) { - return m.family6.DeleteDNATRule(rule) - } - return m.family4.DeleteDNATRule(rule) -} - // UpdateSet updates the set with the given prefixes func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error { m.mutex.Lock() diff --git a/client/firewall/iptables/manager_linux_test.go b/client/firewall/iptables/manager_linux_test.go index 9f53352e1..8435bf6a5 100644 --- a/client/firewall/iptables/manager_linux_test.go +++ b/client/firewall/iptables/manager_linux_test.go @@ -497,16 +497,6 @@ func TestIptablesCloseRemovesAllState(t *testing.T) { require.NoError(t, manager.AddNatRule(pair), "add nat rule") require.NoError(t, manager.EnableRouting(), "enable routing") - // A DNAT redirect, which also holds a forwarding reference. - dnat := fw.ForwardRule{ - Protocol: fw.ProtocolTCP, - DestinationPort: fw.Port{Values: []uint16{8080}}, - TranslatedAddress: netip.MustParseAddr("10.20.0.44"), - TranslatedPort: fw.Port{Values: []uint16{80}}, - } - _, err = manager.AddDNATRule(dnat) - require.NoError(t, err, "add dnat rule") - require.NotEqual(t, before, snapshotIptables(t, ipv4Client), "the manager must have installed state") // Everything above stays in place, so Close is what has to remove it. diff --git a/client/firewall/manager/firewall.go b/client/firewall/manager/firewall.go index 0eb376875..f8de1e2b5 100644 --- a/client/firewall/manager/firewall.go +++ b/client/firewall/manager/firewall.go @@ -172,12 +172,6 @@ type Manager interface { DisableRouting() error - // AddDNATRule adds outbound DNAT rule for forwarding external traffic to the NetBird network. - AddDNATRule(ForwardRule) (Rule, error) - - // DeleteDNATRule deletes the outbound DNAT rule. - DeleteDNATRule(Rule) error - // UpdateSet updates the set with the given prefixes UpdateSet(hash Set, prefixes []netip.Prefix) error diff --git a/client/firewall/manager/forward_rule.go b/client/firewall/manager/forward_rule.go deleted file mode 100644 index c2e9e5c60..000000000 --- a/client/firewall/manager/forward_rule.go +++ /dev/null @@ -1,27 +0,0 @@ -package manager - -import ( - "fmt" - "net/netip" -) - -// ForwardRule todo figure out better place to this to avoid circular imports -type ForwardRule struct { - Protocol Protocol - DestinationPort Port - TranslatedAddress netip.Addr - TranslatedPort Port -} - -func (r ForwardRule) ID() RuleID { - id := fmt.Sprintf("%s;%s;%s;%s", - r.Protocol, - r.DestinationPort.String(), - r.TranslatedAddress.String(), - r.TranslatedPort.String()) - return RuleID(id) -} - -func (r ForwardRule) String() string { - return fmt.Sprintf("protocol: %s, destinationPort: %s, translatedAddress: %s, translatedPort: %s", r.Protocol, r.DestinationPort.String(), r.TranslatedAddress.String(), r.TranslatedPort.String()) -} diff --git a/client/firewall/nftables/dnat_linux.go b/client/firewall/nftables/dnat_linux.go index 8eae694a2..c179d60cc 100644 --- a/client/firewall/nftables/dnat_linux.go +++ b/client/firewall/nftables/dnat_linux.go @@ -9,332 +9,11 @@ import ( "github.com/google/nftables" "github.com/google/nftables/binaryutil" "github.com/google/nftables/expr" - "github.com/google/nftables/xt" - "github.com/hashicorp/go-multierror" log "github.com/sirupsen/logrus" - nberrors "github.com/netbirdio/netbird/client/errors" firewall "github.com/netbirdio/netbird/client/firewall/manager" ) -func (r *family) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { - ruleID := rule.ID() - if _, exists := r.rules[ruleID+dnatSuffix]; exists { - return rule, nil - } - - protoNum, err := r.af.protoNum(rule.Protocol) - if err != nil { - return nil, fmt.Errorf("convert protocol to number: %w", err) - } - - // Request forwarding before queueing rules: addDnatRedirect/addDnatMasq - // buffer netlink messages on r.conn that the next caller's Flush would - // commit if we returned without flushing them ourselves. - if err := r.ipFwdState.RequestForwarding(r.isV6()); err != nil { - return nil, fmt.Errorf("enable forwarding: %w", err) - } - - if err := r.addDnatRedirect(rule, protoNum, ruleID); err != nil { - r.releaseForwarding() - return nil, err - } - - if err := r.addDnatMasq(rule, protoNum, ruleID); err != nil { - r.releaseForwarding() - delete(r.rules, ruleID+dnatSuffix) - return nil, err - } - - // Unlike iptables, there's no point in adding "out" rules in the forward chain here as our policy is ACCEPT. - // To overcome DROP policies in other chains, we'd have to add rules to the chains there. - // We also cannot just add "oif accept" there and filter in our own table as we don't know what is supposed to be allowed. - // TODO: find chains with drop policies and add rules there - - if err := r.conn.Flush(); err != nil { - r.releaseForwarding() - delete(r.rules, ruleID+dnatSuffix) - delete(r.rules, ruleID+snatSuffix) - return nil, fmt.Errorf("flush rules: %w", err) - } - - return &rule, nil -} - -func (r *family) addDnatRedirect(rule firewall.ForwardRule, protoNum uint8, ruleID firewall.RuleID) error { - dnatExprs := []expr.Any{ - &expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1}, - &expr.Cmp{ - Op: expr.CmpOpNeq, - Register: 1, - Data: ifname(r.wgIface.Name()), - }, - &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, - &expr.Cmp{ - Op: expr.CmpOpEq, - Register: 1, - Data: []byte{protoNum}, - }, - &expr.Payload{ - DestRegister: 1, - Base: expr.PayloadBaseTransportHeader, - Offset: 2, - Len: 2, - }, - } - portExprs, err := r.applyPort(&rule.DestinationPort, false) - if err != nil { - return fmt.Errorf("apply destination port: %w", err) - } - dnatExprs = append(dnatExprs, portExprs...) - - // shifted translated port is not supported in nftables, so we hand this over to xtables - if rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2 { - if rule.TranslatedPort.Values[0] != rule.DestinationPort.Values[0] || - rule.TranslatedPort.Values[1] != rule.DestinationPort.Values[1] { - return r.addXTablesRedirect(dnatExprs, ruleID, rule) - } - } - - additionalExprs, regProtoMin, regProtoMax, err := r.handleTranslatedPort(rule) - if err != nil { - return err - } - dnatExprs = append(dnatExprs, additionalExprs...) - - dnatExprs = append(dnatExprs, - &expr.NAT{ - Type: expr.NATTypeDestNAT, - Family: uint32(r.af.tableFamily), - RegAddrMin: 1, - RegProtoMin: regProtoMin, - RegProtoMax: regProtoMax, - }, - ) - - dnatRule := &nftables.Rule{ - Table: r.workTable, - Chain: r.chains[chainNameRoutingRdr], - Exprs: dnatExprs, - UserData: []byte(ruleID + dnatSuffix), - } - r.conn.AddRule(dnatRule) - r.rules[ruleID+dnatSuffix] = dnatRule - - return nil -} - -func (r *family) handleTranslatedPort(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) { - switch { - case rule.TranslatedPort.IsRange && len(rule.TranslatedPort.Values) == 2: - return r.handlePortRange(rule) - case len(rule.TranslatedPort.Values) == 0: - return r.handleAddressOnly(rule) - case len(rule.TranslatedPort.Values) == 1: - return r.handleSinglePort(rule) - default: - return nil, 0, 0, fmt.Errorf("invalid translated port: %v", rule.TranslatedPort) - } -} - -func (r *family) handlePortRange(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) { - exprs := []expr.Any{ - &expr.Immediate{ - Register: 1, - Data: rule.TranslatedAddress.AsSlice(), - }, - &expr.Immediate{ - Register: 2, - Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[0]), - }, - &expr.Immediate{ - Register: 3, - Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[1]), - }, - } - return exprs, 2, 3, nil -} - -func (r *family) handleAddressOnly(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) { - exprs := []expr.Any{ - &expr.Immediate{ - Register: 1, - Data: rule.TranslatedAddress.AsSlice(), - }, - } - return exprs, 0, 0, nil -} - -func (r *family) handleSinglePort(rule firewall.ForwardRule) ([]expr.Any, uint32, uint32, error) { - exprs := []expr.Any{ - &expr.Immediate{ - Register: 1, - Data: rule.TranslatedAddress.AsSlice(), - }, - &expr.Immediate{ - Register: 2, - Data: binaryutil.BigEndian.PutUint16(rule.TranslatedPort.Values[0]), - }, - } - return exprs, 2, 0, nil -} - -func (r *family) addXTablesRedirect(dnatExprs []expr.Any, ruleID firewall.RuleID, rule firewall.ForwardRule) error { - dnatExprs = append(dnatExprs, - &expr.Counter{}, - &expr.Target{ - Name: "DNAT", - Rev: 2, - Info: &xt.NatRange2{ - NatRange: xt.NatRange{ - Flags: uint(xt.NatRangeMapIPs | xt.NatRangeProtoSpecified | xt.NatRangeProtoOffset), - MinIP: rule.TranslatedAddress.AsSlice(), - MaxIP: rule.TranslatedAddress.AsSlice(), - MinPort: rule.TranslatedPort.Values[0], - MaxPort: rule.TranslatedPort.Values[1], - }, - BasePort: rule.DestinationPort.Values[0], - }, - }, - ) - - natTable := &nftables.Table{ - Name: tableNat, - Family: r.af.tableFamily, - } - dnatRule := &nftables.Rule{ - Table: natTable, - Chain: &nftables.Chain{ - Name: chainNameNatPrerouting, - Table: natTable, - Type: nftables.ChainTypeNAT, - Hooknum: nftables.ChainHookPrerouting, - Priority: nftables.ChainPriorityNATDest, - }, - Exprs: dnatExprs, - UserData: []byte(ruleID + dnatSuffix), - } - r.conn.AddRule(dnatRule) - r.rules[ruleID+dnatSuffix] = dnatRule - - return nil -} - -func (r *family) addDnatMasq(rule firewall.ForwardRule, protoNum uint8, ruleID firewall.RuleID) error { - portExprs, err := r.applyPort(&rule.TranslatedPort, false) - if err != nil { - return fmt.Errorf("apply translated port: %w", err) - } - - masqExprs := []expr.Any{ - &expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1}, - &expr.Cmp{ - Op: expr.CmpOpEq, - Register: 1, - Data: ifname(r.wgIface.Name()), - }, - &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, - &expr.Cmp{ - Op: expr.CmpOpEq, - Register: 1, - Data: []byte{protoNum}, - }, - &expr.Payload{ - DestRegister: 1, - Base: expr.PayloadBaseNetworkHeader, - Offset: r.af.dstAddrOffset, - Len: r.af.addrLen, - }, - &expr.Cmp{ - Op: expr.CmpOpEq, - Register: 1, - Data: rule.TranslatedAddress.AsSlice(), - }, - } - - masqExprs = append(masqExprs, portExprs...) - masqExprs = append(masqExprs, &expr.Masq{}) - - masqRule := &nftables.Rule{ - Table: r.workTable, - Chain: r.chains[chainNameRoutingNat], - Exprs: masqExprs, - UserData: []byte(ruleID + snatSuffix), - } - r.conn.AddRule(masqRule) - r.rules[ruleID+snatSuffix] = masqRule - - return nil -} - -func (r *family) DeleteDNATRule(rule firewall.Rule) error { - ruleID := rule.ID() - - if err := r.refreshRulesMap(); err != nil { - return fmt.Errorf(refreshRulesMapError, err) - } - - var merr *multierror.Error - var needsFlush bool - var found bool - - if dnatRule, exists := r.rules[ruleID+dnatSuffix]; exists { - found = true - if dnatRule.Handle == 0 { - log.Warnf("dnat rule %s has no handle, removing stale entry", ruleID+dnatSuffix) - delete(r.rules, ruleID+dnatSuffix) - } else if err := r.conn.DelRule(dnatRule); err != nil { - merr = multierror.Append(merr, fmt.Errorf("delete dnat rule: %w", err)) - } else { - needsFlush = true - } - } - - if masqRule, exists := r.rules[ruleID+snatSuffix]; exists { - found = true - if masqRule.Handle == 0 { - log.Warnf("snat rule %s has no handle, removing stale entry", ruleID+snatSuffix) - delete(r.rules, ruleID+snatSuffix) - } else if err := r.conn.DelRule(masqRule); err != nil { - merr = multierror.Append(merr, fmt.Errorf("delete snat rule: %w", err)) - } else { - needsFlush = true - } - } - - if needsFlush { - if err := r.conn.Flush(); err != nil { - merr = multierror.Append(merr, fmt.Errorf(flushError, err)) - } - } - - if merr != nil { - return nberrors.FormatErrorOrNil(merr) - } - - delete(r.rules, ruleID+dnatSuffix) - delete(r.rules, ruleID+snatSuffix) - - // Release once, only if the rule was present and removed. - if found { - r.releaseForwarding() - } - - return nil -} - -// releaseForwarding drops one IP forwarding reference, logging any error. -func (r *family) releaseForwarding() { - if err := r.ipFwdState.ReleaseForwarding(r.isV6()); err != nil { - log.Errorf("release IP forwarding: %v", err) - } -} - -// isV6 reports whether this family handles the IPv6 table. -func (r *family) isV6() bool { - return r.af.tableFamily == nftables.TableFamilyIPv6 -} - func (r *family) AddInboundDNAT(localAddr netip.Addr, protocol firewall.Protocol, originalPort, translatedPort uint16) error { ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort)) diff --git a/client/firewall/nftables/dnat_refcount_linux_test.go b/client/firewall/nftables/dnat_refcount_linux_test.go deleted file mode 100644 index cdc24e77f..000000000 --- a/client/firewall/nftables/dnat_refcount_linux_test.go +++ /dev/null @@ -1,249 +0,0 @@ -//go:build privileged - -package nftables - -import ( - "net/netip" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - fw "github.com/netbirdio/netbird/client/firewall/manager" - "github.com/netbirdio/netbird/client/iface" - "github.com/netbirdio/netbird/client/iface/wgaddr" -) - -func nftRefcountIfaceV4() *iFaceMock { - return &iFaceMock{ - NameFunc: func() string { return "wt-refcount" }, - AddressFunc: func() wgaddr.Address { - return wgaddr.Address{ - IP: netip.MustParseAddr("100.96.0.1"), - Network: netip.MustParsePrefix("100.96.0.0/16"), - } - }, - } -} - -func nftRefcountIfaceDual() *iFaceMock { - return &iFaceMock{ - NameFunc: func() string { return "wt-refcount" }, - AddressFunc: func() wgaddr.Address { - return wgaddr.Address{ - IP: netip.MustParseAddr("100.96.0.1"), - Network: netip.MustParsePrefix("100.96.0.0/16"), - IPv6: netip.MustParseAddr("fd00::1"), - IPv6Net: netip.MustParsePrefix("fd00::/64"), - } - }, - } -} - -func newNftRefcountManager(t *testing.T, dual bool) *Manager { - t.Helper() - if check() != NFTABLES { - t.Skip("nftables not supported on this system") - } - var ifMock *iFaceMock - if dual { - ifMock = nftRefcountIfaceDual() - } else { - ifMock = nftRefcountIfaceV4() - } - m, err := Create(ifMock, iface.DefaultMTU) - require.NoError(t, err, "create manager") - require.NoError(t, m.Init(nil), "init manager") - t.Cleanup(func() { - require.NoError(t, m.Close(nil), "close manager") - }) - return m -} - -func dnatV4(port uint16) fw.ForwardRule { - return fw.ForwardRule{ - Protocol: fw.ProtocolTCP, - DestinationPort: fw.Port{Values: []uint16{port}}, - TranslatedAddress: netip.MustParseAddr("100.96.0.2"), - TranslatedPort: fw.Port{Values: []uint16{80}}, - } -} - -func dnatV6(port uint16) fw.ForwardRule { - return fw.ForwardRule{ - Protocol: fw.ProtocolTCP, - DestinationPort: fw.Port{Values: []uint16{port}}, - TranslatedAddress: netip.MustParseAddr("fd00::2"), - TranslatedPort: fw.Port{Values: []uint16{80}}, - } -} - -// TestNftablesDNAT_RefcountBalancedV4 verifies that Add/Delete pairs leave the -// v4 refcount at zero. -func TestNftablesDNAT_RefcountBalancedV4(t *testing.T) { - m := newNftRefcountManager(t, false) - state := m.family4.ipFwdState - - r1, err := m.AddDNATRule(dnatV4(8081)) - require.NoError(t, err, "add v4 dnat 1") - v4, v6 := state.Counts() - assert.Equal(t, 1, v4, "v4 refcount after first add") - assert.Equal(t, 0, v6, "v6 refcount unchanged") - - r2, err := m.AddDNATRule(dnatV4(8082)) - require.NoError(t, err, "add v4 dnat 2") - v4, v6 = state.Counts() - assert.Equal(t, 2, v4, "v4 refcount after second add") - assert.Equal(t, 0, v6, "v6 refcount unchanged") - - require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat 1") - v4, v6 = state.Counts() - assert.Equal(t, 1, v4, "v4 refcount after first delete") - assert.Equal(t, 0, v6, "v6 refcount unchanged") - - require.NoError(t, m.DeleteDNATRule(r2), "delete v4 dnat 2") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4, "v4 refcount after second delete") - assert.Equal(t, 0, v6, "v6 refcount unchanged") -} - -// TestNftablesDNAT_RefcountBalancedV6 verifies the v6 path increments v6 only -// and decrements back to zero on Delete. -func TestNftablesDNAT_RefcountBalancedV6(t *testing.T) { - m := newNftRefcountManager(t, true) - require.NotNil(t, m.family6, "v6 family") - require.Same(t, m.family4.ipFwdState, m.family6.ipFwdState, "shared state") - state := m.family4.ipFwdState - - r1, err := m.AddDNATRule(dnatV6(9091)) - require.NoError(t, err, "add v6 dnat 1") - v4, v6 := state.Counts() - assert.Equal(t, 0, v4, "v4 refcount unchanged") - assert.Equal(t, 1, v6, "v6 refcount after first add") - - r2, err := m.AddDNATRule(dnatV6(9092)) - require.NoError(t, err, "add v6 dnat 2") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4) - assert.Equal(t, 2, v6, "v6 refcount after second add") - - require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat 1") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4, "v4 refcount unchanged") - assert.Equal(t, 1, v6, "v6 refcount after first delete") - - require.NoError(t, m.DeleteDNATRule(r2), "delete v6 dnat 2") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4) - assert.Equal(t, 0, v6, "v6 refcount after second delete") -} - -// TestNftablesDNAT_DuplicateAddNoLeak verifies that a duplicate Add (same -// ForwardRule) does not double-increment the refcount. -func TestNftablesDNAT_DuplicateAddNoLeak(t *testing.T) { - m := newNftRefcountManager(t, true) - state := m.family4.ipFwdState - - rule := dnatV4(8083) - r1, err := m.AddDNATRule(rule) - require.NoError(t, err, "add v4 dnat") - v4, _ := state.Counts() - assert.Equal(t, 1, v4) - - // duplicate add: same rule ID, must be a no-op for the refcount. - _, err = m.AddDNATRule(rule) - require.NoError(t, err, "duplicate add") - v4, _ = state.Counts() - assert.Equal(t, 1, v4, "duplicate add must not increment") - - require.NoError(t, m.DeleteDNATRule(r1), "delete v4 dnat") - v4, _ = state.Counts() - assert.Equal(t, 0, v4, "single delete must drop to zero") -} - -// TestNftablesDNAT_DeleteMissingNoUnderflow verifies deleting a rule that was -// never added does not underflow the refcount. -func TestNftablesDNAT_DeleteMissingNoUnderflow(t *testing.T) { - m := newNftRefcountManager(t, true) - state := m.family4.ipFwdState - - // Construct a Rule reference for something never added. The router stores - // rules by ID(), and DeleteDNATRule looks them up in r.rules; a missing - // entry must be a no-op rather than calling Release. - phantom := dnatV4(8099) - require.NoError(t, m.DeleteDNATRule(&phantom), "delete missing v4 dnat") - v4, v6 := state.Counts() - assert.Equal(t, 0, v4, "v4 refcount unaffected by missing delete") - assert.Equal(t, 0, v6, "v6 refcount unaffected") - - phantom6 := dnatV6(9099) - require.NoError(t, m.DeleteDNATRule(&phantom6), "delete missing v6 dnat") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4) - assert.Equal(t, 0, v6, "v6 refcount unaffected by missing delete") - - // And after a phantom delete, a real add still results in count=1. - r1, err := m.AddDNATRule(dnatV4(8100)) - require.NoError(t, err, "add v4 dnat after phantom delete") - v4, _ = state.Counts() - assert.Equal(t, 1, v4, "real add still increments after phantom delete") - require.NoError(t, m.DeleteDNATRule(r1)) -} - -// TestNftablesRouting_RepeatedEnableSingleReference verifies that EnableRouting -// (called on every network-map update) holds at most one reference per family -// and a single DisableRouting drops both back to zero. -func TestNftablesRouting_RepeatedEnableSingleReference(t *testing.T) { - m := newNftRefcountManager(t, true) - state := m.family4.ipFwdState - - require.NoError(t, m.EnableRouting(), "first enable") - require.NoError(t, m.EnableRouting(), "second enable") - require.NoError(t, m.EnableRouting(), "third enable") - v4, v6 := state.Counts() - assert.Equal(t, 1, v4, "repeated enable holds a single v4 reference") - assert.Equal(t, 1, v6, "repeated enable holds a single v6 reference") - - require.NoError(t, m.DisableRouting(), "disable") - v4, v6 = state.Counts() - assert.Equal(t, 0, v4, "single disable releases the v4 reference") - assert.Equal(t, 0, v6, "single disable releases the v6 reference") -} - -// TestNftablesRouting_DisableKeepsDNATReference verifies that an unpaired -// DisableRouting does not release references held by active DNAT rules. -func TestNftablesRouting_DisableKeepsDNATReference(t *testing.T) { - m := newNftRefcountManager(t, true) - state := m.family4.ipFwdState - - r1, err := m.AddDNATRule(dnatV6(9095)) - require.NoError(t, err, "add v6 dnat") - - require.NoError(t, m.DisableRouting(), "unpaired disable") - _, v6 := state.Counts() - assert.Equal(t, 1, v6, "DNAT-held reference survives unpaired DisableRouting") - - require.NoError(t, m.DeleteDNATRule(r1), "delete v6 dnat") - _, v6 = state.Counts() - assert.Equal(t, 0, v6, "delete releases the DNAT reference") -} - -// TestNftablesDNAT_DoubleDeleteNoUnderflow verifies that deleting the same rule -// twice does not underflow the refcount (the second delete is a no-op). -func TestNftablesDNAT_DoubleDeleteNoUnderflow(t *testing.T) { - m := newNftRefcountManager(t, true) - state := m.family4.ipFwdState - - r1, err := m.AddDNATRule(dnatV6(9093)) - require.NoError(t, err) - _, v6 := state.Counts() - assert.Equal(t, 1, v6) - - require.NoError(t, m.DeleteDNATRule(r1), "first delete") - _, v6 = state.Counts() - assert.Equal(t, 0, v6) - - require.NoError(t, m.DeleteDNATRule(r1), "second delete must be no-op") - _, v6 = state.Counts() - assert.Equal(t, 0, v6, "double delete must not underflow") -} diff --git a/client/firewall/nftables/family_linux.go b/client/firewall/nftables/family_linux.go index 7a5df3ed7..4169c9d2d 100644 --- a/client/firewall/nftables/family_linux.go +++ b/client/firewall/nftables/family_linux.go @@ -24,7 +24,6 @@ const ( tableRaw = "raw" tableSecurity = "security" - chainNameNatPrerouting = "PREROUTING" chainNameRoutingFw = "netbird-rt-fwd" chainNameRoutingNat = "netbird-rt-postrouting" chainNameRoutingRdr = "netbird-rt-redirect" @@ -47,9 +46,6 @@ const ( userDataAcceptForwardRuleOif = "frwacceptoif" userDataAcceptInputRule = "inputaccept" - dnatSuffix firewall.RuleID = "_dnat" - snatSuffix firewall.RuleID = "_snat" - // ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation. ipv4TCPHeaderSize = 40 // ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation. @@ -167,10 +163,6 @@ func (r *family) Reset() error { merr = multierror.Append(merr, err) } - if err := r.removeNatPreroutingRules(); err != nil { - merr = multierror.Append(merr, fmt.Errorf("remove filter prerouting rules: %w", err)) - } - return nberrors.FormatErrorOrNil(merr) } diff --git a/client/firewall/nftables/filter_linux.go b/client/firewall/nftables/filter_linux.go index ebd238063..bb3ac1dfe 100644 --- a/client/firewall/nftables/filter_linux.go +++ b/client/firewall/nftables/filter_linux.go @@ -197,11 +197,6 @@ func (r *family) hasRule(id firewall.RuleID) bool { return ok } -func (r *family) hasDNATRule(id firewall.RuleID) bool { - _, ok := r.rules[id+dnatSuffix] - return ok -} - // DeleteFilterRule removes a previously installed filter rule. Source // set references are recovered from the stored rule's expressions via // findSets and dropped from the shared refcounter. diff --git a/client/firewall/nftables/manager_linux.go b/client/firewall/nftables/manager_linux.go index 87651761f..75405e213 100644 --- a/client/firewall/nftables/manager_linux.go +++ b/client/firewall/nftables/manager_linux.go @@ -252,7 +252,7 @@ func (m *Manager) DeleteFilterRule(rule firewall.Rule) error { m.mutex.Lock() defer m.mutex.Unlock() - fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule, false) + fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule) if err != nil { return err } @@ -260,11 +260,8 @@ func (m *Manager) DeleteFilterRule(rule firewall.Rule) error { } // familyForRuleID picks the family holding the rule with the given id, using -// the supplied lookup. With refresh set, a miss in both cached maps reloads -// the NAT/DNAT rule maps from the kernel once and re-checks before falling -// back to the v4 family. Filter rules are tracked only in memory and have no -// kernel-backed reload, so their callers pass refresh as false. -func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall.RuleID) bool, refresh bool) (*family, error) { +// the supplied lookup, and falls back to the v4 family on a miss. +func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall.RuleID) bool) (*family, error) { if has(m.family4, id) { return m.family4, nil } @@ -274,18 +271,6 @@ func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall if has(m.family6, id) { return m.family6, nil } - if !refresh { - return m.family4, nil - } - if err := m.family4.refreshRulesMap(); err != nil { - return nil, fmt.Errorf("refresh v4 rules: %w", err) - } - if err := m.family6.refreshRulesMap(); err != nil { - return nil, fmt.Errorf("refresh v6 rules: %w", err) - } - if has(m.family6, id) && !has(m.family4, id) { - return m.family6, nil - } return m.family4, nil } @@ -450,32 +435,6 @@ func (m *Manager) Flush() error { return nil } -// AddDNATRule adds a DNAT rule -func (m *Manager) AddDNATRule(rule firewall.ForwardRule) (firewall.Rule, error) { - m.mutex.Lock() - defer m.mutex.Unlock() - - if rule.TranslatedAddress.Is6() { - if !m.hasIPv6() { - return nil, fmt.Errorf("add DNAT rule: %w", firewall.ErrIPv6NotInitialized) - } - return m.family6.AddDNATRule(rule) - } - return m.family4.AddDNATRule(rule) -} - -// DeleteDNATRule deletes a DNAT rule -func (m *Manager) DeleteDNATRule(rule firewall.Rule) error { - m.mutex.Lock() - defer m.mutex.Unlock() - - r, err := m.familyForRuleID(rule.ID(), (*family).hasDNATRule, true) - if err != nil { - return err - } - return r.DeleteDNATRule(rule) -} - // UpdateSet updates the set with the given prefixes func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error { m.mutex.Lock() diff --git a/client/firewall/nftables/manager_linux_test.go b/client/firewall/nftables/manager_linux_test.go index 0ca56409e..4d6eec3c1 100644 --- a/client/firewall/nftables/manager_linux_test.go +++ b/client/firewall/nftables/manager_linux_test.go @@ -378,18 +378,6 @@ func TestNftablesManagerCompatibilityWithIptables(t *testing.T) { err = manager.AddNatRule(pair) require.NoError(t, err, "failed to add NAT rule") - dnatRule, err := manager.AddDNATRule(fw.ForwardRule{ - Protocol: fw.ProtocolTCP, - DestinationPort: fw.Port{Values: []uint16{8080}}, - TranslatedAddress: netip.MustParseAddr("100.96.0.2"), - TranslatedPort: fw.Port{Values: []uint16{80}}, - }) - require.NoError(t, err, "failed to add DNAT rule") - - t.Cleanup(func() { - require.NoError(t, manager.DeleteDNATRule(dnatRule), "failed to delete DNAT rule") - }) - stdout, stderr = runIptablesSave(t) verifyIptablesOutput(t, stdout, stderr) } @@ -453,18 +441,6 @@ func TestNftablesManagerIPv6CompatibilityWithIp6tables(t *testing.T) { }) require.NoError(t, err, "add v6 NAT rule") - dnatRule, err := manager.AddDNATRule(fw.ForwardRule{ - Protocol: fw.ProtocolTCP, - DestinationPort: fw.Port{Values: []uint16{8080}}, - TranslatedAddress: netip.MustParseAddr("fd00::2"), - TranslatedPort: fw.Port{Values: []uint16{80}}, - }) - require.NoError(t, err, "add v6 DNAT rule") - - t.Cleanup(func() { - require.NoError(t, manager.DeleteDNATRule(dnatRule), "delete v6 DNAT rule") - }) - stdout, stderr := runIptablesSave(t) verifyIptablesOutput(t, stdout, stderr) diff --git a/client/firewall/nftables/routing_linux.go b/client/firewall/nftables/routing_linux.go index d619c5543..e98471e8f 100644 --- a/client/firewall/nftables/routing_linux.go +++ b/client/firewall/nftables/routing_linux.go @@ -459,41 +459,6 @@ func (r *family) RemoveAllLegacyRouteRules() error { return nberrors.FormatErrorOrNil(merr) } -func (r *family) removeNatPreroutingRules() error { - table := &nftables.Table{ - Name: tableNat, - Family: r.af.tableFamily, - } - chain := &nftables.Chain{ - Name: chainNameNatPrerouting, - Table: table, - Hooknum: nftables.ChainHookPrerouting, - Priority: nftables.ChainPriorityNATDest, - Type: nftables.ChainTypeNAT, - } - rules, err := r.conn.GetRules(table, chain) - if err != nil { - return fmt.Errorf("get rules from nat table: %w", err) - } - - var merr *multierror.Error - - // Delete rules that have our UserData suffix - for _, rule := range rules { - if len(rule.UserData) == 0 || !strings.HasSuffix(string(rule.UserData), string(dnatSuffix)) { - continue - } - if err := r.conn.DelRule(rule); err != nil { - merr = multierror.Append(merr, fmt.Errorf("delete rule %s: %w", rule.UserData, err)) - } - } - - if err := r.conn.Flush(); err != nil { - merr = multierror.Append(merr, fmt.Errorf(flushError, err)) - } - return nberrors.FormatErrorOrNil(merr) -} - func (r *family) RemoveNatRule(pair firewall.RouterPair) error { if err := r.refreshRulesMap(); err != nil { return fmt.Errorf(refreshRulesMapError, err) diff --git a/client/firewall/uspfilter/nat.go b/client/firewall/uspfilter/nat.go index 06312aabf..49c26766a 100644 --- a/client/firewall/uspfilter/nat.go +++ b/client/firewall/uspfilter/nat.go @@ -486,16 +486,6 @@ func incrementalUpdate(oldChecksum uint16, oldBytes, newBytes []byte) uint16 { return ^uint16(sum) } -// AddDNATRule adds outbound DNAT rule for forwarding external traffic to NetBird network. -func (m *Manager) AddDNATRule(firewall.ForwardRule) (firewall.Rule, error) { - return nil, errNotSupported -} - -// DeleteDNATRule deletes outbound DNAT rule. -func (m *Manager) DeleteDNATRule(firewall.Rule) error { - return errNotSupported -} - // addPortRedirection adds a port redirection rule. func (m *Manager) addPortRedirection(targetIP netip.Addr, protocol gopacket.LayerType, originalPort, translatedPort uint16) error { m.portDNATMutex.Lock() diff --git a/client/internal/auth/sessionwatch/watcher.go b/client/internal/auth/sessionwatch/watcher.go index e685c28d0..496903044 100644 --- a/client/internal/auth/sessionwatch/watcher.go +++ b/client/internal/auth/sessionwatch/watcher.go @@ -90,8 +90,9 @@ type StatusRecorder interface { // fallback T-FinalWarningLead dialog (suppressed when the user dismissed // the first one for the same deadline). Safe for concurrent use. type Watcher struct { - lead time.Duration - finalLead time.Duration + lead time.Duration + finalLead time.Duration + deadlineOnly bool mu sync.Mutex current time.Time @@ -102,6 +103,7 @@ type Watcher struct { dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal closed bool recorder StatusRecorder + nowFn func() time.Time } // New returns a watcher with the package defaults WarningLead and @@ -122,9 +124,17 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher { lead: lead, finalLead: final, recorder: recorder, + nowFn: time.Now, } } +// NewDeadlineOnly returns a watcher that validates and records deadlines but arms no warning timers. +func NewDeadlineOnly(recorder StatusRecorder) *Watcher { + w := New(recorder) + w.deadlineOnly = true + return w +} + // Update sets the latest deadline. Pass the zero time to clear (e.g. when // a Sync push from the server omits the field because login expiration // was disabled). @@ -181,7 +191,7 @@ func (w *Watcher) Update(deadline time.Time) error { w.finalFiredAt = time.Time{} w.dismissedAt = time.Time{} - if deadline.After(now) { + if deadline.After(now) && !w.deadlineOnly { w.armTimerLocked(deadline) } recorder := w.recorder @@ -303,6 +313,11 @@ func (w *Watcher) fire(armedFor time.Time) { w.mu.Unlock() return } + now := w.nowFn() + if isLate(now, armedFor, max(w.finalLead, 0)) { + w.fireLateLocked(armedFor, now) + return + } w.firedAt = armedFor recorder := w.recorder w.mu.Unlock() @@ -331,6 +346,14 @@ func (w *Watcher) fireFinal(armedFor time.Time) { log.Infof("auth session final-warning skipped (dismissed by user)") return } + now := w.nowFn() + if isLate(now, armedFor, 0) { + w.finalFiredAt = armedFor + w.mu.Unlock() + log.Infof("auth session final-warning skipped for deadline %s (passed %s ago)", + armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second)) + return + } w.finalFiredAt = armedFor recorder := w.recorder w.mu.Unlock() @@ -341,6 +364,39 @@ func (w *Watcher) fireFinal(armedFor time.Time) { publishWarning(recorder, armedFor, true) } +// fireLateLocked handles a T-WarningLead callback that fired inside the +// final-warning window: it sends the final warning in its place while the +// deadline has not passed and the user has not dismissed it, so a resume +// with time left still warns. The caller must hold w.mu; this helper +// releases it. +func (w *Watcher) fireLateLocked(armedFor, now time.Time) { + w.firedAt = armedFor + switch { + case w.dismissedAt.Equal(armedFor): + w.mu.Unlock() + log.Infof("auth session expiry soon warning skipped (dismissed by user)") + return + case w.finalFiredAt.Equal(armedFor): + w.mu.Unlock() + log.Infof("auth session expiry soon warning skipped (final warning already fired)") + return + case isLate(now, armedFor, 0): + w.mu.Unlock() + log.Infof("auth session expiry soon warning skipped for deadline %s (passed %s ago)", + armedFor.Format(time.RFC3339), now.Round(0).Sub(armedFor).Round(time.Second)) + return + } + w.finalFiredAt = armedFor + recorder := w.recorder + w.mu.Unlock() + if recorder == nil { + return + } + log.Infof("auth session expiry soon warning fired inside the final-warning window, sending final warning for deadline %s", + armedFor.Format(time.RFC3339)) + publishWarning(recorder, armedFor, true) +} + // armOneShotLocked schedules cb at fireAt. When fireAt is already in the // past it dispatches on the next scheduler tick so a state-change recorder // notification (invoked after w.mu is released) lands first. Caller must @@ -380,3 +436,11 @@ func publishWarning(recorder StatusRecorder, deadline time.Time, final bool) { meta, ) } + +// isLate reports whether the wall clock now has already reached armedFor +// minus cutoffLead. The timers run on the monotonic clock, which can stall +// while the host sleeps, so a timer can fire long after the window it was +// armed for. +func isLate(now, armedFor time.Time, cutoffLead time.Duration) bool { + return !now.Round(0).Before(armedFor.Add(-cutoffLead).Round(0)) +} diff --git a/client/internal/auth/sessionwatch/watcher_test.go b/client/internal/auth/sessionwatch/watcher_test.go index 4b49a94b6..cb2800978 100644 --- a/client/internal/auth/sessionwatch/watcher_test.go +++ b/client/internal/auth/sessionwatch/watcher_test.go @@ -527,3 +527,201 @@ func TestDismissBeforeUpdateIsNoop(t *testing.T) { } t.Fatalf("final-warning did not publish after no-op pre-Update Dismiss, events=%+v", r.snapshot()) } + +func TestIsLate(t *testing.T) { + armedFor := time.Date(2026, 10, 1, 12, 0, 0, 0, time.UTC) + lead := 2 * time.Minute + tests := []struct { + name string + now time.Time + cutoffLead time.Duration + want bool + }{ + {"before cutoff", armedFor.Add(-3 * time.Minute), lead, false}, + {"at cutoff", armedFor.Add(-lead), lead, true}, + {"after cutoff", armedFor.Add(-time.Minute), lead, true}, + {"zero lead before deadline", armedFor.Add(-time.Second), 0, false}, + {"zero lead at deadline", armedFor, 0, true}, + {"zero lead after deadline", armedFor.Add(time.Second), 0, true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isLate(tt.now, armedFor, tt.cutoffLead); got != tt.want { + t.Fatalf("isLate(%s, %s, %s) = %v, want %v", tt.now, armedFor, tt.cutoffLead, got, tt.want) + } + }) + } +} + +func TestIsLateIgnoresMonotonicReading(t *testing.T) { + now := time.Now() + wallOnly := now.Round(0) + if isLate(now, wallOnly.Add(time.Second), 0) { + t.Fatalf("now with monotonic reading must compare as wall clock before a later wall-only deadline") + } + if !isLate(now, wallOnly, 0) { + t.Fatalf("now with monotonic reading must compare as wall clock at an equal wall-only deadline") + } +} + +func TestLateTimerFiring(t *testing.T) { + tests := []struct { + name string + final bool + beforeDl time.Duration + wantWarns int + wantFinals int + }{ + {"warning on resume inside window", false, 3 * time.Minute, 1, 0}, + {"warning promoted to final inside final window", false, time.Minute, 0, 1}, + {"warning skipped past deadline", false, -time.Minute, 0, 0}, + {"final on resume before deadline", true, time.Minute, 0, 1}, + {"final skipped past deadline", true, -time.Minute, 0, 0}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + r := &fakeRecorder{} + w := New(r) + defer w.Close() + + // The deadline is an hour out so the real timers never fire + // during the test; the late callback is invoked directly with an + // injected clock that simulates a resume near the deadline. + d := time.Now().Add(time.Hour).Round(0) + w.nowFn = func() time.Time { return d.Add(-tt.beforeDl) } + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + if tt.final { + w.fireFinal(d) + } else { + w.fire(d) + } + + events := r.snapshot() + if got := countWhere(events, event.isWarning); got != tt.wantWarns { + t.Fatalf("expected %d warning publishes, got %d: %+v", tt.wantWarns, got, events) + } + if got := countWhere(events, event.isFinalWarning); got != tt.wantFinals { + t.Fatalf("expected %d final-warning publishes, got %d: %+v", tt.wantFinals, got, events) + } + }) + } +} + +func TestPromotedFinalWarningIsNotRepeated(t *testing.T) { + r := &fakeRecorder{} + w := New(r) + defer w.Close() + + d := time.Now().Add(time.Hour).Round(0) + now := d.Add(-time.Minute) + w.nowFn = func() time.Time { return now } + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + w.fire(d) + // The final timer was suspended too, so it fires even later than the + // warning timer, here still just before the deadline. + now = d.Add(-30 * time.Second) + w.fireFinal(d) + + events := r.snapshot() + if got := countWhere(events, event.isFinalWarning); got != 1 { + t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events) + } + if got := countWhere(events, event.isWarning); got != 0 { + t.Fatalf("expected no regular warning publish, got %d: %+v", got, events) + } +} + +func TestPromotionRespectsDismiss(t *testing.T) { + r := &fakeRecorder{} + w := New(r) + defer w.Close() + + d := time.Now().Add(time.Hour).Round(0) + w.nowFn = func() time.Time { return d.Add(-time.Minute) } + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + w.Dismiss() + w.fire(d) + + events := r.snapshot() + if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 { + t.Fatalf("expected no publish after dismiss, got %d: %+v", got, events) + } +} + +func TestPromotionSkippedWhenFinalAlreadyFired(t *testing.T) { + r := &fakeRecorder{} + w := New(r) + defer w.Close() + + // Both timers fall in the past after a long suspend and are dispatched + // with a zero delay, so the final callback can run before the warning one. + d := time.Now().Add(time.Hour).Round(0) + w.nowFn = func() time.Time { return d.Add(-time.Minute) } + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + + w.fireFinal(d) + w.fire(d) + + events := r.snapshot() + if got := countWhere(events, event.isFinalWarning); got != 1 { + t.Fatalf("expected exactly 1 final-warning publish, got %d: %+v", got, events) + } + if got := countWhere(events, event.isWarning); got != 0 { + t.Fatalf("expected no regular warning publish, got %d: %+v", got, events) + } +} + +func TestDeadlineOnlyRecordsDeadlineWithoutWarnings(t *testing.T) { + r := &fakeRecorder{} + w := NewDeadlineOnly(r) + defer w.Close() + + // With the default leads this deadline would otherwise fire both + // timers on the next tick. + d := time.Now().Add(50 * time.Millisecond).Round(0) + if err := w.Update(d); err != nil { + t.Fatalf("Update: %v", err) + } + if got := r.deadline(); !got.Equal(d) { + t.Fatalf("expected recorder deadline %v, got %v", d, got) + } + + time.Sleep(100 * time.Millisecond) + + events := r.snapshot() + if got := countWhere(events, func(e event) bool { return e.kind == publish }); got != 0 { + t.Fatalf("expected no publish in deadline-only mode, got %d: %+v", got, events) + } + if w.timer != nil || w.finalTimer != nil { + t.Fatal("expected no timers armed in deadline-only mode") + } +} + +func TestDeadlineOnlyStillRejectsOutOfRangeDeadlines(t *testing.T) { + r := &fakeRecorder{} + w := NewDeadlineOnly(r) + defer w.Close() + + if err := w.Update(time.Now().Add(time.Hour)); err != nil { + t.Fatalf("Update: %v", err) + } + + err := w.Update(time.Now().Add(-maxPastHorizon - time.Hour)) + if !errors.Is(err, ErrDeadlineInPast) { + t.Fatalf("expected ErrDeadlineInPast, got %v", err) + } + if got := r.deadline(); !got.IsZero() { + t.Fatalf("expected recorder cleared after rejection, got %v", got) + } +} diff --git a/client/internal/debug/debug_test.go b/client/internal/debug/debug_test.go index 6a810bccc..0f74490f2 100644 --- a/client/internal/debug/debug_test.go +++ b/client/internal/debug/debug_test.go @@ -846,6 +846,7 @@ func TestAddConfig_AllFieldsCovered(t *testing.T) { "ClientCertKeyPair": "non-config: parsed cert pair, not serialized", "Name": "non-config: profile name is not needed for debug purposes", "policy": "non-config: in-memory MDM policy snapshot, surfaced via Config.Policy() / GetConfigResponse.MDMManagedFields", + "probing": "non-config: marks a throwaway copy built to be diffed against; never set on a config anyone runs with", "DebugBundleUploadURL": "sensitive: MDM-provided upload URL may carry credentials or query tokens; kept out of the shared bundle", } diff --git a/client/internal/engine.go b/client/internal/engine.go index 7e9375771..4d731cbd7 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -42,7 +42,6 @@ import ( dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config" "github.com/netbirdio/netbird/client/internal/dnsfwd" "github.com/netbirdio/netbird/client/internal/expose" - "github.com/netbirdio/netbird/client/internal/ingressgw" "github.com/netbirdio/netbird/client/internal/lazyconn" "github.com/netbirdio/netbird/client/internal/metrics" "github.com/netbirdio/netbird/client/internal/netflow" @@ -262,11 +261,10 @@ type Engine struct { statusRecorder *peer.Status - firewall firewallManager.Manager - routeManager routemanager.Manager - acl acl.Manager - dnsForwardMgr *dnsfwd.Manager - ingressGatewayMgr *ingressgw.Manager + firewall firewallManager.Manager + routeManager routemanager.Manager + acl acl.Manager + dnsForwardMgr *dnsfwd.Manager dnsServer dns.Server @@ -448,13 +446,6 @@ func (e *Engine) stopLocked() { e.cleanupSSHConfig() - if e.ingressGatewayMgr != nil { - if err := e.ingressGatewayMgr.Close(); err != nil { - log.Warnf("failed to cleanup forward rules: %v", err) - } - e.ingressGatewayMgr = nil - } - if e.srWatcher != nil { e.srWatcher.Close() } @@ -1627,13 +1618,6 @@ func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap) error { e.updateDNSForwarder(dnsRouteFeatureFlag, fwdEntries) done() - // Ingress forward rules - done = e.phase("forward_rules") - if _, err := e.updateForwardRules(networkMap.GetForwardingRules()); err != nil { - log.Errorf("failed to update forward rules, err: %v", err) - } - done() - log.Debugf("got peers update from Management Service, total peers to connect to = %d", len(networkMap.GetRemotePeers())) done = e.phase("offline_peers") @@ -2733,74 +2717,6 @@ func (e *Engine) setForwarderCapture(pc device.PacketCapture) { } } -func (e *Engine) updateForwardRules(rules []*mgmProto.ForwardingRule) ([]firewallManager.ForwardRule, error) { - if e.firewall == nil { - log.Warn("firewall is disabled, not updating forwarding rules") - return nil, nil - } - - if len(rules) == 0 { - if e.ingressGatewayMgr == nil { - return nil, nil - } - - err := e.ingressGatewayMgr.Close() - e.ingressGatewayMgr = nil - e.statusRecorder.SetIngressGwMgr(nil) - return nil, err - } - - if e.ingressGatewayMgr == nil { - mgr := ingressgw.NewManager(e.firewall) - e.ingressGatewayMgr = mgr - e.statusRecorder.SetIngressGwMgr(mgr) - } - - var merr *multierror.Error - forwardingRules := make([]firewallManager.ForwardRule, 0, len(rules)) - for _, rule := range rules { - proto, err := acl.ConvertToFirewallProtocol(rule.GetProtocol()) - if err != nil { - merr = multierror.Append(merr, fmt.Errorf("failed to convert protocol '%s': %w", rule.GetProtocol(), err)) - continue - } - - dstPortInfo, err := convertPortInfo(rule.GetDestinationPort()) - if err != nil { - merr = multierror.Append(merr, fmt.Errorf("invalid destination port '%v': %w", rule.GetDestinationPort(), err)) - continue - } - - translateIP, err := convertToIP(rule.GetTranslatedAddress()) - if err != nil { - merr = multierror.Append(merr, fmt.Errorf("failed to convert translated address '%s': %w", rule.GetTranslatedAddress(), err)) - continue - } - - translatePort, err := convertPortInfo(rule.GetTranslatedPort()) - if err != nil { - merr = multierror.Append(merr, fmt.Errorf("invalid translate port '%v': %w", rule.GetTranslatedPort(), err)) - continue - } - - forwardRule := firewallManager.ForwardRule{ - Protocol: proto, - DestinationPort: *dstPortInfo, - TranslatedAddress: translateIP, - TranslatedPort: *translatePort, - } - - forwardingRules = append(forwardingRules, forwardRule) - } - - log.Infof("updating forwarding rules: %d", len(forwardingRules)) - if err := e.ingressGatewayMgr.Update(forwardingRules); err != nil { - log.Errorf("failed to update forwarding rules: %v", err) - } - - return forwardingRules, nberrors.FormatErrorOrNil(merr) -} - // toExcludedLazyPeers returns the peers that must have an always-active // connection: those that are not lazy by policy (the per-peer lazy state or the // account flag, subject to the local override). diff --git a/client/internal/engine_privileged_test.go b/client/internal/engine_privileged_test.go index 2db0cd5ed..4449b5788 100644 --- a/client/internal/engine_privileged_test.go +++ b/client/internal/engine_privileged_test.go @@ -42,7 +42,6 @@ import ( nbcache "github.com/netbirdio/netbird/management/server/cache" "github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator" - "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" @@ -523,8 +522,8 @@ func startManagement(t *testing.T, dataDir, testFile string) (*grpc.Server, stri updateManager := update_channel.NewPeersUpdateManager(metrics) requestBuffer := server.NewAccountRequestBuffer(context.Background(), store) - networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config, nil) - accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore) + networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", manager.NewEphemeralManager(store, peersManager), config, nil) + accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, settingsMockManager, permissionsManager, false, cacheStore) if err != nil { return nil, "", err } diff --git a/client/internal/engine_sessionwatch.go b/client/internal/engine_sessionwatch.go index a46d73f87..05b46a465 100644 --- a/client/internal/engine_sessionwatch.go +++ b/client/internal/engine_sessionwatch.go @@ -1,4 +1,4 @@ -//go:build !js +//go:build !js && !android package internal @@ -7,10 +7,12 @@ import ( "github.com/netbirdio/netbird/client/internal/peer" ) -// newSessionWatcher returns the real SSO session expiry watcher for every -// non-wasm build. The js/wasm build gets a no-op stub from -// engine_sessionwatch_js.go so the sessionwatch package (and its timer -// machinery) never links into the wasm binary. +// newSessionWatcher returns the real SSO session expiry watcher. The js/wasm +// build gets a no-op stub from engine_sessionwatch_js.go so the sessionwatch +// package (and its timer machinery) never links into the wasm binary; the +// android build gets a deadline-only watcher from +// engine_sessionwatch_android.go because the app schedules the warnings +// itself. func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher { return sessionwatch.New(recorder) } diff --git a/client/internal/engine_sessionwatch_android.go b/client/internal/engine_sessionwatch_android.go new file mode 100644 index 000000000..8317f9165 --- /dev/null +++ b/client/internal/engine_sessionwatch_android.go @@ -0,0 +1,12 @@ +//go:build android + +package internal + +import ( + "github.com/netbirdio/netbird/client/internal/auth/sessionwatch" + "github.com/netbirdio/netbird/client/internal/peer" +) + +func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher { + return sessionwatch.NewDeadlineOnly(recorder) +} diff --git a/client/internal/ingressgw/manager.go b/client/internal/ingressgw/manager.go deleted file mode 100644 index 605543d1c..000000000 --- a/client/internal/ingressgw/manager.go +++ /dev/null @@ -1,111 +0,0 @@ -package ingressgw - -import ( - "fmt" - "sync" - - "github.com/hashicorp/go-multierror" - log "github.com/sirupsen/logrus" - - nberrors "github.com/netbirdio/netbird/client/errors" - firewall "github.com/netbirdio/netbird/client/firewall/manager" -) - -type DNATFirewall interface { - AddDNATRule(fwdRule firewall.ForwardRule) (firewall.Rule, error) - DeleteDNATRule(rule firewall.Rule) error -} - -type RulePair struct { - firewall.ForwardRule - firewall.Rule -} - -type Manager struct { - dnatFirewall DNATFirewall - - rules map[firewall.RuleID]RulePair - rulesMu sync.Mutex -} - -func NewManager(dnatFirewall DNATFirewall) *Manager { - return &Manager{ - dnatFirewall: dnatFirewall, - rules: make(map[firewall.RuleID]RulePair), - } -} - -func (h *Manager) Update(forwardRules []firewall.ForwardRule) error { - h.rulesMu.Lock() - defer h.rulesMu.Unlock() - - var mErr *multierror.Error - - toDelete := make(map[firewall.RuleID]RulePair, len(h.rules)) - for id, r := range h.rules { - toDelete[id] = r - } - - // Process new/updated rules - for _, fwdRule := range forwardRules { - id := fwdRule.ID() - if _, ok := h.rules[id]; ok { - delete(toDelete, id) - continue - } - - rule, err := h.dnatFirewall.AddDNATRule(fwdRule) - if err != nil { - mErr = multierror.Append(mErr, fmt.Errorf("add forward rule '%s': %v", fwdRule.String(), err)) - continue - } - if rule == nil { - mErr = multierror.Append(mErr, fmt.Errorf("add forward rule '%s': backend returned no rule", fwdRule.String())) - continue - } - log.Infof("forward rule has been added '%s'", fwdRule) - h.rules[id] = RulePair{ - ForwardRule: fwdRule, - Rule: rule, - } - } - - // Remove deleted rules - for id, rulePair := range toDelete { - if err := h.dnatFirewall.DeleteDNATRule(rulePair.Rule); err != nil { - mErr = multierror.Append(mErr, fmt.Errorf("failed to delete forward rule '%s': %v", rulePair.ForwardRule.String(), err)) - } - log.Infof("forward rule has been deleted '%s'", rulePair.ForwardRule) - delete(h.rules, id) - } - - return nberrors.FormatErrorOrNil(mErr) -} - -func (h *Manager) Close() error { - h.rulesMu.Lock() - defer h.rulesMu.Unlock() - - log.Infof("clean up all (%d) forward rules", len(h.rules)) - var mErr *multierror.Error - for _, rule := range h.rules { - if err := h.dnatFirewall.DeleteDNATRule(rule.Rule); err != nil { - mErr = multierror.Append(mErr, fmt.Errorf("failed to delete forward rule '%s': %v", rule, err)) - } - } - - h.rules = make(map[firewall.RuleID]RulePair) - return nberrors.FormatErrorOrNil(mErr) -} - -func (h *Manager) Rules() []firewall.ForwardRule { - h.rulesMu.Lock() - defer h.rulesMu.Unlock() - - rules := make([]firewall.ForwardRule, 0, len(h.rules)) - for _, rulePair := range h.rules { - rules = append(rules, rulePair.ForwardRule) - } - - return rules -} diff --git a/client/internal/ingressgw/manager_test.go b/client/internal/ingressgw/manager_test.go deleted file mode 100644 index 0cd40fcc4..000000000 --- a/client/internal/ingressgw/manager_test.go +++ /dev/null @@ -1,281 +0,0 @@ -package ingressgw - -import ( - "fmt" - "net/netip" - "testing" - - firewall "github.com/netbirdio/netbird/client/firewall/manager" -) - -var ( - _ firewall.Rule = (*MocFwRule)(nil) - _ DNATFirewall = &MockDNATFirewall{} -) - -type MocFwRule struct { - id firewall.RuleID -} - -func (m *MocFwRule) ID() firewall.RuleID { - return m.id -} - -type MockDNATFirewall struct { - throwError bool -} - -func (m *MockDNATFirewall) AddDNATRule(fwdRule firewall.ForwardRule) (firewall.Rule, error) { - if m.throwError { - return nil, fmt.Errorf("moc error") - } - - fwRule := &MocFwRule{ - id: fwdRule.ID(), - } - return fwRule, nil -} - -func (m *MockDNATFirewall) DeleteDNATRule(rule firewall.Rule) error { - if m.throwError { - return fmt.Errorf("moc error") - } - return nil -} - -func (m *MockDNATFirewall) forceToThrowErrors() { - m.throwError = true -} - -func TestManager_AddRule(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - - updates := []firewall.ForwardRule{ - { - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - }, - { - Protocol: firewall.ProtocolUDP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - }} - - if err := mgr.Update(updates); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != len(updates) { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} - -func TestManager_UpdateRule(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - ruleUDP := firewall.ForwardRule{ - Protocol: firewall.ProtocolUDP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.2"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleUDP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 1 { - t.Errorf("unexpected rules count: %d", len(rules)) - } - - if rules[0].TranslatedAddress.String() != ruleUDP.TranslatedAddress.String() { - t.Errorf("unexpected rule: %v", rules[0]) - } - - if rules[0].TranslatedPort.String() != ruleUDP.TranslatedPort.String() { - t.Errorf("unexpected rule: %v", rules[0]) - } - - if rules[0].DestinationPort.String() != ruleUDP.DestinationPort.String() { - t.Errorf("unexpected rule: %v", rules[0]) - } - - if rules[0].Protocol != ruleUDP.Protocol { - t.Errorf("unexpected rule: %v", rules[0]) - } -} - -func TestManager_ExtendRules(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - ruleUDP := firewall.ForwardRule{ - Protocol: firewall.ProtocolUDP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.2"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP, ruleUDP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 2 { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} - -func TestManager_UnderlingError(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - ruleUDP := firewall.ForwardRule{ - Protocol: firewall.ProtocolUDP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.2"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - fw.forceToThrowErrors() - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP, ruleUDP}); err == nil { - t.Errorf("expected error") - } - - rules := mgr.Rules() - if len(rules) != 1 { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} - -func TestManager_Cleanup(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - if err := mgr.Update([]firewall.ForwardRule{}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 0 { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} - -func TestManager_DeleteBrokenRule(t *testing.T) { - fw := &MockDNATFirewall{} - - // force to throw errors when Add DNAT Rule - fw.forceToThrowErrors() - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err == nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 0 { - t.Errorf("unexpected rules count: %d", len(rules)) - } - - // simulate that to remove a broken rule - if err := mgr.Update([]firewall.ForwardRule{}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - if err := mgr.Close(); err != nil { - t.Errorf("unexpected error: %v", err) - } -} - -func TestManager_Close(t *testing.T) { - fw := &MockDNATFirewall{} - mgr := NewManager(fw) - - port, _ := firewall.NewPort(8080) - ruleTCP := firewall.ForwardRule{ - Protocol: firewall.ProtocolTCP, - DestinationPort: *port, - TranslatedAddress: netip.MustParseAddr("172.16.254.1"), - TranslatedPort: *port, - } - - if err := mgr.Update([]firewall.ForwardRule{ruleTCP}); err != nil { - t.Errorf("unexpected error: %v", err) - } - - if err := mgr.Close(); err != nil { - t.Errorf("unexpected error: %v", err) - } - - rules := mgr.Rules() - if len(rules) != 0 { - t.Errorf("unexpected rules count: %d", len(rules)) - } -} diff --git a/client/internal/message_convert.go b/client/internal/message_convert.go deleted file mode 100644 index 60f19e228..000000000 --- a/client/internal/message_convert.go +++ /dev/null @@ -1,43 +0,0 @@ -package internal - -import ( - "errors" - "fmt" - "net" - "net/netip" - - firewallManager "github.com/netbirdio/netbird/client/firewall/manager" - mgmProto "github.com/netbirdio/netbird/shared/management/proto" -) - -func convertPortInfo(portInfo *mgmProto.PortInfo) (*firewallManager.Port, error) { - if portInfo == nil { - return nil, errors.New("portInfo cannot be nil") - } - - if portInfo.GetPort() != 0 { - return firewallManager.NewPort(int(portInfo.GetPort())) - } - - if portInfo.GetRange() != nil { - return firewallManager.NewPort(int(portInfo.GetRange().Start), int(portInfo.GetRange().End)) - } - - return nil, fmt.Errorf("invalid portInfo: %v", portInfo) -} - -func convertToIP(rawIP []byte) (netip.Addr, error) { - if rawIP == nil { - return netip.Addr{}, errors.New("input bytes cannot be nil") - } - - if len(rawIP) != net.IPv4len && len(rawIP) != net.IPv6len { - return netip.Addr{}, fmt.Errorf("invalid IP length: %d", len(rawIP)) - } - - if len(rawIP) == net.IPv4len { - return netip.AddrFrom4([4]byte(rawIP)), nil - } - - return netip.AddrFrom16([16]byte(rawIP)), nil -} diff --git a/client/internal/peer/conn.go b/client/internal/peer/conn.go index d73144773..17823e043 100644 --- a/client/internal/peer/conn.go +++ b/client/internal/peer/conn.go @@ -828,7 +828,8 @@ func (conn *Conn) evalStatus() ConnStatus { // // The result is a tri-state: // - ConnStatusConnected: all available transports are up -// - ConnStatusPartiallyConnected: relay is up but ICE is still pending/reconnecting +// - ConnStatusPartiallyConnected: one transport carries the traffic and the other does +// not: relay up with ICE down, or ICE up with the shared relay transport down // - ConnStatusDisconnected: no working transport func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { defer func() { @@ -845,13 +846,14 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { } return evalConnStatus(connStatusInputs{ - forceRelay: IsForceRelayed(), - peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(), - relayConnected: conn.statusRelay.Get() == worker.StatusConnected, - remoteSupportsICE: conn.handshaker.RemoteICESupported(), - iceWorkerCreated: iceWorkerCreated, - iceStatusConnecting: conn.statusICE.Get() != worker.StatusDisconnected, - iceInProgress: iceInProgress, + forceRelay: IsForceRelayed(), + peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(), + relayConnected: conn.statusRelay.Get() == worker.StatusConnected, + relayTransportConnected: conn.workerRelay.IsTransportConnected(), + remoteSupportsICE: conn.handshaker.RemoteICESupported(), + iceWorkerCreated: iceWorkerCreated, + iceStatusConnected: conn.statusICE.Get() == worker.StatusConnected, + iceInProgress: iceInProgress, }) } @@ -1060,19 +1062,21 @@ func evalConnStatus(in connStatusInputs) guard.ConnStatus { return boolToConnStatus(relayUsedAndUp) } - // ICE counts as "up" when the status is anything other than Disconnected, OR - // when a negotiation is currently in progress (so we don't spam offers while one is in flight). - iceUp := in.iceStatusConnecting || in.iceInProgress + // ICE counts as "running" when either connected or attempting to connect. + iceRunning := in.iceStatusConnected || in.iceInProgress // Relay side is acceptable if the peer doesn't rely on relay, or relay is connected. relayOK := !in.peerUsesRelay || in.relayConnected switch { - case iceUp && relayOK: + case iceRunning && relayOK: return guard.ConnStatusConnected case relayUsedAndUp: // Relay is up but ICE is down — partially connected. return guard.ConnStatusPartiallyConnected + case in.iceStatusConnected && !in.relayTransportConnected: + // ICE is up and the shared relay transport is down — offers cannot restore it. + return guard.ConnStatusPartiallyConnected default: return guard.ConnStatusDisconnected } diff --git a/client/internal/peer/conn_status.go b/client/internal/peer/conn_status.go index d6ad37b70..acf271534 100644 --- a/client/internal/peer/conn_status.go +++ b/client/internal/peer/conn_status.go @@ -17,13 +17,14 @@ const ( // tri-state connection classification. Extracted so the decision logic can be unit-tested // without constructing full Worker/Handshaker objects. type connStatusInputs struct { - forceRelay bool // NB_FORCE_RELAY or JS/WASM - peerUsesRelay bool // remote peer advertises relay support AND local has relay - relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay) - remoteSupportsICE bool // remote peer sent ICE credentials - iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode) - iceStatusConnecting bool // statusICE is anything other than Disconnected - iceInProgress bool // a negotiation is currently in flight + forceRelay bool // NB_FORCE_RELAY or JS/WASM + peerUsesRelay bool // remote peer advertises relay support AND local has relay + relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay) + relayTransportConnected bool // the relay transport shared by all peers on that server is up + remoteSupportsICE bool // remote peer sent ICE credentials + iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode) + iceStatusConnected bool // statusICE reports Connected + iceInProgress bool // a negotiation is currently in flight } // ConnStatus describe the status of a peer's connection diff --git a/client/internal/peer/conn_status_eval_test.go b/client/internal/peer/conn_status_eval_test.go index 66393cafe..a239196dc 100644 --- a/client/internal/peer/conn_status_eval_test.go +++ b/client/internal/peer/conn_status_eval_test.go @@ -30,6 +30,21 @@ func TestEvalConnStatus_ForceRelay(t *testing.T) { }, want: guard.ConnStatusDisconnected, }, + { + name: "force relay, relay up but the shared transport reports down", + in: connStatusInputs{ + forceRelay: true, + peerUsesRelay: true, + relayConnected: true, + relayTransportConnected: false, + // The ICE inputs are set so that the force-relay return is the only branch + // that can produce Connected here: without it the peer would fall through to + // relayUsedAndUp and report PartiallyConnected. + remoteSupportsICE: true, + iceWorkerCreated: true, + }, + want: guard.ConnStatusConnected, + }, { name: "force relay, peer does NOT use relay - disconnected forever", in: connStatusInputs{ @@ -123,24 +138,28 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = true - in.iceStatusConnecting = true + in.relayTransportConnected = true + in.iceStatusConnected = true }, want: guard.ConnStatusConnected, }, { - name: "ICE connected, peer does NOT use relay", + name: "ICE connected, peer does NOT use relay, shared transport down", mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = false - in.iceStatusConnecting = true + in.relayTransportConnected = false + in.iceStatusConnected = true }, + // A peer that does not rely on relay is unaffected by the shared transport: + // relayOK is true, so the first arm matches before the transport is considered. want: guard.ConnStatusConnected, }, { name: "ICE InProgress only, peer does NOT use relay", mutator: func(in *connStatusInputs) { in.peerUsesRelay = false - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = true }, want: guard.ConnStatusConnected, @@ -150,7 +169,8 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = true - in.iceStatusConnecting = false + in.relayTransportConnected = true + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusPartiallyConnected, @@ -160,21 +180,60 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = false - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusDisconnected, }, { - name: "ICE up, peer uses relay but relay down -> partial (relay required, ICE ignored)", + name: "ICE connected, relay down for this peer but the shared transport is up -> disconnected", mutator: func(in *connStatusInputs) { in.peerUsesRelay = true in.relayConnected = false - in.iceStatusConnecting = true + in.relayTransportConnected = true + in.iceStatusConnected = true + }, + // The transport is fine, so the peer itself is unreachable over relay: it may have + // moved to another server, and only an offer carries its new relay address. + want: guard.ConnStatusDisconnected, + }, + { + name: "ICE connected, the shared relay transport is down -> partial", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = true + }, + // ICE carries the traffic and the relay transport is restored by the relay client's + // own guard, not by offers, so this must not trigger the aggressive retry. + want: guard.ConnStatusPartiallyConnected, + }, + { + name: "ICE only negotiating while the shared relay transport is down -> disconnected", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = false + in.iceInProgress = true + }, + // A negotiation in flight is not a working transport, so this peer has no path at + // all and must keep the aggressive retry. Calling it partially connected spends the + // ICE retry budget and parks the guard on the hourly ticker, and nothing wakes it + // when the negotiation then fails: onICEStateDisconnected is only reached once ICE + // has reached Connected (worker_ice.go onConnectionStateChange). + want: guard.ConnStatusDisconnected, + }, + { + name: "ICE down and the shared relay transport is down -> disconnected", + mutator: func(in *connStatusInputs) { + in.peerUsesRelay = true + in.relayConnected = false + in.relayTransportConnected = false + in.iceStatusConnected = false + in.iceInProgress = false }, - // relayOK = false (peer uses relay but it's down), iceUp = true - // first switch arm fails (relayOK false), relayUsedAndUp = false (relay down), - // falls into default: Disconnected. want: guard.ConnStatusDisconnected, }, { @@ -182,7 +241,7 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) { mutator: func(in *connStatusInputs) { in.peerUsesRelay = false in.relayConnected = true // not actually used since peer doesn't rely on it - in.iceStatusConnecting = false + in.iceStatusConnected = false in.iceInProgress = false }, want: guard.ConnStatusDisconnected, diff --git a/client/internal/peer/guard/guard.go b/client/internal/peer/guard/guard.go index 73bab2a89..15028d91c 100644 --- a/client/internal/peer/guard/guard.go +++ b/client/internal/peer/guard/guard.go @@ -14,7 +14,8 @@ type ConnStatus int const ( // ConnStatusDisconnected means neither ICE nor Relay is connected. ConnStatusDisconnected ConnStatus = iota - // ConnStatusPartiallyConnected means Relay is connected but ICE is not. + // ConnStatusPartiallyConnected means one transport is usable and the other is not: + // relay connected with ICE down, or ICE connected with the shared relay transport down. ConnStatusPartiallyConnected // ConnStatusConnected means all required connections are established. ConnStatusConnected @@ -87,8 +88,9 @@ func (g *Guard) SetICEConnDisconnected() { // - Connected: no action, the peer is fully reachable. // - Disconnected (neither ICE nor Relay): retries aggressively with exponential backoff (800ms doubling // up to timeout), never gives up. This ensures rapid recovery when the peer has no connectivity at all. -// - PartiallyConnected (Relay up, ICE not): retries up to 3 times with exponential backoff, then switches -// to one attempt per hour. This limits signaling traffic when relay already provides connectivity. +// - PartiallyConnected (one transport usable, the other not): retries up to 3 times +// with exponential backoff, then switches to one attempt per hour. This limits +// signaling traffic while the peer still has a working path. // // External events (relay/ICE disconnect, signal/relay reconnect, candidate changes) reset the retry // counter and backoff ticker, giving ICE a fresh chance after network conditions change. diff --git a/client/internal/peer/notifier.go b/client/internal/peer/notifier.go index 1ee1d32ea..564098bd4 100644 --- a/client/internal/peer/notifier.go +++ b/client/internal/peer/notifier.go @@ -12,6 +12,8 @@ type notifier struct { serverStateLock sync.Mutex listenersLock sync.Mutex listener Listener + peerListWake chan struct{} + peerListStop chan struct{} currentClientState bool lastNotification ClientState lastNumberOfPeers int @@ -62,7 +64,6 @@ func (n *notifier) setNetworkAvailable(available bool) { func (n *notifier) setListener(listener Listener) { n.serverStateLock.Lock() lastNotification := n.effectiveState(n.lastNotification) - numOfPeers := n.lastNumberOfPeers fqdnAddress := n.lastFqdnAddress address := n.lastIPAddress n.serverStateLock.Unlock() @@ -70,17 +71,19 @@ func (n *notifier) setListener(listener Listener) { n.listenersLock.Lock() defer n.listenersLock.Unlock() + n.stopPeerListDelivererLocked() n.listener = listener listener.OnAddressChanged(fqdnAddress, address) notifyListener(listener, lastNotification) - // run on go routine to avoid on Java layer to call go functions on same thread - go listener.OnPeersListChanged(numOfPeers) + n.startPeerListDelivererLocked(listener) + n.wakePeerListDelivererLocked() } func (n *notifier) removeListener() { n.listenersLock.Lock() defer n.listenersLock.Unlock() + n.stopPeerListDelivererLocked() n.listener = nil } @@ -178,15 +181,56 @@ func (n *notifier) peerListChanged(numOfPeers int) { n.serverStateLock.Unlock() n.listenersLock.Lock() - listener := n.listener - n.listenersLock.Unlock() + defer n.listenersLock.Unlock() + n.wakePeerListDelivererLocked() +} - if listener == nil { +func (n *notifier) startPeerListDelivererLocked(listener Listener) { + wake := make(chan struct{}, 1) + stop := make(chan struct{}) + n.peerListWake = wake + n.peerListStop = stop + go n.deliverPeerListChanges(listener, wake, stop) +} + +func (n *notifier) stopPeerListDelivererLocked() { + if n.peerListStop == nil { return } + close(n.peerListStop) + n.peerListStop = nil + n.peerListWake = nil +} - // run on go routine to avoid on Java layer to call go functions on same thread - go listener.OnPeersListChanged(numOfPeers) +func (n *notifier) wakePeerListDelivererLocked() { + if n.peerListWake == nil { + return + } + select { + case n.peerListWake <- struct{}{}: + default: + } +} + +func (n *notifier) deliverPeerListChanges(listener Listener, wake <-chan struct{}, stop <-chan struct{}) { + for { + select { + case <-stop: + return + case <-wake: + } + select { + case <-stop: + return + default: + } + + n.serverStateLock.Lock() + numOfPeers := n.lastNumberOfPeers + n.serverStateLock.Unlock() + + listener.OnPeersListChanged(numOfPeers) + } } func (n *notifier) localAddressChanged(fqdn, address string) { diff --git a/client/internal/peer/notifier_test.go b/client/internal/peer/notifier_test.go index a73016b05..f81866214 100644 --- a/client/internal/peer/notifier_test.go +++ b/client/internal/peer/notifier_test.go @@ -2,7 +2,9 @@ package peer import ( "sync" + "sync/atomic" "testing" + "time" ) type mocListener struct { @@ -115,3 +117,156 @@ func Test_notifier_RemoveListener(t *testing.T) { t.Errorf("invalid state: %d", listener.peers) } } + +type coalescingListener struct { + final int + calls atomic.Int32 + inFlight atomic.Int32 + maxInFlight atomic.Int32 + last atomic.Int32 + done chan struct{} + entered chan struct{} + release chan struct{} + once sync.Once +} + +func (l *coalescingListener) OnStateChanged(ClientState) {} +func (l *coalescingListener) OnConnected() {} +func (l *coalescingListener) OnDisconnected() {} +func (l *coalescingListener) OnConnecting() {} +func (l *coalescingListener) OnDisconnecting() {} +func (l *coalescingListener) OnAddressChanged(string, string) {} + +func (l *coalescingListener) OnPeersListChanged(size int) { + current := l.inFlight.Add(1) + for { + seen := l.maxInFlight.Load() + if current <= seen || l.maxInFlight.CompareAndSwap(seen, current) { + break + } + } + if l.calls.Add(1) == 1 && l.entered != nil { + close(l.entered) + } + if l.release != nil { + <-l.release + } + time.Sleep(time.Millisecond) + l.last.Store(int32(size)) + l.inFlight.Add(-1) + if size == l.final { + l.once.Do(func() { close(l.done) }) + } +} + +func Test_notifier_PeerListChangedCoalesces(t *testing.T) { + const events = 1000 + listener := &coalescingListener{final: events, done: make(chan struct{})} + n := newNotifier() + n.setListener(listener) + + for i := 1; i <= events; i++ { + n.peerListChanged(i) + } + + select { + case <-listener.done: + case <-time.After(5 * time.Second): + t.Fatalf("last peer count not delivered, last seen: %d", listener.last.Load()) + } + + if got := listener.maxInFlight.Load(); got != 1 { + t.Errorf("concurrent deliveries: %d, expected 1", got) + } + if got := listener.calls.Load(); got >= events { + t.Errorf("deliveries not coalesced: %d calls for %d events", got, events) + } +} + +func Test_notifier_SetListenerStopsPreviousDeliverer(t *testing.T) { + old := &coalescingListener{final: -1} + replacement := &coalescingListener{final: 7, done: make(chan struct{})} + n := newNotifier() + n.setListener(old) + oldStop := n.peerListStop + + n.peerListChanged(7) + n.setListener(replacement) + + select { + case <-oldStop: + default: + t.Fatal("old deliverer not stopped on listener replacement") + } + waitFor(t, replacement.done, "replacement listener not notified") +} + +func Test_notifier_RemoveListenerStopsDeliverer(t *testing.T) { + n := newNotifier() + n.setListener(&coalescingListener{final: -1}) + stop := n.peerListStop + + n.removeListener() + + select { + case <-stop: + default: + t.Fatal("deliverer not stopped on listener removal") + } +} + +func Test_notifier_DelivererExitsAfterInFlightCallback(t *testing.T) { + listener := &coalescingListener{ + final: -1, + entered: make(chan struct{}), + release: make(chan struct{}), + } + n := newNotifier() + wake := make(chan struct{}, 1) + stop := make(chan struct{}) + exited := make(chan struct{}) + go func() { + n.deliverPeerListChanges(listener, wake, stop) + close(exited) + }() + + wake <- struct{}{} + waitFor(t, listener.entered, "listener not called") + + n.peerListChanged(7) + wake <- struct{}{} + close(stop) + close(listener.release) + + waitFor(t, exited, "deliverer did not exit after stop") + if got := listener.calls.Load(); got != 1 { + t.Errorf("deliverer ran %d callbacks after stop, expected only the in-flight one", got) + } + if got := listener.last.Load(); got == 7 { + t.Errorf("deliverer delivered the peer count queued after stop") + } +} + +func Test_notifier_DelivererPrefersStopOverPendingWake(t *testing.T) { + listener := &coalescingListener{final: -1} + n := newNotifier() + wake := make(chan struct{}, 1) + stop := make(chan struct{}) + + wake <- struct{}{} + close(stop) + n.deliverPeerListChanges(listener, wake, stop) + + if got := listener.calls.Load(); got != 0 { + t.Errorf("deliverer ran %d callbacks with stop closed, expected 0", got) + } +} + +func waitFor(t *testing.T, ch <-chan struct{}, msg string) { + t.Helper() + select { + case <-ch: + case <-time.After(5 * time.Second): + t.Fatal(msg) + } +} diff --git a/client/internal/peer/status.go b/client/internal/peer/status.go index 826bf6fe0..6c44178e1 100644 --- a/client/internal/peer/status.go +++ b/client/internal/peer/status.go @@ -18,9 +18,7 @@ import ( "google.golang.org/protobuf/types/known/durationpb" "google.golang.org/protobuf/types/known/timestamppb" - firewall "github.com/netbirdio/netbird/client/firewall/manager" "github.com/netbirdio/netbird/client/iface/configurer" - "github.com/netbirdio/netbird/client/internal/ingressgw" "github.com/netbirdio/netbird/client/internal/relay" "github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/route" @@ -161,7 +159,6 @@ type FullStatus struct { RosenpassState RosenpassState Relays []relay.ProbeResult NSGroupStates []NSGroupState - NumOfForwardingRules int LazyConnectionEnabled bool Events []*proto.SystemEvent } @@ -247,8 +244,6 @@ type Status struct { // read it without taking mux. networksRevision atomic.Uint64 - ingressGwMgr *ingressgw.Manager - routeIDLookup routeIDLookup wgIface WGIfaceStatus } @@ -276,12 +271,6 @@ func (d *Status) SetRelayMgr(manager *relayClient.Manager) { d.relayMgr = manager } -func (d *Status) SetIngressGwMgr(ingressGwMgr *ingressgw.Manager) { - d.mux.Lock() - defer d.mux.Unlock() - d.ingressGwMgr = ingressGwMgr -} - // ReplaceOfflinePeers replaces func (d *Status) ReplaceOfflinePeers(replacement []State) { d.mux.Lock() @@ -332,18 +321,6 @@ func (d *Status) GetPeer(peerPubKey string) (State, error) { return state, nil } -func (d *Status) PeerByIP(ip string) (string, bool) { - d.mux.RLock() - defer d.mux.RUnlock() - - for _, state := range d.peers { - if state.IP == ip { - return state.FQDN, true - } - } - return "", false -} - // PeerStateByIP returns the full peer State for the given tunnel IP. // Matches against either the IPv4 (State.IP) or IPv6 (State.IPv6) tunnel // address so dual-stack peers are reachable on either family. Only @@ -1163,16 +1140,6 @@ func (d *Status) GetRelayStates() []relay.ProbeResult { return relayStates } -func (d *Status) ForwardingRules() []firewall.ForwardRule { - d.mux.RLock() - defer d.mux.RUnlock() - if d.ingressGwMgr == nil { - return nil - } - - return d.ingressGwMgr.Rules() -} - func (d *Status) GetDNSStates() []NSGroupState { d.mux.RLock() defer d.mux.RUnlock() @@ -1207,7 +1174,6 @@ func (d *Status) GetFullStatus() FullStatus { Relays: d.GetRelayStates(), RosenpassState: d.GetRosenpassState(), NSGroupStates: d.GetDNSStates(), - NumOfForwardingRules: len(d.ForwardingRules()), LazyConnectionEnabled: d.GetLazyConnection(), } @@ -1579,7 +1545,6 @@ func (fs FullStatus) ToProto() *proto.FullStatus { pbFullStatus.LocalPeerState.WgPort = int32(fs.LocalPeerState.WgPort) pbFullStatus.LocalPeerState.RosenpassPermissive = fs.RosenpassState.Permissive pbFullStatus.LocalPeerState.RosenpassEnabled = fs.RosenpassState.Enabled - pbFullStatus.NumberOfForwardingRules = int32(fs.NumOfForwardingRules) pbFullStatus.LazyConnectionEnabled = fs.LazyConnectionEnabled pbFullStatus.LocalPeerState.Networks = maps.Keys(fs.LocalPeerState.Routes) diff --git a/client/internal/peer/worker_relay.go b/client/internal/peer/worker_relay.go index fc3489992..694207847 100644 --- a/client/internal/peer/worker_relay.go +++ b/client/internal/peer/worker_relay.go @@ -101,6 +101,10 @@ func (w *WorkerRelay) RelayIsSupportedLocally() bool { return w.relayManager.HasRelayAddress() } +func (w *WorkerRelay) IsTransportConnected() bool { + return w.relayManager.Ready() +} + func (w *WorkerRelay) CloseConn() { w.relayLock.Lock() conn := w.relayedConn diff --git a/client/internal/profilemanager/config.go b/client/internal/profilemanager/config.go index 412f81b5c..ac1b90a62 100644 --- a/client/internal/profilemanager/config.go +++ b/client/internal/profilemanager/config.go @@ -10,7 +10,6 @@ import ( "os" "os/user" "path/filepath" - "reflect" "runtime" "slices" "strings" @@ -198,6 +197,11 @@ type Config struct { MTU uint16 + // probing marks a config that exists only to be compared against and then + // thrown away, so apply() can skip the work that feeds no verdict. + // Unexported, so it never reaches the JSON. + probing bool + // policy is the MDM policy that produced the currently-set values // for any MDM-enforced fields. Set by ApplyMDMPolicy on every // invocation. Never persisted to disk. Callers query enforcement @@ -300,9 +304,11 @@ func fileExists(path string) (bool, error) { return false, err } -// createNewConfig creates a new config generating a new Wireguard key and saving to file -func createNewConfig(input ConfigInput) (*Config, error) { - config := &Config{ +// newConfigSkeleton returns the field values a brand-new profile config starts +// from, before apply() fills in the rest. Shared with the dry-run baseline so +// the two cannot disagree about what "a new config" means. +func newConfigSkeleton() *Config { + return &Config{ // defaults to false only for new (post 0.26) configurations ServerSSHAllowed: util.False(), // Remote jobs are an explicit opt-in and default off, including for @@ -310,6 +316,91 @@ func createNewConfig(input ConfigInput) (*Config, error) { RemoteJobsAllowed: util.False(), WgPort: iface.DefaultWgPort, } +} + +// resolveUnsetDefaults is the single place where an optional field that carries +// no value gets one, and the only place that states what each of those defaults +// is. apply() runs it before it compares anything, and that ordering is the +// point: with the values named, every comparison below it diffs values instead +// of presence. +// +// Presence-based comparison is what broke `netbird up` for a client configured +// through the environment. These fields mean "the effective default" when they +// hold nothing — every consumer already reads a nil as the value resolved here, +// the SSH toggles in engine_ssh.go and the network monitor in +// createEngineConfig — so naming them changes nothing about what runs. But +// while they stayed nil, an input restating the default read as a change, and +// since the CLI sends every flag whose value came from an environment variable +// on each `netbird up`, a client with NB_ENABLE_SSH_ROOT=false restated it +// every time and the update-settings gate refused it. +// +// Filling a field in is not a settings change, so a caller measuring change +// must not read the returned bool as one: see WouldChange, which runs a pass +// for this and discards its verdict. +// +// ServerSSHAllowed is the one field whose default depends on the config's age. +// A brand-new profile gets false from newConfigSkeleton, which runs before +// this, so what is resolved here is only the legacy case: a config written by a +// version that had no such field keeps SSH on, for backwards compatibility. +func (config *Config) resolveUnsetDefaults() (updated bool) { + // Fields that default to false on every platform. + for _, field := range []**bool{ + &config.EnableSSHRoot, + &config.EnableSSHSFTP, + &config.EnableSSHLocalPortForwarding, + &config.EnableSSHRemotePortForwarding, + &config.DisableSSHAuth, + // Remote jobs are an explicit opt-in: unlike SSH, a pre-existing config + // with no value defaults to disabled rather than being turned on. + &config.RemoteJobsAllowed, + } { + if *field == nil { + *field = util.False() + updated = true + } + } + + if config.DisableNotifications == nil { + log.Infof("setting notifications to disabled by default") + config.DisableNotifications = util.True() + updated = true + } + + if config.SSHJWTCacheTTL == nil { + // A zero TTL disables the JWT cache, which is what no value meant. + config.SSHJWTCacheTTL = new(int) + updated = true + } + + if config.NetworkMonitor == nil { + // network monitoring is on by default on windows and darwin clients + enabled := runtime.GOOS == "windows" || runtime.GOOS == "darwin" + config.NetworkMonitor = &enabled + updated = true + } + + if config.ServerSSHAllowed == nil { + if runtime.GOOS == "android" { + // default to disabled SSH on Android for security + log.Infof("setting SSH server to false by default on Android") + config.ServerSSHAllowed = util.False() + } else { + // enables SSH for configs from old versions to preserve backwards compatibility + log.Infof("falling back to enabled SSH server for pre-existing configuration") + config.ServerSSHAllowed = util.True() + } + updated = true + } + + return updated +} + +// createNewConfig resolves a new config in memory, with no identity: whoever +// needs the peer's keys calls EnsureIdentity and persists the result, so a read +// that lands on a missing file cannot hand back a config carrying keys that +// nothing will ever write down. +func createNewConfig(input ConfigInput) (*Config, error) { + config := newConfigSkeleton() if _, err := config.apply(input); err != nil { return nil, err @@ -318,6 +409,52 @@ func createNewConfig(input ConfigInput) (*Config, error) { return config, nil } +// createProvisionedConfig is createNewConfig plus the peer's identity, for the +// callers that go on to persist the config or to connect with it. +func createProvisionedConfig(input ConfigInput) (*Config, error) { + config, err := createNewConfig(input) + if err != nil { + return nil, err + } + + if _, err := config.EnsureIdentity(); err != nil { + return nil, err + } + + return config, nil +} + +// EnsureIdentity generates the keys that identify this peer if the config does +// not carry them yet, reporting whether it had to generate any. +// +// It is deliberately not part of apply(). Everything apply() fills in is a +// default it can recompute on the next read, but a generated key is not: it +// has to be persisted, or the peer comes back with a different WireGuard +// identity and re-registers. Having apply() generate keys is what forced every +// read of a config to write it back — so identity provisioning is its own step +// now, and the callers that perform it write the result out explicitly. +func (config *Config) EnsureIdentity() (bool, error) { + generated := false + + if config.PrivateKey == "" { + log.Infof("generated new Wireguard key") + config.PrivateKey = generateKey() + generated = true + } + + if config.SSHKey == "" { + log.Infof("generated new SSH key") + pem, err := ssh.GeneratePrivateKey(ssh.ED25519) + if err != nil { + return generated, err + } + config.SSHKey = string(pem) + generated = true + } + + return generated, nil +} + func (config *Config) apply(input ConfigInput) (updated bool, err error) { if config.Name != "" { sanitized, err := sanitizeDisplayName(config.Name) @@ -329,6 +466,13 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } } + + // Every optional field gets its value here, before anything below compares + // one. See resolveUnsetDefaults for why that ordering is the point. + if config.resolveUnsetDefaults() { + updated = true + } + if config.ManagementURL == nil { log.Infof("using default Management URL %s", DefaultManagementURL) config.ManagementURL, err = parseURL("Management URL", DefaultManagementURL) @@ -336,20 +480,21 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { return false, err } } - if input.ManagementURL != "" && input.ManagementURL != config.ManagementURL.String() { - log.Infof("new Management URL provided, updated to %#v (old value %#v)", - input.ManagementURL, config.ManagementURL.String()) + // The comparison is on the endpoint the URL addresses, not on its + // spelling: the same endpoint can be written several ways (an implicit + // :443, a trailing slash, a different host case), and treating an + // equivalent URL as new would rewrite the config and report a settings + // change where the configuration does not actually change. + if input.ManagementURL != "" { URL, err := parseURL("Management URL", input.ManagementURL) if err != nil { return false, err } - config.ManagementURL = URL - updated = true - } else if config.ManagementURL == nil { - log.Infof("using default Management URL %s", DefaultManagementURL) - config.ManagementURL, err = parseURL("Management URL", DefaultManagementURL) - if err != nil { - return false, err + if !SameServiceURL(URL, config.ManagementURL) { + log.Infof("new Management URL provided, updated to %#v (old value %#v)", + URL.String(), config.ManagementURL.String()) + config.ManagementURL = URL + updated = true } } @@ -360,31 +505,20 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { return false, err } } - if input.AdminURL != "" && input.AdminURL != config.AdminURL.String() { - log.Infof("new Admin Panel URL provided, updated to %#v (old value %#v)", - input.AdminURL, config.AdminURL.String()) + // The admin panel is opened, not dialed, so unlike the Management URL its + // path is part of what identifies it: a panel served under /netbird is not + // the one served at the root. + if input.AdminURL != "" { newURL, err := parseURL("Admin Panel URL", input.AdminURL) if err != nil { return updated, err } - config.AdminURL = newURL - updated = true - } - - if config.PrivateKey == "" { - log.Infof("generated new Wireguard key") - config.PrivateKey = generateKey() - updated = true - } - - if config.SSHKey == "" { - log.Infof("generated new SSH key") - pem, err := ssh.GeneratePrivateKey(ssh.ED25519) - if err != nil { - return false, err + if !SameServiceURLIncludingPath(newURL, config.AdminURL) { + log.Infof("new Admin Panel URL provided, updated to %#v (old value %#v)", + newURL.String(), config.AdminURL.String()) + config.AdminURL = newURL + updated = true } - config.SSHKey = string(pem) - updated = true } if input.WireguardPort != nil && *input.WireguardPort != config.WgPort { @@ -405,7 +539,14 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.NATExternalIPs != nil && !reflect.DeepEqual(config.NATExternalIPs, input.NATExternalIPs) { + // slices.Equal, not reflect.DeepEqual, and for the same reason the DNS + // labels below use it: DeepEqual calls a nil slice and an empty one + // different, while both mean "no NAT mappings". A profile stores the + // absent list as JSON null and reads it back nil, and `netbird up` sends + // CleanNATExternalIPs — an empty list — whenever NB_EXTERNAL_IP_MAP is set + // to nothing, so the two met on every start and the gate read a no-op as a + // settings change. + if input.NATExternalIPs != nil && !slices.Equal(config.NATExternalIPs, input.NATExternalIPs) { log.Infof("updating NAT External IP [ %s ] (old value: [ %s ])", strings.Join(input.NATExternalIPs, " "), strings.Join(config.NATExternalIPs, " ")) @@ -443,21 +584,12 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.NetworkMonitor != nil && (config.NetworkMonitor == nil || *input.NetworkMonitor != *config.NetworkMonitor) { + if input.NetworkMonitor != nil && *input.NetworkMonitor != *config.NetworkMonitor { log.Infof("switching Network Monitor to %t", *input.NetworkMonitor) config.NetworkMonitor = input.NetworkMonitor updated = true } - if config.NetworkMonitor == nil { - // enable network monitoring by default on windows and darwin clients - if runtime.GOOS == "windows" || runtime.GOOS == "darwin" { - enabled := true - config.NetworkMonitor = &enabled - updated = true - } - } - if input.CustomDNSAddress != nil && string(input.CustomDNSAddress) != config.CustomDNSAddress { log.Infof("updating custom DNS address %#v (old value %#v)", string(input.CustomDNSAddress), config.CustomDNSAddress) @@ -490,7 +622,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.ServerSSHAllowed != nil && (config.ServerSSHAllowed == nil || *input.ServerSSHAllowed != *config.ServerSSHAllowed) { + if input.ServerSSHAllowed != nil && *input.ServerSSHAllowed != *config.ServerSSHAllowed { if *input.ServerSSHAllowed { log.Infof("enabling SSH server") } else { @@ -498,20 +630,9 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { } config.ServerSSHAllowed = input.ServerSSHAllowed updated = true - } else if config.ServerSSHAllowed == nil { - if runtime.GOOS == "android" { - // default to disabled SSH on Android for security - log.Infof("setting SSH server to false by default on Android") - config.ServerSSHAllowed = util.False() - } else { - // enables SSH for configs from old versions to preserve backwards compatibility - log.Infof("falling back to enabled SSH server for pre-existing configuration") - config.ServerSSHAllowed = util.True() - } - updated = true } - if input.RemoteJobsAllowed != nil && (config.RemoteJobsAllowed == nil || *input.RemoteJobsAllowed != *config.RemoteJobsAllowed) { + if input.RemoteJobsAllowed != nil && *input.RemoteJobsAllowed != *config.RemoteJobsAllowed { if *input.RemoteJobsAllowed { log.Infof("enabling remote jobs") } else { @@ -519,14 +640,9 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { } config.RemoteJobsAllowed = input.RemoteJobsAllowed updated = true - } else if config.RemoteJobsAllowed == nil { - // Remote jobs are an explicit opt-in: unlike SSH, a pre-existing config - // with no value defaults to disabled rather than being turned on. - config.RemoteJobsAllowed = util.False() - updated = true } - if input.EnableSSHRoot != nil && (config.EnableSSHRoot == nil || *input.EnableSSHRoot != *config.EnableSSHRoot) { + if input.EnableSSHRoot != nil && *input.EnableSSHRoot != *config.EnableSSHRoot { if *input.EnableSSHRoot { log.Infof("enabling SSH root login") } else { @@ -536,7 +652,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.EnableSSHSFTP != nil && (config.EnableSSHSFTP == nil || *input.EnableSSHSFTP != *config.EnableSSHSFTP) { + if input.EnableSSHSFTP != nil && *input.EnableSSHSFTP != *config.EnableSSHSFTP { if *input.EnableSSHSFTP { log.Infof("enabling SSH SFTP subsystem") } else { @@ -546,7 +662,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.EnableSSHLocalPortForwarding != nil && (config.EnableSSHLocalPortForwarding == nil || *input.EnableSSHLocalPortForwarding != *config.EnableSSHLocalPortForwarding) { + if input.EnableSSHLocalPortForwarding != nil && *input.EnableSSHLocalPortForwarding != *config.EnableSSHLocalPortForwarding { if *input.EnableSSHLocalPortForwarding { log.Infof("enabling SSH local port forwarding") } else { @@ -556,7 +672,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.EnableSSHRemotePortForwarding != nil && (config.EnableSSHRemotePortForwarding == nil || *input.EnableSSHRemotePortForwarding != *config.EnableSSHRemotePortForwarding) { + if input.EnableSSHRemotePortForwarding != nil && *input.EnableSSHRemotePortForwarding != *config.EnableSSHRemotePortForwarding { if *input.EnableSSHRemotePortForwarding { log.Infof("enabling SSH remote port forwarding") } else { @@ -566,7 +682,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.DisableSSHAuth != nil && (config.DisableSSHAuth == nil || *input.DisableSSHAuth != *config.DisableSSHAuth) { + if input.DisableSSHAuth != nil && *input.DisableSSHAuth != *config.DisableSSHAuth { if *input.DisableSSHAuth { log.Infof("disabling SSH authentication") } else { @@ -576,7 +692,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.SSHJWTCacheTTL != nil && (config.SSHJWTCacheTTL == nil || *input.SSHJWTCacheTTL != *config.SSHJWTCacheTTL) { + if input.SSHJWTCacheTTL != nil && *input.SSHJWTCacheTTL != *config.SSHJWTCacheTTL { log.Infof("updating SSH JWT cache TTL to %d seconds", *input.SSHJWTCacheTTL) config.SSHJWTCacheTTL = input.SSHJWTCacheTTL updated = true @@ -659,13 +775,16 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if input.SyncMessageVersion != nil && *input.SyncMessageVersion != *config.SyncMessageVersion { + // Assigning the pointer, not writing through it: a config that carries no + // version yet would otherwise be a nil dereference, and a panic inside a + // request handler is not a way to fail. + if input.SyncMessageVersion != nil && (config.SyncMessageVersion == nil || *input.SyncMessageVersion != *config.SyncMessageVersion) { log.Infof("setting SyncMessageVersion to %v", *input.SyncMessageVersion) - *config.SyncMessageVersion = *input.SyncMessageVersion + config.SyncMessageVersion = input.SyncMessageVersion updated = true } - if input.DisableNotifications != nil && (config.DisableNotifications == nil || *input.DisableNotifications != *config.DisableNotifications) { + if input.DisableNotifications != nil && *input.DisableNotifications != *config.DisableNotifications { if *input.DisableNotifications { log.Infof("disabling notifications") } else { @@ -675,24 +794,24 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) { updated = true } - if config.DisableNotifications == nil { - disabled := true - config.DisableNotifications = &disabled - log.Infof("setting notifications to disabled by default") - updated = true - } - - if input.ClientCertKeyPath != "" { + // Compared, not just assigned: restating the path a config already holds + // changes nothing, and reporting it as an update makes a caller that + // re-sends its own configuration look like one asking to change it. + if input.ClientCertKeyPath != "" && input.ClientCertKeyPath != config.ClientCertKeyPath { config.ClientCertKeyPath = input.ClientCertKeyPath updated = true } - if input.ClientCertPath != "" { + if input.ClientCertPath != "" && input.ClientCertPath != config.ClientCertPath { config.ClientCertPath = input.ClientCertPath updated = true } - if config.ClientCertPath != "" && config.ClientCertKeyPath != "" { + // Not on a probe: the loaded pair feeds the connection, never the + // comparison, and this would otherwise run on every gated SetConfig and + // Login — twice per request — including those that are refused or change + // nothing, logging an error per request when the files are missing. + if !config.probing && config.ClientCertPath != "" && config.ClientCertKeyPath != "" { cert, err := tls.LoadX509KeyPair(config.ClientCertPath, config.ClientCertKeyPath) if err != nil { log.Error("Failed to load mTLS cert/key pair: ", err) @@ -886,6 +1005,49 @@ func ParseServiceURL(serviceName, serviceURL string) (*url.URL, error) { return parseURL(serviceName, serviceURL) } +// SameServiceURL reports whether two service URLs address the same endpoint: +// same scheme, same host compared case-insensitively as DNS names are, and +// same effective port, where an absent port means the scheme's default. +// +// This is the one comparison every caller deciding "did this URL change?" must +// use. A string comparison answers a different question: "https://host", +// "https://host/" and "https://HOST:443" are one endpoint written three ways, +// and reading them as three values makes a client that restates its own +// management URL look like a client asking to be repointed. A nil operand +// matches only another nil one. +// +// The path plays no part: a management URL is dialed, and only its host and +// port are. util.SameServiceURL is this comparison plus the path, which is +// what SameServiceURLIncludingPath needs and delegates to. +func SameServiceURL(a, b *url.URL) bool { + if a == nil || b == nil { + return a == b + } + + return strings.EqualFold(a.Scheme, b.Scheme) && + strings.EqualFold(a.Hostname(), b.Hostname()) && + util.ServiceURLPort(a) == util.ServiceURLPort(b) +} + +// SameServiceURLIncludingPath is SameServiceURL plus everything a URL carries +// past its endpoint: path, query, fragment and userinfo. +// +// Use it for a URL that gets opened rather than dialed. The admin panel can +// live under a path, so two URLs with the same endpoint and different paths are +// two different panels — where for a URL the client dials over gRPC only the +// endpoint is ever used. Equivalent spellings still compare equal: a missing +// path and "/" are the same root, and so is a trailing slash on any path. +func SameServiceURLIncludingPath(a, b *url.URL) bool { + if a == nil || b == nil { + return a == b + } + + return util.SameServiceURL(a, b) && + a.RawQuery == b.RawQuery && + a.Fragment == b.Fragment && + a.User.String() == b.User.String() +} + func parseURL(serviceName, serviceURL string) (*url.URL, error) { parsedMgmtURL, err := url.ParseRequestURI(serviceURL) if err != nil { @@ -930,6 +1092,84 @@ func isPreSharedKeyHidden(preSharedKey *string) bool { return false } +// WouldChange reports whether applying input would modify any field the +// config persists, leaving the receiver untouched. It is the dry-run half of +// UpdateConfig and reuses the very same diff logic (Config.apply), so a +// caller asking "is this a settings change?" cannot drift from what an +// actual update would do, nor go stale when a new field is added. +// +// A redacted pre-shared key is collapsed to "unset" exactly as +// UpdateOrCreateConfig does, so a UI that round-trips the mask is not read as +// a request for a new key. +// +// A nil receiver means the profile holds no config yet, so the baseline is the +// config the daemon would create for it: input values matching those defaults +// change nothing, anything else does. +func (config *Config) WouldChange(input ConfigInput) (bool, error) { + probe := config.clone() + if probe == nil { + baseline, err := newDryRunBaseline(input.ConfigPath) + if err != nil { + return true, fmt.Errorf("build default config baseline: %w", err) + } + probe = baseline + } + probe.probing = true + + // Normalize before measuring. apply() reports two different things through + // one bool: an input that changed a value, and a field it had to fill in + // because the config carried none. Only the first is a settings change, so + // the filling-in gets a pass of its own whose verdict is discarded, and the + // pass that answers the caller runs against a config with nothing left to + // fill in. + // + // Readers already hand out normalized configs — readConfig applies an empty + // input for this very reason — so this is normally a no-op. But a gate that + // refuses a request must not depend on where its caller got the config + // from, and it must not start reading "this profile predates a field" as + // "the caller asked for a change" the day someone adds one. + if _, err := probe.apply(ConfigInput{ConfigPath: input.ConfigPath}); err != nil { + return true, fmt.Errorf("normalize the config to diff against: %w", err) + } + + if isPreSharedKeyHidden(input.PreSharedKey) { + input.PreSharedKey = nil + } + + return probe.apply(input) +} + +// newDryRunBaseline builds the config a brand-new profile would start from, for +// a dry run to compare an input against. It is createNewConfig without the +// identity: this config exists only to be compared against and thrown away, and +// no ConfigInput field maps to either key. +func newDryRunBaseline(configPath string) (*Config, error) { + baseline := newConfigSkeleton() + + if _, err := baseline.apply(ConfigInput{ConfigPath: configPath}); err != nil { + return nil, err + } + + return baseline, nil +} + +// clone returns a copy of the config that apply can be run against without the +// original observing the writes, or nil for a nil receiver. Only what apply +// mutates in place needs detaching, which is the slices it replaces or appends +// to: every pointer field it touches is reassigned rather than written through, +// and ClientCertKeyPair is only overwritten. +func (config *Config) clone() *Config { + if config == nil { + return nil + } + + probe := *config + probe.IFaceBlackList = slices.Clone(config.IFaceBlackList) + probe.NATExternalIPs = slices.Clone(config.NATExternalIPs) + probe.DNSLabels = slices.Clone(config.DNSLabels) + return &probe +} + // UpdateConfig update existing configuration according to input configuration and return with the configuration func UpdateConfig(input ConfigInput) (*Config, error) { configExists, err := fileExists(input.ConfigPath) @@ -940,6 +1180,14 @@ func UpdateConfig(input ConfigInput) (*Config, error) { return nil, fmt.Errorf("config file %s does not exist", input.ConfigPath) } + // A UI that round-trips the mask GetConfig hands it back is asking to keep + // the stored key, not to set the mask as the new one. UpdateOrCreateConfig + // and DirectUpdateOrCreateConfig already collapse it; this one did not, so + // the same round-trip through SetConfig replaced the key with asterisks. + if isPreSharedKeyHidden(input.PreSharedKey) { + input.PreSharedKey = nil + } + return update(input) } @@ -951,7 +1199,7 @@ func UpdateOrCreateConfig(input ConfigInput) (*Config, error) { } if !configExists { log.Infof("generating new config %s", input.ConfigPath) - cfg, err := createNewConfig(input) + cfg, err := createProvisionedConfig(input) if err != nil { return nil, err } @@ -976,12 +1224,20 @@ func update(input ConfigInput) (*Config, error) { return nil, err } + // A write path is a provisioning point: a stored profile can legitimately + // carry no identity (a mobile logout clears the keys in place), and the + // next config write is what has to mint a new one. Reads leave that alone. + identityGenerated, err := config.EnsureIdentity() + if err != nil { + return nil, err + } + updated, err := config.apply(input) if err != nil { return nil, err } - if updated { + if updated || identityGenerated { if err := util.WriteJson(context.Background(), input.ConfigPath, config); err != nil { return nil, err } @@ -990,8 +1246,8 @@ func update(input ConfigInput) (*Config, error) { return config, nil } -// GetConfig read config file and return with Config and if it was created. Errors out if it does not exist -func GetConfig(configPath string) (*Config, error) { +// GetExistingConfig reads and returns the config if it exists on disk. Fails otherwise. +func GetExistingConfig(configPath string) (*Config, error) { return readConfig(configPath, false) } @@ -1074,17 +1330,27 @@ func UpdateOldManagementURL(ctx context.Context, config *Config, configPath stri return newConfig, nil } -// CreateInMemoryConfig generate a new config but do not write out it to the store +// CreateInMemoryConfig generate a new config but do not write out it to the store. +// It carries an identity: callers connect with what they get back. func CreateInMemoryConfig(input ConfigInput) (*Config, error) { - return createNewConfig(input) + return createProvisionedConfig(input) } -// ReadConfig read config file and return with Config. If it is not exists create a new with default values -func ReadConfig(configPath string) (*Config, error) { +// ReadConfigOrDefault reads the profile config at configPath, or resolves the +// default config in memory when the file does not exist. It never writes, and +// never mints an identity — EnsureIdentity is where that happens, so the +// caller that provisions is also the one that persists. +func ReadConfigOrDefault(configPath string) (*Config, error) { return readConfig(configPath, true) } -// ReadConfig read config file and return with Config. If it is not exists create a new with default values +// readConfig reads the profile config at configPath. createIfMissing resolves a +// default config in memory when the file is absent, rather than erroring. +// +// Reads are pure. This used to write the config back whenever apply() had to +// fill in a default the file was missing, which quietly made every reader a +// writer: a gate deciding whether to refuse a request, a UI listing profiles, +// a mobile getter reading a single preference. func readConfig(configPath string, createIfMissing bool) (*Config, error) { configExists, err := fileExists(configPath) if err != nil { @@ -1102,12 +1368,8 @@ func readConfig(configPath string, createIfMissing bool) (*Config, error) { return nil, err } // initialize through apply() without changes - if changed, err := config.apply(ConfigInput{}); err != nil { + if _, err := config.apply(ConfigInput{}); err != nil { return nil, err - } else if changed { - if err = WriteOutConfig(configPath, config); err != nil { - return nil, err - } } return config, nil @@ -1115,13 +1377,7 @@ func readConfig(configPath string, createIfMissing bool) (*Config, error) { return nil, fmt.Errorf("config file %s does not exist", configPath) } - cfg, err := createNewConfig(ConfigInput{ConfigPath: configPath}) - if err != nil { - return nil, err - } - - err = WriteOutConfig(configPath, cfg) - return cfg, err + return createNewConfig(ConfigInput{ConfigPath: configPath}) } // WriteOutConfig write put the prepared config to the given path @@ -1144,7 +1400,7 @@ func DirectUpdateOrCreateConfig(input ConfigInput) (*Config, error) { } if !configExists { log.Infof("generating new config %s", input.ConfigPath) - cfg, err := createNewConfig(input) + cfg, err := createProvisionedConfig(input) if err != nil { return nil, err } @@ -1171,12 +1427,18 @@ func directUpdate(input ConfigInput) (*Config, error) { return nil, err } + // Same provisioning point as update(); see the note there. + identityGenerated, err := config.EnsureIdentity() + if err != nil { + return nil, err + } + updated, err := config.apply(input) if err != nil { return nil, err } - if updated { + if updated || identityGenerated { if err := util.DirectWriteJson(context.Background(), input.ConfigPath, config); err != nil { return nil, err } @@ -1198,7 +1460,16 @@ func ConfigToJSON(config *Config) (string, error) { // ConfigFromJSON deserializes a JSON string to a Config struct. // This is useful for restoring config from alternative storage mechanisms. -// After unmarshaling, defaults are applied to ensure the config is fully initialized. +// After unmarshaling, defaults are applied to ensure the config is fully +// initialized. +// +// The peer identity is deliberately none of its business, in either direction. +// It does not generate one: a read cannot hand back keys that nothing will +// write down (see ReadConfigOrDefault). Nor does it refuse a document that +// carries none, because a config legitimately has no identity between a logout +// and the next login — mobile logout clears both keys in place — and this is +// also the deserializer the iOS SDK copies a config through. Whoever goes on +// to connect is where an absent identity has to be answered. func ConfigFromJSON(jsonStr string) (*Config, error) { config := &Config{} err := json.Unmarshal([]byte(jsonStr), config) diff --git a/client/internal/profilemanager/config_json_test.go b/client/internal/profilemanager/config_json_test.go new file mode 100644 index 000000000..9a6d820c4 --- /dev/null +++ b/client/internal/profilemanager/config_json_test.go @@ -0,0 +1,44 @@ +package profilemanager + +import ( + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" +) + +// The serialized form is how the tvOS SDK stores a profile and how the iOS SDK +// copies one in memory, so it must round-trip whatever a profile legitimately +// holds — including no identity at all, which is the state mobile logout leaves +// behind when it clears both keys in place. Refusing that document here broke +// logout, profile switching and the login that follows them. +func TestConfigFromJSONRoundTripsALoggedOutProfile(t *testing.T) { + path := filepath.Join(t.TempDir(), "exported.json") + stored, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL}) + require.NoError(t, err) + require.NotEmpty(t, stored.PrivateKey, "a provisioned config is the fixture this test starts from") + require.NotEmpty(t, stored.SSHKey) + + exported, err := ConfigToJSON(stored) + require.NoError(t, err) + + restored, err := ConfigFromJSON(exported) + require.NoError(t, err, "a config exported after a login must load") + require.Equal(t, stored.PrivateKey, restored.PrivateKey, "the restored peer is not the stored one") + require.Equal(t, stored.SSHKey, restored.SSHKey) + + // What mobile logout leaves on disk. + loggedOut := stored.clone() + loggedOut.PrivateKey = "" + loggedOut.SSHKey = "" + + document, err := ConfigToJSON(loggedOut) + require.NoError(t, err) + + reloaded, err := ConfigFromJSON(document) + require.NoError(t, err, "a logged-out profile must still load") + require.Empty(t, reloaded.PrivateKey, "loading must not mint a key nothing will write down") + require.Empty(t, reloaded.SSHKey) + require.Equal(t, stored.ManagementURL.String(), reloaded.ManagementURL.String(), + "the rest of the profile survives the logout") +} diff --git a/client/internal/profilemanager/config_optional_fields_test.go b/client/internal/profilemanager/config_optional_fields_test.go new file mode 100644 index 000000000..9b74e2217 --- /dev/null +++ b/client/internal/profilemanager/config_optional_fields_test.go @@ -0,0 +1,131 @@ +package profilemanager + +import ( + "encoding/json" + "os" + "path/filepath" + "reflect" + "testing" + + "github.com/stretchr/testify/require" +) + +// optionalBoolFields lists the *bool fields of Config by name, derived from the +// type so a field added later is covered without touching these tests. +func optionalBoolFields() []string { + pointerToBool := reflect.TypeOf((*bool)(nil)) + + var fields []string + configType := reflect.TypeOf(Config{}) + for i := range configType.NumField() { + field := configType.Field(i) + if field.Type == pointerToBool && field.Tag.Get("json") != "-" { + fields = append(fields, field.Name) + } + } + return fields +} + +func requireNoUnsetOptionalBool(t *testing.T, config *Config, context string) { + t.Helper() + + value := reflect.ValueOf(*config) + for _, name := range optionalBoolFields() { + require.False(t, value.FieldByName(name).IsNil(), + "%s left %s unset, so its readers have to invent a default and a diff of it compares presence instead of value", context, name) + } +} + +// An optional bool must not be tristate. While one can be nil, true or false, +// every reader has to invent the meaning of nil, and — the reason this test +// exists — a diff of the config ends up comparing presence rather than value: +// that is what made the update-settings gate refuse `netbird up` for a client +// restating its own defaults. apply() is where a config becomes complete, so +// the invariant belongs to it: no *bool may come out of apply() unset. +func TestApplyLeavesNoOptionalBoolUnset(t *testing.T) { + require.NotEmpty(t, optionalBoolFields(), "the invariant is only meaningful while Config has optional bools") + + t.Run("a config built from scratch", func(t *testing.T) { + config := newConfigSkeleton() + _, err := config.apply(ConfigInput{}) + require.NoError(t, err) + + requireNoUnsetOptionalBool(t, config, "apply on a new config") + }) + + t.Run("a config file that predates every optional field", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "legacy.json") + require.NoError(t, os.WriteFile(path, []byte(`{"WgIface":"wt0"}`), 0o600)) + + config, err := GetExistingConfig(path) + require.NoError(t, err) + + requireNoUnsetOptionalBool(t, config, "a read of a legacy config") + }) + + t.Run("a config file that stores them as null", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "null.json") + _, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path}) + require.NoError(t, err) + unsetOnDisk(t, path, optionalBoolFields()...) + + config, err := GetExistingConfig(path) + require.NoError(t, err) + + requireNoUnsetOptionalBool(t, config, "a read of a config storing nulls") + }) +} + +// The same invariant on disk: what a write leaves in the file is what the next +// client to read it starts from, so no write may store a null. +func TestNoWriteStoresAnUnsetOptionalBool(t *testing.T) { + requireNoNullOnDisk := func(t *testing.T, path string, context string) { + t.Helper() + + raw, err := os.ReadFile(path) + require.NoError(t, err) + + var stored map[string]json.RawMessage + require.NoError(t, json.Unmarshal(raw, &stored)) + + for _, name := range optionalBoolFields() { + value, present := stored[name] + require.True(t, present, "%s did not store %s at all", context, name) + require.NotEqual(t, "null", string(value), "%s stored %s as null", context, name) + } + } + + t.Run("UpdateOrCreateConfig", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "created.json") + _, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL}) + require.NoError(t, err) + + requireNoNullOnDisk(t, path, "UpdateOrCreateConfig") + }) + + t.Run("UpdateConfig over a config storing nulls", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "stored.json") + _, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path}) + require.NoError(t, err) + unsetOnDisk(t, path, optionalBoolFields()...) + + _, err = UpdateConfig(ConfigInput{ConfigPath: path, ManagementURL: "https://mgmt.example.com"}) + require.NoError(t, err) + + requireNoNullOnDisk(t, path, "UpdateConfig") + }) + + // Renaming used to copy the file back through a bare Unmarshal, which + // preserved the nulls a pre-fix client had written. + t.Run("RenameProfile", func(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + unsetOnDisk(t, created.Path, optionalBoolFields()...) + + require.NoError(t, sm.RenameProfile(created.ID, username, "office")) + + requireNoNullOnDisk(t, created.Path, "RenameProfile") + }) + }) +} diff --git a/client/internal/profilemanager/config_probe_test.go b/client/internal/profilemanager/config_probe_test.go new file mode 100644 index 000000000..35a179a84 --- /dev/null +++ b/client/internal/profilemanager/config_probe_test.go @@ -0,0 +1,96 @@ +package profilemanager + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// writeCertPair writes a throwaway certificate and key, so apply() has +// something real to load rather than a missing file it would only log about. +func writeCertPair(t *testing.T) (certPath, keyPath string) { + t.Helper() + + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + template := x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "probe-test"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + } + der, err := x509.CreateCertificate(rand.Reader, &template, &template, &key.PublicKey, key) + require.NoError(t, err) + + keyDER, err := x509.MarshalECPrivateKey(key) + require.NoError(t, err) + + dir := t.TempDir() + certPath = filepath.Join(dir, "client.crt") + keyPath = filepath.Join(dir, "client.key") + require.NoError(t, os.WriteFile(certPath, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), 0o600)) + require.NoError(t, os.WriteFile(keyPath, pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}), 0o600)) + return certPath, keyPath +} + +// The dry run behind the update-settings gate must not read the mTLS pair off +// disk. The loaded pair feeds the connection, never the comparison, and the +// gate runs it on every SetConfig and Login — twice per request — including the +// ones it refuses. +func TestProbeDoesNotLoadTheCertificatePair(t *testing.T) { + certPath, keyPath := writeCertPair(t) + + t.Run("a real apply loads it", func(t *testing.T) { + config := newConfigSkeleton() + config.ClientCertPath, config.ClientCertKeyPath = certPath, keyPath + + _, err := config.apply(ConfigInput{}) + require.NoError(t, err) + require.NotNil(t, config.ClientCertKeyPair, "the connection would have no client certificate") + }) + + t.Run("a probe does not", func(t *testing.T) { + config := newConfigSkeleton() + config.ClientCertPath, config.ClientCertKeyPath = certPath, keyPath + config.probing = true + + _, err := config.apply(ConfigInput{}) + require.NoError(t, err) + require.Nil(t, config.ClientCertKeyPair, "the dry run read the certificate off disk") + }) + + // And the verdict is the same either way, which is the only thing the gate + // asks of the probe. + t.Run("the verdict is unaffected", func(t *testing.T) { + path := filepath.Join(t.TempDir(), "mtls.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: DefaultManagementURL, + ClientCertPath: certPath, + ClientCertKeyPath: keyPath, + }) + require.NoError(t, err) + + stored, err := GetExistingConfig(path) + require.NoError(t, err) + + changed, err := stored.WouldChange(ConfigInput{ClientCertPath: certPath, ClientCertKeyPath: keyPath}) + require.NoError(t, err) + require.False(t, changed, "restating the stored certificate paths is not a change") + + changed, err = stored.WouldChange(ConfigInput{ClientCertPath: filepath.Join(t.TempDir(), "other.crt")}) + require.NoError(t, err) + require.True(t, changed, "a different certificate path is a change") + }) +} diff --git a/client/internal/profilemanager/config_test.go b/client/internal/profilemanager/config_test.go index 248920b5e..a461aa71f 100644 --- a/client/internal/profilemanager/config_test.go +++ b/client/internal/profilemanager/config_test.go @@ -196,7 +196,7 @@ func TestWireguardPortZeroExplicit(t *testing.T) { assert.Equal(t, 0, config.WgPort, "WgPort should be 0 when explicitly set by user") // Verify it persists - readConfig, err := GetConfig(configPath) + readConfig, err := GetExistingConfig(configPath) require.NoError(t, err) assert.Equal(t, 0, readConfig.WgPort, "WgPort should remain 0 after reading from file") } diff --git a/client/internal/profilemanager/config_would_change_test.go b/client/internal/profilemanager/config_would_change_test.go new file mode 100644 index 000000000..6b140030f --- /dev/null +++ b/client/internal/profilemanager/config_would_change_test.go @@ -0,0 +1,529 @@ +package profilemanager + +import ( + "encoding/json" + "os" + "path/filepath" + "runtime" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/iface" + "github.com/netbirdio/netbird/shared/management/domain" +) + +func seededConfig(t *testing.T) *Config { + t.Helper() + + path := filepath.Join(t.TempDir(), "seeded.json") + cfg, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://api.netbird.io:443", + PreSharedKey: strPointer("stored-key"), + }) + require.NoError(t, err) + return cfg +} + +func strPointer(s string) *string { return &s } + +func intPtr(i int) *int { return &i } + +func TestWouldChange(t *testing.T) { + tests := []struct { + name string + input ConfigInput + want bool + }{ + {name: "empty input", input: ConfigInput{}, want: false}, + {name: "same management URL", input: ConfigInput{ManagementURL: "https://api.netbird.io:443"}, want: false}, + {name: "management URL without its default port", input: ConfigInput{ManagementURL: "https://api.netbird.io"}, want: false}, + {name: "different management URL", input: ConfigInput{ManagementURL: "https://other.example:443"}, want: true}, + {name: "same pre-shared key", input: ConfigInput{PreSharedKey: strPointer("stored-key")}, want: false}, + {name: "redacted pre-shared key", input: ConfigInput{PreSharedKey: strPointer("**********")}, want: false}, + {name: "different pre-shared key", input: ConfigInput{PreSharedKey: strPointer("other-key")}, want: true}, + {name: "new interface blacklist entry", input: ConfigInput{ExtraIFaceBlackList: []string{"nb-probe0"}}, want: true}, + {name: "blacklist entry already present", input: ConfigInput{ExtraIFaceBlackList: []string{"lo"}}, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := seededConfig(t) + + changed, err := cfg.WouldChange(tt.input) + require.NoError(t, err) + require.Equal(t, tt.want, changed) + }) + } +} + +// The dry run must not be observable on the config it is run against: it +// decides whether a write is allowed, it does not perform one. +func TestWouldChangeLeavesTheConfigAlone(t *testing.T) { + cfg := seededConfig(t) + blacklist := len(cfg.IFaceBlackList) + + changed, err := cfg.WouldChange(ConfigInput{ + ManagementURL: "https://other.example:443", + PreSharedKey: strPointer("other-key"), + ExtraIFaceBlackList: []string{"nb-probe0"}, + DNSLabels: domain.FromPunycodeList([]string{"probe"}), + NATExternalIPs: []string{"1.2.3.4"}, + }) + require.NoError(t, err) + require.True(t, changed) + + require.Equal(t, "https://api.netbird.io:443", cfg.ManagementURL.String()) + require.Equal(t, "stored-key", cfg.PreSharedKey) + require.Len(t, cfg.IFaceBlackList, blacklist) + require.Empty(t, cfg.DNSLabels) + require.Empty(t, cfg.NATExternalIPs) +} + +// A nil config means the profile holds nothing yet, so the baseline is what +// the daemon would create for it. +func TestWouldChangeWithoutAStoredConfig(t *testing.T) { + var cfg *Config + + changed, err := cfg.WouldChange(ConfigInput{}) + require.NoError(t, err) + require.False(t, changed, "a request carrying nothing cannot change anything") + + changed, err = cfg.WouldChange(ConfigInput{ManagementURL: DefaultManagementURL}) + require.NoError(t, err) + require.False(t, changed, "the default management URL is what would be written anyway") + + changed, err = cfg.WouldChange(ConfigInput{ManagementURL: "https://other.example:443"}) + require.NoError(t, err) + require.True(t, changed) +} + +func TestWouldChangeReportsAnInvalidInput(t *testing.T) { + cfg := seededConfig(t) + + _, err := cfg.WouldChange(ConfigInput{ManagementURL: "not-a-url"}) + require.Error(t, err) +} + +// Reads must not write. A config file missing a field apply() fills in (MTU, +// here) is what used to trigger the write-back. +func TestReadsDoNotWriteTheConfigBack(t *testing.T) { + denormalized := []byte(`{"WgIface":"wt0"}`) + + for name, read := range map[string]func(string) (*Config, error){ + "GetExistingConfig": GetExistingConfig, + "ReadConfigOrDefault": ReadConfigOrDefault, + } { + t.Run(name, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "profile.json") + require.NoError(t, os.WriteFile(path, denormalized, 0o600)) + + cfg, err := read(path) + require.NoError(t, err) + require.Equal(t, uint16(iface.DefaultMTU), cfg.MTU, "the returned config is still normalized in memory") + require.Empty(t, cfg.PrivateKey, "a read must not mint an identity either") + + after, err := os.ReadFile(path) + require.NoError(t, err) + require.Equal(t, string(denormalized), string(after), "%s rewrote the config file", name) + }) + } +} + +// ReadConfigOrDefault resolves a default config for a profile that has no file +// yet, and that must not create the file either. +func TestReadConfigDoesNotCreateTheFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "absent.json") + + cfg, err := ReadConfigOrDefault(path) + require.NoError(t, err) + require.Equal(t, DefaultManagementURL, cfg.ManagementURL.String()) + + _, err = os.Stat(path) + require.True(t, os.IsNotExist(err), "ReadConfigOrDefault created the config file") +} + +// The identity is the one thing a read cannot recompute, so it is provisioned +// on request and its caller persists it. +func TestEnsureIdentity(t *testing.T) { + cfg := newConfigSkeleton() + + generated, err := cfg.EnsureIdentity() + require.NoError(t, err) + require.True(t, generated) + require.NotEmpty(t, cfg.PrivateKey) + require.NotEmpty(t, cfg.SSHKey) + + key := cfg.PrivateKey + generated, err = cfg.EnsureIdentity() + require.NoError(t, err) + require.False(t, generated, "a config that already has an identity keeps it") + require.Equal(t, key, cfg.PrivateKey) +} + +// One endpoint written several ways is one endpoint. A gate that compared +// spellings refused a client restating its own management URL with a trailing +// slash, which is a normal way to write it. +func TestSameServiceURL(t *testing.T) { + tests := []struct { + a, b string + want bool + }{ + {a: "https://mgmt.example.com", b: "https://mgmt.example.com:443", want: true}, + {a: "https://mgmt.example.com", b: "https://mgmt.example.com/", want: true}, + {a: "https://mgmt.example.com/", b: "https://mgmt.example.com:443/", want: true}, + {a: "https://MGMT.example.com", b: "https://mgmt.example.com", want: true}, + {a: "http://mgmt.example.com", b: "http://mgmt.example.com:80", want: true}, + {a: "https://mgmt.example.com", b: "http://mgmt.example.com", want: false}, + {a: "https://mgmt.example.com", b: "https://mgmt.example.com:8443", want: false}, + {a: "https://mgmt.example.com", b: "https://other.example.com", want: false}, + } + + for _, tt := range tests { + t.Run(tt.a+" vs "+tt.b, func(t *testing.T) { + a, err := ParseServiceURL("a", tt.a) + require.NoError(t, err) + b, err := ParseServiceURL("b", tt.b) + require.NoError(t, err) + + require.Equal(t, tt.want, SameServiceURL(a, b)) + require.Equal(t, tt.want, SameServiceURL(b, a), "the comparison must be symmetric") + }) + } +} + +// The same spellings, through the dry run the update-settings gate uses. +func TestWouldChangeIgnoresURLSpelling(t *testing.T) { + path := filepath.Join(t.TempDir(), "seeded.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://mgmt.example.com", + }) + require.NoError(t, err) + + cfg, err := GetExistingConfig(path) + require.NoError(t, err) + + for _, spelling := range []string{ + "https://mgmt.example.com", + "https://mgmt.example.com/", + "https://mgmt.example.com:443", + "https://mgmt.example.com:443/", + "https://MGMT.example.com", + } { + changed, err := cfg.WouldChange(ConfigInput{ManagementURL: spelling}) + require.NoError(t, err) + require.False(t, changed, "%q is the stored endpoint written differently", spelling) + } + + changed, err := cfg.WouldChange(ConfigInput{ManagementURL: "https://mgmt.example.com:8443"}) + require.NoError(t, err) + require.True(t, changed, "a different port is a different endpoint") +} + +// The dry-run baseline exists to be compared against and discarded, so it must +// not mint keys — the CLI's login backoff loop would otherwise log a fresh +// "generated new Wireguard key" on every attempt. +func TestDryRunBaselineDoesNotGenerateKeys(t *testing.T) { + baseline, err := newDryRunBaseline(filepath.Join(t.TempDir(), "absent.json")) + require.NoError(t, err) + + require.Empty(t, baseline.PrivateKey, "generated a WireGuard key for a throwaway config") + require.Empty(t, baseline.SSHKey, "generated an SSH key for a throwaway config") + + // Everything the comparison actually looks at is still the default config. + require.Equal(t, DefaultManagementURL, baseline.ManagementURL.String()) + require.Equal(t, uint16(iface.DefaultMTU), baseline.MTU) + require.Equal(t, iface.DefaultWgPort, baseline.WgPort) +} + +// A stored profile can carry no identity — a mobile logout clears the keys in +// place — so the next config write has to mint one, which is what keeps the +// following login from dialing management with an empty key. +func TestUpdateConfigProvisionsAMissingIdentity(t *testing.T) { + path := filepath.Join(t.TempDir(), "logged-out.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://api.netbird.io:443", + }) + require.NoError(t, err) + + // Stand in for the logout, which zeroes the keys and writes the config out. + loggedOut, err := GetExistingConfig(path) + require.NoError(t, err) + loggedOut.PrivateKey = "" + loggedOut.SSHKey = "" + require.NoError(t, WriteOutConfig(path, loggedOut)) + + cfg, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path}) + require.NoError(t, err) + require.NotEmpty(t, cfg.PrivateKey, "the write path did not provision an identity") + require.NotEmpty(t, cfg.SSHKey) + + persisted, err := GetExistingConfig(path) + require.NoError(t, err) + require.Equal(t, cfg.PrivateKey, persisted.PrivateKey, "the provisioned identity was not persisted") +} + +// A config that carries no sync message version must not make the dry run +// panic: the gate runs inside a request handler, where failing closed is the +// worst acceptable outcome. +func TestWouldChangeWithoutAStoredSyncMessageVersion(t *testing.T) { + cfg := seededConfig(t) + require.Nil(t, cfg.SyncMessageVersion, "the fixture is only useful while the field starts out unset") + + version := 2 + changed, err := cfg.WouldChange(ConfigInput{SyncMessageVersion: &version}) + require.NoError(t, err) + require.True(t, changed) + require.Nil(t, cfg.SyncMessageVersion, "the dry run set the version on the stored config") +} + +// Restating the certificate paths a config already holds is not a change, for +// the same reason restating any other value is not. +func TestWouldChangeIgnoresRestatedCertificatePaths(t *testing.T) { + path := filepath.Join(t.TempDir(), "mtls.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://api.netbird.io:443", + ClientCertPath: "/etc/netbird/client.crt", + ClientCertKeyPath: "/etc/netbird/client.key", + }) + require.NoError(t, err) + + cfg, err := GetExistingConfig(path) + require.NoError(t, err) + + changed, err := cfg.WouldChange(ConfigInput{ + ClientCertPath: "/etc/netbird/client.crt", + ClientCertKeyPath: "/etc/netbird/client.key", + }) + require.NoError(t, err) + require.False(t, changed, "the stored certificate paths were restated") + + changed, err = cfg.WouldChange(ConfigInput{ClientCertPath: "/etc/netbird/other.crt"}) + require.NoError(t, err) + require.True(t, changed, "a different certificate path is a change") +} + +// A read that lands on a missing file must not hand back keys: nothing would +// write them down, so the caller would connect with an identity that changes on +// the next run and registers a second peer. +func TestReadConfigOrDefaultCarriesNoIdentity(t *testing.T) { + cfg, err := ReadConfigOrDefault(filepath.Join(t.TempDir(), "absent.json")) + require.NoError(t, err) + + require.Empty(t, cfg.PrivateKey, "a read minted a WireGuard key") + require.Empty(t, cfg.SSHKey, "a read minted an SSH key") + + // So the caller's own EnsureIdentity is the one that reports the work, and + // therefore the one that triggers the write. + generated, err := cfg.EnsureIdentity() + require.NoError(t, err) + require.True(t, generated, "the provisioning caller could not tell it had to persist the identity") +} + +// CreateInMemoryConfig is the opposite contract: its callers connect with what +// they get back, so it does carry an identity. +func TestCreateInMemoryConfigCarriesAnIdentity(t *testing.T) { + cfg, err := CreateInMemoryConfig(ConfigInput{ManagementURL: "https://api.netbird.io:443"}) + require.NoError(t, err) + + require.NotEmpty(t, cfg.PrivateKey) + require.NotEmpty(t, cfg.SSHKey) +} + +// The admin panel is opened, not dialed, so its path identifies it. Comparing +// it as a bare endpoint left a custom panel URL unable to change. +func TestAdminURLPathIsPartOfTheIdentity(t *testing.T) { + path := filepath.Join(t.TempDir(), "panel.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + AdminURL: "https://app.example.com/netbird", + }) + require.NoError(t, err) + + cfg, err := GetExistingConfig(path) + require.NoError(t, err) + require.Equal(t, "https://app.example.com:443/netbird", cfg.AdminURL.String()) + + // Equivalent spellings of the same panel are still not a change. + for _, same := range []string{ + "https://app.example.com/netbird", + "https://app.example.com:443/netbird", + "https://app.example.com/netbird/", + "https://APP.example.com/netbird", + } { + changed, err := cfg.WouldChange(ConfigInput{AdminURL: same}) + require.NoError(t, err) + require.False(t, changed, "%q is the stored panel written differently", same) + } + + // A different path is a different panel, and it must be persisted. + changed, err := cfg.WouldChange(ConfigInput{AdminURL: "https://app.example.com/other"}) + require.NoError(t, err) + require.True(t, changed, "a different panel path is a change") + + updated, err := UpdateConfig(ConfigInput{ConfigPath: path, AdminURL: "https://app.example.com/other"}) + require.NoError(t, err) + require.Equal(t, "https://app.example.com:443/other", updated.AdminURL.String(), "the new panel path was not persisted") +} + +// unsetOnDisk rewrites the stored config so the named fields carry a JSON null, +// which is how a profile written before apply() resolved them looks on disk. +// It synthesizes that state: no write produces it any more. +func unsetOnDisk(t *testing.T, path string, fields ...string) { + t.Helper() + + raw, err := os.ReadFile(path) + require.NoError(t, err) + + var stored map[string]json.RawMessage + require.NoError(t, json.Unmarshal(raw, &stored)) + + for _, field := range fields { + _, present := stored[field] + require.True(t, present, "%s is not a field of the stored config", field) + stored[field] = json.RawMessage("null") + } + + rewritten, err := json.Marshal(stored) + require.NoError(t, err) + require.NoError(t, os.WriteFile(path, rewritten, 0600)) +} + +// Seven fields mean "the effective default" when they hold no value, and every +// profile written before apply() resolved them holds them as null. Restating +// that default is asking for no change — and the CLI restates it on every +// `netbird up`, because a flag set through an environment variable is a flag +// pflag reports as Changed. Judging those restatements as changes made the +// update-settings gate refuse `netbird up` outright for a client configured +// through the environment, which is the shape of a Kubernetes deployment. +// +// A login now writes those fields set, so the fixture puts the null state back +// on disk with unsetOnDisk instead of getting it from a login. +func TestWouldChangeIgnoresRestatedDefaultsOfUnsetFields(t *testing.T) { + networkMonitorDefault := runtime.GOOS == "windows" || runtime.GOOS == "darwin" + + tests := []struct { + field string + theDefault ConfigInput + theOtherWay ConfigInput + }{ + {"EnableSSHRoot", + ConfigInput{EnableSSHRoot: boolPtr(false)}, ConfigInput{EnableSSHRoot: boolPtr(true)}}, + {"EnableSSHSFTP", + ConfigInput{EnableSSHSFTP: boolPtr(false)}, ConfigInput{EnableSSHSFTP: boolPtr(true)}}, + {"EnableSSHLocalPortForwarding", + ConfigInput{EnableSSHLocalPortForwarding: boolPtr(false)}, ConfigInput{EnableSSHLocalPortForwarding: boolPtr(true)}}, + {"EnableSSHRemotePortForwarding", + ConfigInput{EnableSSHRemotePortForwarding: boolPtr(false)}, ConfigInput{EnableSSHRemotePortForwarding: boolPtr(true)}}, + {"DisableSSHAuth", + ConfigInput{DisableSSHAuth: boolPtr(false)}, ConfigInput{DisableSSHAuth: boolPtr(true)}}, + {"SSHJWTCacheTTL", + ConfigInput{SSHJWTCacheTTL: intPtr(0)}, ConfigInput{SSHJWTCacheTTL: intPtr(300)}}, + {"NetworkMonitor", + ConfigInput{NetworkMonitor: boolPtr(networkMonitorDefault)}, ConfigInput{NetworkMonitor: boolPtr(!networkMonitorDefault)}}, + } + + for _, tt := range tests { + t.Run(tt.field, func(t *testing.T) { + path := filepath.Join(t.TempDir(), "unset.json") + _, err := UpdateOrCreateConfig(ConfigInput{ + ConfigPath: path, + ManagementURL: "https://api.netbird.io:443", + }) + require.NoError(t, err) + unsetOnDisk(t, path, tt.field) + + cfg, err := GetExistingConfig(path) + require.NoError(t, err) + + changed, err := cfg.WouldChange(tt.theDefault) + require.NoError(t, err) + require.False(t, changed, "restating the default of an unset %s was judged a change", tt.field) + + // The gate still has to refuse a request that does ask for something. + changed, err = cfg.WouldChange(tt.theOtherWay) + require.NoError(t, err) + require.True(t, changed, "asking for a non-default %s is a change", tt.field) + }) + } +} + +// The verdict must not depend on where the caller got the config from. Readers +// normalize what they hand out, but apply() signals "I filled in a default" +// through the same bool as "the input changed something", so a config that +// never passed through a read would otherwise report a change for an input +// that asks for nothing. +func TestWouldChangeNormalizesBeforeMeasuring(t *testing.T) { + rawConfig := func(t *testing.T) *Config { + t.Helper() + + cfg := &Config{WgIface: iface.WgInterfaceDefault} + require.Nil(t, cfg.ServerSSHAllowed, "the fixture is only useful while the config is not normalized") + require.Nil(t, cfg.EnableSSHRoot) + require.Empty(t, cfg.IFaceBlackList) + return cfg + } + + changed, err := rawConfig(t).WouldChange(ConfigInput{}) + require.NoError(t, err) + require.False(t, changed, "an input carrying nothing cannot change anything") + + changed, err = rawConfig(t).WouldChange(ConfigInput{EnableSSHRoot: boolPtr(false)}) + require.NoError(t, err) + require.False(t, changed, "the default of a field the config never held is not a change") + + changed, err = rawConfig(t).WouldChange(ConfigInput{EnableSSHRoot: boolPtr(true)}) + require.NoError(t, err) + require.True(t, changed, "a non-default value is still a change") +} + +// A zero-padded port addresses the same port. The normalization itself belongs +// to util.ServiceURLPort and is tested there; this asserts that the comparison +// this package hands its callers inherits it. +func TestServiceURLPortIsNormalizedNumerically(t *testing.T) { + padded, err := ParseServiceURL("padded", "https://mgmt.example.com:0443") + require.NoError(t, err) + plain, err := ParseServiceURL("plain", "https://mgmt.example.com:443") + require.NoError(t, err) + + require.True(t, SameServiceURL(padded, plain)) +} + +// A list the profile does not have and a list the request empties are the same +// thing: no NAT mappings, no DNS labels. The profile stores an absent list as +// JSON null and reads it back as a nil slice, while `netbird up` sends the +// emptied list — CleanNATExternalIPs / CleanDNSLabels — whenever the matching +// environment variable is set to nothing, which a deployment template does by +// default. Judging nil and empty as different made the gate refuse that start, +// which is the very deadlock this branch exists to remove, on another field. +func TestWouldChangeIgnoresAnEmptiedListThatWasAlreadyAbsent(t *testing.T) { + path := filepath.Join(t.TempDir(), "lists.json") + _, err := UpdateOrCreateConfig(ConfigInput{ConfigPath: path, ManagementURL: DefaultManagementURL}) + require.NoError(t, err) + + stored, err := GetExistingConfig(path) + require.NoError(t, err) + require.Nil(t, stored.NATExternalIPs, "the fixture is only useful while the stored list is absent") + require.Nil(t, stored.DNSLabels) + + changed, err := stored.WouldChange(ConfigInput{NATExternalIPs: make([]string, 0)}) + require.NoError(t, err) + require.False(t, changed, "emptying a NAT list the profile never had is not a change") + + changed, err = stored.WouldChange(ConfigInput{DNSLabels: domain.List{}}) + require.NoError(t, err) + require.False(t, changed, "emptying a DNS label list the profile never had is not a change") + + // A list that does hold something still moves when the request empties it. + withEntries, err := UpdateConfig(ConfigInput{ConfigPath: path, NATExternalIPs: []string{"1.2.3.4"}}) + require.NoError(t, err) + require.Equal(t, []string{"1.2.3.4"}, withEntries.NATExternalIPs) + + changed, err = withEntries.WouldChange(ConfigInput{NATExternalIPs: make([]string, 0)}) + require.NoError(t, err) + require.True(t, changed, "clearing a NAT list that had an entry is a change") +} diff --git a/client/internal/profilemanager/service.go b/client/internal/profilemanager/service.go index ec287f01a..e58f421fd 100644 --- a/client/internal/profilemanager/service.go +++ b/client/internal/profilemanager/service.go @@ -313,7 +313,11 @@ func (s *ServiceManager) AddProfile(displayName, username string) (*Profile, err } profPath := filepath.Join(configDir, id.String()+".json") - cfg, err := createNewConfig(ConfigInput{ConfigPath: profPath}) + // Provisioned, not bare: this config goes straight to disk, and a profile + // file with no identity is one whose first reader has to mint the keys and + // remember to write them back. Before identity generation moved out of + // apply() into EnsureIdentity, createNewConfig produced them here too. + cfg, err := createProvisionedConfig(ConfigInput{ConfigPath: profPath}) if err != nil { return nil, fmt.Errorf("failed to create new config: %w", err) } @@ -330,6 +334,19 @@ func (s *ServiceManager) AddProfile(displayName, username string) (*Profile, err }, nil } +// RenameProfile changes a profile's display name. It rewrites the whole +// profile file, not just the name: the config is read through the normalizing +// reader, so apply()'s resolved values — the optional booleans, the interface +// blacklist, the DNS route interval — are persisted along with the new name. +// +// That is deliberate. A write that skipped apply() is what left profiles on +// disk carrying null where a value was meant, and made a diff of the config +// compare presence instead of value. Two consequences worth knowing: the +// platform-dependent defaults resolved here are the renaming host's +// (ServerSSHAllowed and the network monitor differ per OS), and a profile +// whose stored name does not survive sanitizeDisplayName now fails to rename +// rather than being rewritten — though apply() rejects such a profile on every +// other read too, so it was already unusable. func (s *ServiceManager) RenameProfile(id ID, username string, newName string) error { displayName, err := sanitizeDisplayName(newName) if err != nil { @@ -356,17 +373,17 @@ func (s *ServiceManager) RenameProfile(id ID, username string, newName string) e return ErrProfileNotFound } - data, err := os.ReadFile(target.Path) + // Through the reader, not a bare Unmarshal: this was the one write that + // skipped apply(), so it copied back whatever the file held — including an + // optional field left unset, which every other write resolves to its + // default. Renaming a profile is a poor place to leave that behind. + cfg, err := GetExistingConfig(target.Path) if err != nil { - return err - } - var cfg Config - if err := json.Unmarshal(data, &cfg); err != nil { - return err + return fmt.Errorf("read profile config: %w", err) } cfg.Name = displayName - if err := util.WriteJson(context.Background(), target.Path, cfg); err != nil { + if err := WriteOutConfig(target.Path, cfg); err != nil { return fmt.Errorf("failed to write profile name: %w", err) } return nil diff --git a/client/internal/profilemanager/service_test.go b/client/internal/profilemanager/service_test.go index 5e051b15d..d26ce746a 100644 --- a/client/internal/profilemanager/service_test.go +++ b/client/internal/profilemanager/service_test.go @@ -228,3 +228,27 @@ func TestRemoveProfile_DeletesStateFile(t *testing.T) { assert.True(t, errors.Is(err, os.ErrNotExist), "state file should be removed") }) } + +// A profile file is written here and read back by whoever connects with it, so +// it has to carry the peer's identity. While AddProfile used the bare +// constructor, it wrote a config with no keys: the first reader had to mint +// them, and the paths that read without writing — a gate deciding whether to +// refuse a request, the mobile SDKs loading a stored profile — got a config +// that cannot connect. +func TestAddProfileWritesAnIdentity(t *testing.T) { + withTestSM(t, func(sm *ServiceManager, username string) { + created, err := sm.AddProfile("work", username) + require.NoError(t, err) + + stored, err := GetExistingConfig(created.Path) + require.NoError(t, err) + + require.NotEmpty(t, stored.PrivateKey, "the profile was written without a WireGuard key") + require.NotEmpty(t, stored.SSHKey, "the profile was written without an SSH key") + + // And the identity is the one on disk, not one minted per read. + reread, err := GetExistingConfig(created.Path) + require.NoError(t, err) + require.Equal(t, stored.PrivateKey, reread.PrivateKey) + }) +} diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate.go b/client/internal/routemanager/ipfwdstate/ipfwdstate.go index 3d571e16b..22f7bd07a 100644 --- a/client/internal/routemanager/ipfwdstate/ipfwdstate.go +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate.go @@ -19,8 +19,7 @@ type IPForwardingState struct { // routingV4/routingV6 track whether the routing path currently holds a // reference, so repeated EnableRouting calls (one per network-map update) - // hold at most one reference per family and an unpaired DisableRouting - // can't release references held by DNAT rules. + // hold at most one reference per family. routingV4 bool routingV6 bool @@ -95,31 +94,6 @@ func (f *IPForwardingState) ReleaseRouting() error { return nil } -// RequestForwarding enables the family's forwarding sysctl on first request. -func (f *IPForwardingState) RequestForwarding(v6 bool) error { - f.mu.Lock() - defer f.mu.Unlock() - - if v6 { - return f.requestV6() - } - return f.requestV4() -} - -// ReleaseForwarding decrements the family counter. The last v6 release restores -// what enable captured. v4 stays on: net.ipv4.ip_forward is co-owned by other -// tooling (docker, k8s, libvirt). -func (f *IPForwardingState) ReleaseForwarding(v6 bool) error { - f.mu.Lock() - defer f.mu.Unlock() - - if v6 { - return f.releaseV6() - } - f.releaseV4() - return nil -} - func (f *IPForwardingState) requestV4() error { if f.v4Count == 0 { if err := systemops.EnableV4IPForwarding(); err != nil { diff --git a/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go b/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go index b4615ff02..75209965c 100644 --- a/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go +++ b/client/internal/routemanager/ipfwdstate/ipfwdstate_privileged_linux_test.go @@ -10,8 +10,7 @@ import ( ) // TestRequestRoutingV6ToV4Transition verifies that a v4-only routing request -// releases a previously held routing-owned v6 reference without touching -// references held by DNAT rules. +// releases a previously held routing-owned v6 reference. func TestRequestRoutingV6ToV4Transition(t *testing.T) { f := NewIPForwardingState("wt-fwd-test") @@ -25,13 +24,6 @@ func TestRequestRoutingV6ToV4Transition(t *testing.T) { assert.Equal(t, 1, v4, "v4 reference kept") assert.Equal(t, 0, v6, "routing-owned v6 reference released") - // A DNAT-held reference survives a v4-only routing request. - require.NoError(t, f.RequestForwarding(true), "dnat v6 reference") - require.NoError(t, f.RequestRouting(false), "repeat v4-only request") - _, v6 = f.Counts() - assert.Equal(t, 1, v6, "dnat-held v6 reference survives") - require.NoError(t, f.ReleaseForwarding(true), "release dnat v6 reference") - require.NoError(t, f.ReleaseRouting(), "release routing") v4, v6 = f.Counts() assert.Equal(t, 0, v4, "all v4 references released") diff --git a/client/ios/NetBirdSDK/client.go b/client/ios/NetBirdSDK/client.go index de752a49c..457ae3a7d 100644 --- a/client/ios/NetBirdSDK/client.go +++ b/client/ios/NetBirdSDK/client.go @@ -130,6 +130,10 @@ func NewClient(cfgFile, stateFile, cacheDir, logFilePath, deviceName string, osV // SetConfigFromJSON stores the JSON config that later loads resolve instead of the config file (tvOS). func (c *Client) SetConfigFromJSON(jsonStr string) error { + // Parsed only to reject an unreadable document early; the JSON itself is + // what is stored, and every load re-parses it. A document carrying no peer + // identity is readable and accepted: that is a logged-out profile, and the + // login that follows provisions the keys. if _, err := profilemanager.ConfigFromJSON(jsonStr); err != nil { log.Errorf("SetConfigFromJSON: failed to parse config JSON: %v", err) return err diff --git a/client/ios/NetBirdSDK/login.go b/client/ios/NetBirdSDK/login.go index 3fabee8f9..b182e82b6 100644 --- a/client/ios/NetBirdSDK/login.go +++ b/client/ios/NetBirdSDK/login.go @@ -379,6 +379,35 @@ func (a *Auth) SetConfigFromJSON(jsonStr string) error { } func (a *Auth) setBaseConfig(base *profilemanager.Config) error { + // A logged-out profile carries no keys: the mobile logout clears them in + // place so the next login registers a new peer instead of resurrecting the + // old one. This is that login, and auth.NewAuth parses the WireGuard key + // before the SSO flow even starts, so an absent identity fails the login on + // key size rather than asking the user to sign in. + // + // Minted on the base config, which is the one GetConfigJSON hands back for + // the caller to store — the overlaid copy below is runtime-only. + generated, err := base.EnsureIdentity() + if err != nil { + return fmt.Errorf("ensure profile identity: %w", err) + } + if generated { + if a.cfgPath != "" { + // Non-atomic, like NewAuth's own write: the tvOS App Group sandbox + // blocks the temp-file-and-rename an atomic write needs. + if err := profilemanager.DirectWriteOutConfig(a.cfgPath, base); err != nil { + return fmt.Errorf("write out profile config: %w", err) + } + } else { + // No file to write to — this is the tvOS path, where the profile + // lives in the caller's own store. It persists the new identity by + // calling GetConfigJSON once the login completes; until then the + // keys exist only here, and a login that never completes leaves + // nothing behind. + log.Infof("provisioned a peer identity for a config with no file on disk") + } + } + overlaid, err := copyConfig(base) if err != nil { return err diff --git a/client/ios/NetBirdSDK/preferences.go b/client/ios/NetBirdSDK/preferences.go index 5297920a3..642f9e160 100644 --- a/client/ios/NetBirdSDK/preferences.go +++ b/client/ios/NetBirdSDK/preferences.go @@ -49,7 +49,7 @@ func (p *Preferences) GetManagementURL() (string, error) { return p.configInput.ManagementURL, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return "", err } @@ -67,7 +67,7 @@ func (p *Preferences) GetAdminURL() (string, error) { return p.configInput.AdminURL, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return "", err } @@ -89,7 +89,7 @@ func (p *Preferences) HasPreSharedKey() (bool, error) { return *p.configInput.PreSharedKey != "", nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -115,7 +115,7 @@ func (p *Preferences) GetRosenpassEnabled() (bool, error) { return *p.configInput.RosenpassEnabled, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -136,7 +136,7 @@ func (p *Preferences) GetRosenpassPermissive() (bool, error) { return *p.configInput.RosenpassPermissive, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -149,7 +149,7 @@ func (p *Preferences) GetDisableIPv6() (bool, error) { return *p.configInput.DisableIPv6, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } @@ -168,7 +168,7 @@ func (p *Preferences) GetRemoteJobsAllowed() (bool, error) { return *p.configInput.RemoteJobsAllowed, nil } - cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) + cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath) if err != nil { return false, err } diff --git a/client/mobile/profile_lifecycle_test.go b/client/mobile/profile_lifecycle_test.go new file mode 100644 index 000000000..9612f550d --- /dev/null +++ b/client/mobile/profile_lifecycle_test.go @@ -0,0 +1,97 @@ +package mobile + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/client/internal/profilemanager" +) + +// loadAsTheMobileSDKsDo replays what the iOS SDK does with a stored profile: +// read the config, serialize it, and load it back. Client.SetConfigFromJSON +// stores that document for tvOS, Auth.SetConfigFromJSON authenticates with it, +// and copyConfig round-trips a Config through the same pair to take an +// in-memory copy before applying the MDM overlay. +func loadAsTheMobileSDKsDo(t *testing.T, configPath string) *profilemanager.Config { + t.Helper() + + stored, err := profilemanager.GetExistingConfig(configPath) + require.NoError(t, err, "read the stored profile") + + document, err := profilemanager.ConfigToJSON(stored) + require.NoError(t, err, "serialize the stored profile") + + reloaded, err := profilemanager.ConfigFromJSON(document) + require.NoError(t, err, "load the profile back") + return reloaded +} + +// A profile survives the whole round its user puts it through: created, logged +// out, loaded again, and switched away from and back. +// +// Logout is the step that makes this worth asserting. It clears the peer's +// keys in place so the next login registers a new peer rather than bringing +// the old one back, which leaves a profile that legitimately carries no +// identity — and both mobile SDKs go on loading that profile through the +// serialized form. A load that refused it, or a creation that never wrote an +// identity in the first place, breaks logout and profile switching on iOS and +// Android without any of it being visible from the desktop client. +func TestProfileSurvivesLogoutAndReload(t *testing.T) { + pm := newTestProfileManager(t) + + created, err := pm.AddProfile("work") + require.NoError(t, err) + require.NoError(t, pm.SwitchProfile(profilemanager.DefaultProfileName)) + + // Created: the profile carries the identity it will connect with. + require.NotEmpty(t, privateKeyOf(t, pm, created.ID), "a new profile was written with no identity") + + configPath, err := pm.GetConfigPath(created.ID) + require.NoError(t, err) + + before := loadAsTheMobileSDKsDo(t, configPath) + require.NotEmpty(t, before.PrivateKey) + managementURL := before.ManagementURL.String() + + // Logged out: the identity is gone, on purpose. + require.NoError(t, pm.LogoutProfile(created.ID)) + require.Empty(t, privateKeyOf(t, pm, created.ID), "logout left the peer's key behind") + + // Loaded again: the profile is still readable, and loading it neither + // fails nor mints a key that nothing would write down. + after := loadAsTheMobileSDKsDo(t, configPath) + assert.Empty(t, after.PrivateKey, "loading a logged-out profile minted a key nothing will persist") + assert.Empty(t, after.SSHKey, "loading a logged-out profile minted an SSH key") + assert.Equal(t, managementURL, after.ManagementURL.String(), "the rest of the profile did not survive the logout") + + // Switched away from and back: still the same profile, still loadable. + require.NoError(t, pm.SwitchProfile(created.ID)) + require.NoError(t, pm.SwitchProfile(profilemanager.DefaultProfileName)) + require.NoError(t, pm.SwitchProfile(created.ID)) + + active, err := pm.GetActiveProfile() + require.NoError(t, err) + assert.Equal(t, created.ID, active.ID, "the profile switched to is not the active one") + + assert.Equal(t, managementURL, loadAsTheMobileSDKsDo(t, configPath).ManagementURL.String(), + "the profile did not survive the round of switches") +} + +// The profile the SDKs fall back to gets the same treatment, since it is the +// one a mobile client without an explicit profile runs on. +func TestDefaultProfileSurvivesLogoutAndReload(t *testing.T) { + pm := newTestProfileManager(t) + require.NoError(t, pm.SwitchProfile(profilemanager.DefaultProfileName)) + + configPath, err := pm.GetConfigPath(profilemanager.DefaultProfileName) + require.NoError(t, err) + require.NotEmpty(t, loadAsTheMobileSDKsDo(t, configPath).PrivateKey) + + require.NoError(t, pm.LogoutProfile(profilemanager.DefaultProfileName)) + + reloaded := loadAsTheMobileSDKsDo(t, configPath) + assert.Empty(t, reloaded.PrivateKey, "loading the logged-out default profile minted a key") + assert.NotNil(t, reloaded.ManagementURL, "the profile lost its management URL") +} diff --git a/client/mobile/profile_manager.go b/client/mobile/profile_manager.go index 348b7253b..ad79d80c0 100644 --- a/client/mobile/profile_manager.go +++ b/client/mobile/profile_manager.go @@ -192,7 +192,10 @@ func (pm *ProfileManager) LogoutProfile(id string) error { return fmt.Errorf("profile %q does not exist", id) } - config, err := profilemanager.ReadConfig(configPath) + // The existing-file reader, not the generating one: the check above is not + // atomic with this read, so a profile removed in between would otherwise be + // resolved from the defaults here and recreated by the write below. + config, err := profilemanager.GetExistingConfig(configPath) if err != nil { return fmt.Errorf("read profile config: %w", err) } diff --git a/client/netbird.wxs b/client/netbird.wxs index f30a7aa7e..156b4ff27 100644 --- a/client/netbird.wxs +++ b/client/netbird.wxs @@ -76,6 +76,14 @@ + + + +