Merge remote-tracking branch 'origin/main' into fix/pkce-flow-session-extend

# Conflicts:
#	shared/management/proto/management.pb.go
This commit is contained in:
Zoltán Papp
2026-10-07 13:54:40 +02:00
213 changed files with 9278 additions and 8536 deletions
+2
View File
@@ -12,6 +12,8 @@ jobs:
docs-ack: docs-ack:
name: Require docs PR URL or explicit "not needed" name: Require docs PR URL or explicit "not needed"
runs-on: ubuntu-latest 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: steps:
- name: Read PR body - name: Read PR body
+1 -1
View File
@@ -30,7 +30,7 @@ jobs:
# segment by codespell and behave the same across versions; the # segment by codespell and behave the same across versions; the
# recursive "**" form did not take effect with the codespell shipped # recursive "**" form did not take effect with the codespell shipped
# by this action. # 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: golangci:
strategy: strategy:
fail-fast: false fail-fast: false
+2
View File
@@ -7,6 +7,8 @@ on:
jobs: jobs:
check-title: check-title:
runs-on: ubuntu-latest 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: steps:
- name: Validate PR title prefix - name: Validate PR title prefix
uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9.0.0 uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9.0.0
+2
View File
@@ -32,6 +32,7 @@ on:
- all - all
- client-rootless - client-rootless
- reverse-proxy - reverse-proxy
- netbird-server
version: version:
description: "Released version, e.g. v0.80.0" description: "Released version, e.g. v0.80.0"
type: string type: string
@@ -66,6 +67,7 @@ jobs:
components=( components=(
"client-rootless ghcr.io/netbirdio/netbird -rootless-ubi" "client-rootless ghcr.io/netbirdio/netbird -rootless-ubi"
"reverse-proxy ghcr.io/netbirdio/reverse-proxy -ubi" "reverse-proxy ghcr.io/netbirdio/reverse-proxy -ubi"
"netbird-server ghcr.io/netbirdio/netbird-server -ubi"
) )
matrix="[]" matrix="[]"
missing=() missing=()
+1 -1
View File
@@ -199,7 +199,7 @@ jobs:
with: with:
node-version: '22' node-version: '22'
- name: Install proxy web dependencies for license collection - 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 working-directory: proxy/web
run: npm ci --ignore-scripts run: npm ci --ignore-scripts
- name: Set up QEMU - name: Set up QEMU
+3 -2
View File
@@ -36,7 +36,8 @@ jobs:
with: with:
node-version: "22" node-version: "22"
# English (en) is the source of truth for translation keys; every other # English (en) is the source of truth for translation keys. Locales declared
# locale declared in _index.json must carry the exact same key set. # 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 - name: Check translation key parity
run: node client/ui/i18n/check-translations.mjs run: node client/ui/i18n/check-translations.mjs
+37 -2
View File
@@ -385,7 +385,7 @@ dockers_v2:
RELEASE: "{{ .Timestamp }}" RELEASE: "{{ .Timestamp }}"
hooks: hooks:
pre: 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: env:
- GOOS=linux - GOOS=linux
- CGO_ENABLED=0 - CGO_ENABLED=0
@@ -511,6 +511,41 @@ dockers_v2:
"org.opencontainers.image.revision": "{{.FullCommit}}" "org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}" "org.opencontainers.image.source": "{{.GitURL}}"
"maintainer": "dev@netbird.io" "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 - id: netbird-proxy
disable: "{{ .Env.SKIP_DOCKER_PUSH }}" disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
ids: ids:
@@ -552,7 +587,7 @@ dockers_v2:
RELEASE: "{{ .Timestamp }}" RELEASE: "{{ .Timestamp }}"
hooks: hooks:
pre: 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: env:
- GOOS=linux - GOOS=linux
- CGO_ENABLED=0 - CGO_ENABLED=0
+1 -1
View File
@@ -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. 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 BSD 3-Clause License
+1 -2
View File
@@ -104,8 +104,7 @@ type Client struct {
stateChangeMu sync.Mutex stateChangeMu sync.Mutex
stateChangeSubID string stateChangeSubID string
eventSub *peer.EventSubscription // Closed to stop the watch goroutine from delivering buffered ticks to a
// Closed to stop the watch goroutines from delivering buffered items to a
// listener that has been removed or replaced. See stopStateChangeWatchLocked. // listener that has been removed or replaced. See stopStateChangeWatchLocked.
stateChangeDone chan struct{} stateChangeDone chan struct{}
+17 -17
View File
@@ -46,7 +46,7 @@ func (p *Preferences) GetManagementURL() (string, error) {
return p.configInput.ManagementURL, nil return p.configInput.ManagementURL, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return "", err return "", err
} }
@@ -64,7 +64,7 @@ func (p *Preferences) GetAdminURL() (string, error) {
return p.configInput.AdminURL, nil return p.configInput.AdminURL, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return "", err return "", err
} }
@@ -86,7 +86,7 @@ func (p *Preferences) HasPreSharedKey() (bool, error) {
return *p.configInput.PreSharedKey != "", nil return *p.configInput.PreSharedKey != "", nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -112,7 +112,7 @@ func (p *Preferences) GetRosenpassEnabled() (bool, error) {
return *p.configInput.RosenpassEnabled, nil return *p.configInput.RosenpassEnabled, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -133,7 +133,7 @@ func (p *Preferences) GetRosenpassPermissive() (bool, error) {
return *p.configInput.RosenpassPermissive, nil return *p.configInput.RosenpassPermissive, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -149,7 +149,7 @@ func (p *Preferences) GetDisableClientRoutes() (bool, error) {
return *p.configInput.DisableClientRoutes, nil return *p.configInput.DisableClientRoutes, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -170,7 +170,7 @@ func (p *Preferences) GetDisableServerRoutes() (bool, error) {
return *p.configInput.DisableServerRoutes, nil return *p.configInput.DisableServerRoutes, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -188,7 +188,7 @@ func (p *Preferences) GetDisableDNS() (bool, error) {
return *p.configInput.DisableDNS, nil return *p.configInput.DisableDNS, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -206,7 +206,7 @@ func (p *Preferences) GetDisableFirewall() (bool, error) {
return *p.configInput.DisableFirewall, nil return *p.configInput.DisableFirewall, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -227,7 +227,7 @@ func (p *Preferences) GetServerSSHAllowed() (bool, error) {
return *p.configInput.ServerSSHAllowed, nil return *p.configInput.ServerSSHAllowed, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -249,7 +249,7 @@ func (p *Preferences) GetEnableSSHRoot() (bool, error) {
return *p.configInput.EnableSSHRoot, nil return *p.configInput.EnableSSHRoot, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -271,7 +271,7 @@ func (p *Preferences) GetEnableSSHSFTP() (bool, error) {
return *p.configInput.EnableSSHSFTP, nil return *p.configInput.EnableSSHSFTP, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -293,7 +293,7 @@ func (p *Preferences) GetEnableSSHLocalPortForwarding() (bool, error) {
return *p.configInput.EnableSSHLocalPortForwarding, nil return *p.configInput.EnableSSHLocalPortForwarding, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -315,7 +315,7 @@ func (p *Preferences) GetEnableSSHRemotePortForwarding() (bool, error) {
return *p.configInput.EnableSSHRemotePortForwarding, nil return *p.configInput.EnableSSHRemotePortForwarding, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -340,7 +340,7 @@ func (p *Preferences) GetBlockInbound() (bool, error) {
return *p.configInput.BlockInbound, nil return *p.configInput.BlockInbound, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -358,7 +358,7 @@ func (p *Preferences) GetDisableIPv6() (bool, error) {
return *p.configInput.DisableIPv6, nil return *p.configInput.DisableIPv6, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -377,7 +377,7 @@ func (p *Preferences) GetRemoteJobsAllowed() (bool, error) {
return *p.configInput.RemoteJobsAllowed, nil return *p.configInput.RemoteJobsAllowed, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
+10 -81
View File
@@ -6,13 +6,8 @@ import (
"context" "context"
"fmt" "fmt"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/auth" "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. // StateChangeListener receives client state notifications.
@@ -21,16 +16,11 @@ import (
// changed: connection state, the run-loop status label (e.g. NeedsLogin) or // changed: connection state, the run-loop status label (e.g. NeedsLogin) or
// the session deadline. It mirrors the daemon's SubscribeStatus stream // the session deadline. It mirrors the daemon's SubscribeStatus stream
// trigger — on each signal the consumer pulls the fresh values via // trigger — on each signal the consumer pulls the fresh values via
// Status() / SessionExpiresAtUnix(). // Status() / SessionExpiresAtUnix(). The engine arms no expiry-warning
// // timers on Android; the app schedules the warnings from the deadline it
// OnSessionExpiring forwards the engine's session-expiry warnings, fired at // reads here.
// 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.
type StateChangeListener interface { type StateChangeListener interface {
OnStateChanged() OnStateChanged()
OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool)
} }
// Status returns the connect run-loop's status label — the same value the // Status returns the connect run-loop's status label — the same value the
@@ -110,11 +100,11 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
return return
} }
// Both subscriptions are buffered (one pending tick, ten pending events), // The subscription is buffered (one pending tick), so unsubscribing is
// so unsubscribing is not enough to stop callbacks: the loops would drain // not enough to stop callbacks: the loop would drain what is already
// what is already queued and deliver it to a listener the caller has // queued and deliver it to a listener the caller has already removed or
// already removed or replaced. Gate every callback on this registration's // replaced. Gate every callback on this registration's own signal, which
// own signal, which is closed before unsubscribing. // is closed before unsubscribing.
done := make(chan struct{}) done := make(chan struct{})
c.stateChangeDone = done c.stateChangeDone = done
@@ -133,9 +123,6 @@ func (c *Client) SetStateChangeListener(listener StateChangeListener) {
listener.OnStateChanged() listener.OnStateChanged()
} }
}() }()
c.eventSub = c.recorder.SubscribeToEvents()
go watchSessionWarnings(c.eventSub, listener, done)
} }
// RemoveStateChangeListener unregisters the state notification listener. // RemoveStateChangeListener unregisters the state notification listener.
@@ -145,21 +132,6 @@ func (c *Client) RemoveStateChangeListener() {
c.stopStateChangeWatchLocked() 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 // ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and
// asks the management server to extend the session deadline. The tunnel is // asks the management server to extend the session deadline. The tunnel is
// untouched: no resync, no reconnect. Async; the result arrives on the // untouched: no resync, no reconnect. Async; the result arrives on the
@@ -201,8 +173,8 @@ func (c *Client) CancelExtendAuthSession() {
} }
func (c *Client) stopStateChangeWatchLocked() { func (c *Client) stopStateChangeWatchLocked() {
// Signal first, unsubscribe second: closing the channels only stops new // Signal first, unsubscribe second: closing the channel only stops new
// items, and the loops would still hand whatever is buffered to a listener // items, and the loop would still hand whatever is buffered to a listener
// that is no longer registered. // that is no longer registered.
if c.stateChangeDone != nil { if c.stateChangeDone != nil {
close(c.stateChangeDone) close(c.stateChangeDone)
@@ -212,49 +184,6 @@ func (c *Client) stopStateChangeWatchLocked() {
c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID) c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID)
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) { func (c *Client) beginExtend() (context.Context, error) {
-98
View File
@@ -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"
}
}
+20 -7
View File
@@ -9,8 +9,6 @@ import (
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"github.com/spf13/cobra" "github.com/spf13/cobra"
"golang.org/x/term" "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"
"github.com/netbirdio/netbird/client/internal/auth" "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 { err = WithBackOff(func() error {
var backOffErr error var backOffErr error
loginResp, backOffErr = client.Login(ctx, &loginRequest) loginResp, backOffErr = client.Login(ctx, &loginRequest)
if s, ok := gstatus.FromError(backOffErr); ok && (s.Code() == codes.InvalidArgument || if terminalLoginError(backOffErr) {
s.Code() == codes.PermissionDenied ||
s.Code() == codes.NotFound ||
s.Code() == codes.Unimplemented) {
loginErr = backOffErr loginErr = backOffErr
return nil 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 { if err != nil {
return fmt.Errorf("read config file %s: %v", configFilePath, err) 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, // CLI standalone login: profilemanager no longer auto-applies MDM,
// so layer in the OS-native policy here. Desktop builds construct // so layer in the OS-native policy here. Desktop builds construct
// a Loader with no fetcher — the build-tagged loadPlatform reads // a Loader with no fetcher — the build-tagged loadPlatform reads
+39 -3
View File
@@ -20,6 +20,8 @@ import (
"github.com/spf13/cobra" "github.com/spf13/cobra"
"github.com/spf13/pflag" "github.com/spf13/pflag"
"google.golang.org/grpc" "google.golang.org/grpc"
"google.golang.org/grpc/codes"
gstatus "google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/anonymize" "github.com/netbirdio/netbird/client/anonymize"
daddr "github.com/netbirdio/netbird/client/internal/daemonaddr" daddr "github.com/netbirdio/netbird/client/internal/daemonaddr"
@@ -175,7 +177,6 @@ func init() {
rootCmd.AddCommand(versionCmd) rootCmd.AddCommand(versionCmd)
rootCmd.AddCommand(sshCmd) rootCmd.AddCommand(sshCmd)
rootCmd.AddCommand(networksCMD) rootCmd.AddCommand(networksCMD)
rootCmd.AddCommand(forwardingRulesCmd)
rootCmd.AddCommand(debugCmd) rootCmd.AddCommand(debugCmd)
rootCmd.AddCommand(profileCmd) rootCmd.AddCommand(profileCmd)
rootCmd.AddCommand(exposeCmd) rootCmd.AddCommand(exposeCmd)
@@ -183,8 +184,6 @@ func init() {
networksCMD.AddCommand(routesListCmd) networksCMD.AddCommand(routesListCmd)
networksCMD.AddCommand(routesSelectCmd, routesDeselectCmd) networksCMD.AddCommand(routesSelectCmd, routesDeselectCmd)
forwardingRulesCmd.AddCommand(forwardingRulesListCmd)
debugCmd.AddCommand(debugBundleCmd) debugCmd.AddCommand(debugBundleCmd)
debugCmd.AddCommand(logCmd) debugCmd.AddCommand(logCmd)
logCmd.AddCommand(logLevelCmd) logCmd.AddCommand(logLevelCmd)
@@ -285,6 +284,43 @@ func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, e
return grpc.DialContext(ctx, target, opts...) 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. // WithBackOff execute function in backoff cycle.
func WithBackOff(bf func() error) error { func WithBackOff(bf func() error) error {
return backoff.RetryNotify(bf, CLIBackOffSettings, func(err error, duration time.Duration) { return backoff.RetryNotify(bf, CLIBackOffSettings, func(err error, duration time.Duration) {
+3 -4
View File
@@ -6,9 +6,9 @@ import (
"testing" "testing"
"time" "time"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.opentelemetry.io/otel" "go.opentelemetry.io/otel"
"go.uber.org/mock/gomock"
"google.golang.org/grpc" "google.golang.org/grpc"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator" "github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator"
@@ -28,7 +28,6 @@ import (
mgmt "github.com/netbirdio/netbird/management/server" mgmt "github.com/netbirdio/netbird/management/server"
"github.com/netbirdio/netbird/management/server/activity" "github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/groups" "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/permissions"
"github.com/netbirdio/netbird/management/server/settings" "github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/management/server/store" "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) updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := mgmt.NewAccountRequestBuffer(ctx, store) 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 { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
+28 -7
View File
@@ -357,9 +357,17 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
// set the new config // set the new config
req := setupSetConfigReq(customDNSAddressConverted, cmd, activeProf.ID.String(), username.Username) req := setupSetConfigReq(customDNSAddressConverted, cmd, activeProf.ID.String(), username.Username)
if _, err := client.SetConfig(ctx, req); err != nil { if _, err := client.SetConfig(ctx, req); err != nil {
if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unavailable { switch reason, refused := refusedSettingsUpdate(err); {
log.Warnf("setConfig method is not available in the daemon: %s", st.Message()) case refused:
} else { // 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) 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 { err = WithBackOff(func() error {
var backOffErr error var backOffErr error
loginResp, backOffErr = client.Login(ctx, loginRequest) loginResp, backOffErr = client.Login(ctx, loginRequest)
if s, ok := gstatus.FromError(backOffErr); ok && (s.Code() == codes.InvalidArgument || if terminalLoginError(backOffErr) {
s.Code() == codes.PermissionDenied ||
s.Code() == codes.NotFound ||
s.Code() == codes.Unimplemented) {
loginErr = backOffErr loginErr = backOffErr
return nil 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 { func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, profileName, username string) *proto.SetConfigRequest {
var req proto.SetConfigRequest var req proto.SetConfigRequest
req.ProfileName = profileName req.ProfileName = profileName
+85
View File
@@ -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))
})
}
}
-77
View File
@@ -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"
+3 -4
View File
@@ -6,8 +6,8 @@ import (
"testing" "testing"
"time" "time"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"google.golang.org/grpc" "google.golang.org/grpc"
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller" "github.com/netbirdio/netbird/management/internals/controllers/network_map/controller"
@@ -21,7 +21,6 @@ import (
nbcache "github.com/netbirdio/netbird/management/server/cache" nbcache "github.com/netbirdio/netbird/management/server/cache"
"github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/groups"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator" "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/job"
"github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings" "github.com/netbirdio/netbird/management/server/settings"
@@ -146,8 +145,8 @@ func startManagement(t *testing.T, signalAddr string) string {
updateManager := update_channel.NewPeersUpdateManager(metrics) updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := mgmt.NewAccountRequestBuffer(context.Background(), testStore) 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) 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, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore) accountManager, err := mgmt.BuildManager(context.Background(), cfg, testStore, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, settingsMockManager, permissionsManager, false, cacheStore)
require.NoError(t, err) require.NoError(t, err)
secretsManager, err := nbgrpc.NewTimeBasedAuthSecretsManager(updateManager, cfg.TURNConfig, cfg.Relay, settingsMockManager, groupsManager) secretsManager, err := nbgrpc.NewTimeBasedAuthSecretsManager(updateManager, cfg.TURNConfig, cfg.Relay, settingsMockManager, groupsManager)
-166
View File
@@ -8,177 +8,11 @@ import (
"strconv" "strconv"
"strings" "strings"
"github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors"
firewall "github.com/netbirdio/netbird/client/firewall/manager" 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 { 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)) ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
@@ -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")
}
-4
View File
@@ -56,10 +56,6 @@ const (
markManglePost = "mark-mangle-post" markManglePost = "mark-mangle-post"
matchSet = "--match-set" 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 is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
ipv4TCPHeaderSize = 40 ipv4TCPHeaderSize = 40
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation. // ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation.
-9
View File
@@ -81,15 +81,6 @@ func (r *family) hasRule(id nbid.RuleID) bool {
return ok 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
// "<id>_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 // DeleteFilterRule removes a previously installed filter rule. The
// rule's stored chain/table identify where to delete from; source set // rule's stored chain/table identify where to delete from; source set
// references are recovered from the spec via findSets and dropped // references are recovered from the spec via findSets and dropped
-25
View File
@@ -323,31 +323,6 @@ func (m *Manager) DisableRouting() error {
return m.family4.ipFwdState.ReleaseRouting() 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 // UpdateSet updates the set with the given prefixes
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error { func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
m.mutex.Lock() m.mutex.Lock()
@@ -497,16 +497,6 @@ func TestIptablesCloseRemovesAllState(t *testing.T) {
require.NoError(t, manager.AddNatRule(pair), "add nat rule") require.NoError(t, manager.AddNatRule(pair), "add nat rule")
require.NoError(t, manager.EnableRouting(), "enable routing") 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") 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. // Everything above stays in place, so Close is what has to remove it.
-6
View File
@@ -172,12 +172,6 @@ type Manager interface {
DisableRouting() error 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 updates the set with the given prefixes
UpdateSet(hash Set, prefixes []netip.Prefix) error UpdateSet(hash Set, prefixes []netip.Prefix) error
-27
View File
@@ -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())
}
-321
View File
@@ -9,332 +9,11 @@ import (
"github.com/google/nftables" "github.com/google/nftables"
"github.com/google/nftables/binaryutil" "github.com/google/nftables/binaryutil"
"github.com/google/nftables/expr" "github.com/google/nftables/expr"
"github.com/google/nftables/xt"
"github.com/hashicorp/go-multierror"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors"
firewall "github.com/netbirdio/netbird/client/firewall/manager" 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 <iface> 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 { 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)) ruleID := firewall.RuleID(fmt.Sprintf("inbound-dnat-%s-%s-%d-%d", localAddr.String(), protocol, originalPort, translatedPort))
@@ -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")
}
-8
View File
@@ -24,7 +24,6 @@ const (
tableRaw = "raw" tableRaw = "raw"
tableSecurity = "security" tableSecurity = "security"
chainNameNatPrerouting = "PREROUTING"
chainNameRoutingFw = "netbird-rt-fwd" chainNameRoutingFw = "netbird-rt-fwd"
chainNameRoutingNat = "netbird-rt-postrouting" chainNameRoutingNat = "netbird-rt-postrouting"
chainNameRoutingRdr = "netbird-rt-redirect" chainNameRoutingRdr = "netbird-rt-redirect"
@@ -47,9 +46,6 @@ const (
userDataAcceptForwardRuleOif = "frwacceptoif" userDataAcceptForwardRuleOif = "frwacceptoif"
userDataAcceptInputRule = "inputaccept" userDataAcceptInputRule = "inputaccept"
dnatSuffix firewall.RuleID = "_dnat"
snatSuffix firewall.RuleID = "_snat"
// ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation. // ipv4TCPHeaderSize is the minimum IPv4 (20) + TCP (20) header size for MSS calculation.
ipv4TCPHeaderSize = 40 ipv4TCPHeaderSize = 40
// ipv6TCPHeaderSize is the minimum IPv6 (40) + TCP (20) header size for MSS calculation. // 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) 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) return nberrors.FormatErrorOrNil(merr)
} }
-5
View File
@@ -197,11 +197,6 @@ func (r *family) hasRule(id firewall.RuleID) bool {
return ok 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 // DeleteFilterRule removes a previously installed filter rule. Source
// set references are recovered from the stored rule's expressions via // set references are recovered from the stored rule's expressions via
// findSets and dropped from the shared refcounter. // findSets and dropped from the shared refcounter.
+3 -44
View File
@@ -252,7 +252,7 @@ func (m *Manager) DeleteFilterRule(rule firewall.Rule) error {
m.mutex.Lock() m.mutex.Lock()
defer m.mutex.Unlock() defer m.mutex.Unlock()
fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule, false) fam, err := m.familyForRuleID(rule.ID(), (*family).hasRule)
if err != nil { if err != nil {
return err 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 // 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 supplied lookup, and falls back to the v4 family on a miss.
// the NAT/DNAT rule maps from the kernel once and re-checks before falling func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall.RuleID) bool) (*family, error) {
// 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) {
if has(m.family4, id) { if has(m.family4, id) {
return m.family4, nil return m.family4, nil
} }
@@ -274,18 +271,6 @@ func (m *Manager) familyForRuleID(id firewall.RuleID, has func(*family, firewall
if has(m.family6, id) { if has(m.family6, id) {
return m.family6, nil 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 return m.family4, nil
} }
@@ -450,32 +435,6 @@ func (m *Manager) Flush() error {
return nil 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 // UpdateSet updates the set with the given prefixes
func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error { func (m *Manager) UpdateSet(set firewall.Set, prefixes []netip.Prefix) error {
m.mutex.Lock() m.mutex.Lock()
@@ -378,18 +378,6 @@ func TestNftablesManagerCompatibilityWithIptables(t *testing.T) {
err = manager.AddNatRule(pair) err = manager.AddNatRule(pair)
require.NoError(t, err, "failed to add NAT rule") 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) stdout, stderr = runIptablesSave(t)
verifyIptablesOutput(t, stdout, stderr) verifyIptablesOutput(t, stdout, stderr)
} }
@@ -453,18 +441,6 @@ func TestNftablesManagerIPv6CompatibilityWithIp6tables(t *testing.T) {
}) })
require.NoError(t, err, "add v6 NAT rule") 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) stdout, stderr := runIptablesSave(t)
verifyIptablesOutput(t, stdout, stderr) verifyIptablesOutput(t, stdout, stderr)
-35
View File
@@ -459,41 +459,6 @@ func (r *family) RemoveAllLegacyRouteRules() error {
return nberrors.FormatErrorOrNil(merr) 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 { func (r *family) RemoveNatRule(pair firewall.RouterPair) error {
if err := r.refreshRulesMap(); err != nil { if err := r.refreshRulesMap(); err != nil {
return fmt.Errorf(refreshRulesMapError, err) return fmt.Errorf(refreshRulesMapError, err)
-10
View File
@@ -486,16 +486,6 @@ func incrementalUpdate(oldChecksum uint16, oldBytes, newBytes []byte) uint16 {
return ^uint16(sum) 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. // addPortRedirection adds a port redirection rule.
func (m *Manager) addPortRedirection(targetIP netip.Addr, protocol gopacket.LayerType, originalPort, translatedPort uint16) error { func (m *Manager) addPortRedirection(targetIP netip.Addr, protocol gopacket.LayerType, originalPort, translatedPort uint16) error {
m.portDNATMutex.Lock() m.portDNATMutex.Lock()
+67 -3
View File
@@ -90,8 +90,9 @@ type StatusRecorder interface {
// fallback T-FinalWarningLead dialog (suppressed when the user dismissed // fallback T-FinalWarningLead dialog (suppressed when the user dismissed
// the first one for the same deadline). Safe for concurrent use. // the first one for the same deadline). Safe for concurrent use.
type Watcher struct { type Watcher struct {
lead time.Duration lead time.Duration
finalLead time.Duration finalLead time.Duration
deadlineOnly bool
mu sync.Mutex mu sync.Mutex
current time.Time current time.Time
@@ -102,6 +103,7 @@ type Watcher struct {
dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal
closed bool closed bool
recorder StatusRecorder recorder StatusRecorder
nowFn func() time.Time
} }
// New returns a watcher with the package defaults WarningLead and // 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, lead: lead,
finalLead: final, finalLead: final,
recorder: recorder, 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 // 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 // a Sync push from the server omits the field because login expiration
// was disabled). // was disabled).
@@ -181,7 +191,7 @@ func (w *Watcher) Update(deadline time.Time) error {
w.finalFiredAt = time.Time{} w.finalFiredAt = time.Time{}
w.dismissedAt = time.Time{} w.dismissedAt = time.Time{}
if deadline.After(now) { if deadline.After(now) && !w.deadlineOnly {
w.armTimerLocked(deadline) w.armTimerLocked(deadline)
} }
recorder := w.recorder recorder := w.recorder
@@ -303,6 +313,11 @@ func (w *Watcher) fire(armedFor time.Time) {
w.mu.Unlock() w.mu.Unlock()
return return
} }
now := w.nowFn()
if isLate(now, armedFor, max(w.finalLead, 0)) {
w.fireLateLocked(armedFor, now)
return
}
w.firedAt = armedFor w.firedAt = armedFor
recorder := w.recorder recorder := w.recorder
w.mu.Unlock() w.mu.Unlock()
@@ -331,6 +346,14 @@ func (w *Watcher) fireFinal(armedFor time.Time) {
log.Infof("auth session final-warning skipped (dismissed by user)") log.Infof("auth session final-warning skipped (dismissed by user)")
return 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 w.finalFiredAt = armedFor
recorder := w.recorder recorder := w.recorder
w.mu.Unlock() w.mu.Unlock()
@@ -341,6 +364,39 @@ func (w *Watcher) fireFinal(armedFor time.Time) {
publishWarning(recorder, armedFor, true) 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 // armOneShotLocked schedules cb at fireAt. When fireAt is already in the
// past it dispatches on the next scheduler tick so a state-change recorder // past it dispatches on the next scheduler tick so a state-change recorder
// notification (invoked after w.mu is released) lands first. Caller must // 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, 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))
}
@@ -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()) 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)
}
}
+1
View File
@@ -846,6 +846,7 @@ func TestAddConfig_AllFieldsCovered(t *testing.T) {
"ClientCertKeyPair": "non-config: parsed cert pair, not serialized", "ClientCertKeyPair": "non-config: parsed cert pair, not serialized",
"Name": "non-config: profile name is not needed for debug purposes", "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", "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", "DebugBundleUploadURL": "sensitive: MDM-provided upload URL may carry credentials or query tokens; kept out of the shared bundle",
} }
+4 -88
View File
@@ -42,7 +42,6 @@ import (
dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config" dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config"
"github.com/netbirdio/netbird/client/internal/dnsfwd" "github.com/netbirdio/netbird/client/internal/dnsfwd"
"github.com/netbirdio/netbird/client/internal/expose" "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/lazyconn"
"github.com/netbirdio/netbird/client/internal/metrics" "github.com/netbirdio/netbird/client/internal/metrics"
"github.com/netbirdio/netbird/client/internal/netflow" "github.com/netbirdio/netbird/client/internal/netflow"
@@ -262,11 +261,10 @@ type Engine struct {
statusRecorder *peer.Status statusRecorder *peer.Status
firewall firewallManager.Manager firewall firewallManager.Manager
routeManager routemanager.Manager routeManager routemanager.Manager
acl acl.Manager acl acl.Manager
dnsForwardMgr *dnsfwd.Manager dnsForwardMgr *dnsfwd.Manager
ingressGatewayMgr *ingressgw.Manager
dnsServer dns.Server dnsServer dns.Server
@@ -448,13 +446,6 @@ func (e *Engine) stopLocked() {
e.cleanupSSHConfig() 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 { if e.srWatcher != nil {
e.srWatcher.Close() e.srWatcher.Close()
} }
@@ -1627,13 +1618,6 @@ func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap) error {
e.updateDNSForwarder(dnsRouteFeatureFlag, fwdEntries) e.updateDNSForwarder(dnsRouteFeatureFlag, fwdEntries)
done() 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())) log.Debugf("got peers update from Management Service, total peers to connect to = %d", len(networkMap.GetRemotePeers()))
done = e.phase("offline_peers") 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 // 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 // connection: those that are not lazy by policy (the per-peer lazy state or the
// account flag, subject to the local override). // account flag, subject to the local override).
+2 -3
View File
@@ -42,7 +42,6 @@ import (
nbcache "github.com/netbirdio/netbird/management/server/cache" nbcache "github.com/netbirdio/netbird/management/server/cache"
"github.com/netbirdio/netbird/management/server/groups" "github.com/netbirdio/netbird/management/server/groups"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator" "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/job"
"github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings" "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) updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := server.NewAccountRequestBuffer(context.Background(), store) 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) 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, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore) accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, settingsMockManager, permissionsManager, false, cacheStore)
if err != nil { if err != nil {
return nil, "", err return nil, "", err
} }
+7 -5
View File
@@ -1,4 +1,4 @@
//go:build !js //go:build !js && !android
package internal package internal
@@ -7,10 +7,12 @@ import (
"github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/peer"
) )
// newSessionWatcher returns the real SSO session expiry watcher for every // newSessionWatcher returns the real SSO session expiry watcher. The js/wasm
// non-wasm build. The js/wasm build gets a no-op stub from // build gets a no-op stub from engine_sessionwatch_js.go so the sessionwatch
// engine_sessionwatch_js.go so the sessionwatch package (and its timer // package (and its timer machinery) never links into the wasm binary; the
// machinery) never links into the wasm binary. // 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 { func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
return sessionwatch.New(recorder) return sessionwatch.New(recorder)
} }
@@ -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)
}
-111
View File
@@ -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
}
-281
View File
@@ -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))
}
}
-43
View File
@@ -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
}
+16 -12
View File
@@ -828,7 +828,8 @@ func (conn *Conn) evalStatus() ConnStatus {
// //
// The result is a tri-state: // The result is a tri-state:
// - ConnStatusConnected: all available transports are up // - 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 // - ConnStatusDisconnected: no working transport
func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) { func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
defer func() { defer func() {
@@ -845,13 +846,14 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
} }
return evalConnStatus(connStatusInputs{ return evalConnStatus(connStatusInputs{
forceRelay: IsForceRelayed(), forceRelay: IsForceRelayed(),
peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(), peerUsesRelay: conn.workerRelay.IsRelayConnectionSupportedWithPeer(),
relayConnected: conn.statusRelay.Get() == worker.StatusConnected, relayConnected: conn.statusRelay.Get() == worker.StatusConnected,
remoteSupportsICE: conn.handshaker.RemoteICESupported(), relayTransportConnected: conn.workerRelay.IsTransportConnected(),
iceWorkerCreated: iceWorkerCreated, remoteSupportsICE: conn.handshaker.RemoteICESupported(),
iceStatusConnecting: conn.statusICE.Get() != worker.StatusDisconnected, iceWorkerCreated: iceWorkerCreated,
iceInProgress: iceInProgress, iceStatusConnected: conn.statusICE.Get() == worker.StatusConnected,
iceInProgress: iceInProgress,
}) })
} }
@@ -1060,19 +1062,21 @@ func evalConnStatus(in connStatusInputs) guard.ConnStatus {
return boolToConnStatus(relayUsedAndUp) return boolToConnStatus(relayUsedAndUp)
} }
// ICE counts as "up" when the status is anything other than Disconnected, OR // ICE counts as "running" when either connected or attempting to connect.
// when a negotiation is currently in progress (so we don't spam offers while one is in flight). iceRunning := in.iceStatusConnected || in.iceInProgress
iceUp := in.iceStatusConnecting || in.iceInProgress
// Relay side is acceptable if the peer doesn't rely on relay, or relay is connected. // Relay side is acceptable if the peer doesn't rely on relay, or relay is connected.
relayOK := !in.peerUsesRelay || in.relayConnected relayOK := !in.peerUsesRelay || in.relayConnected
switch { switch {
case iceUp && relayOK: case iceRunning && relayOK:
return guard.ConnStatusConnected return guard.ConnStatusConnected
case relayUsedAndUp: case relayUsedAndUp:
// Relay is up but ICE is down — partially connected. // Relay is up but ICE is down — partially connected.
return guard.ConnStatusPartiallyConnected 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: default:
return guard.ConnStatusDisconnected return guard.ConnStatusDisconnected
} }
+8 -7
View File
@@ -17,13 +17,14 @@ const (
// tri-state connection classification. Extracted so the decision logic can be unit-tested // tri-state connection classification. Extracted so the decision logic can be unit-tested
// without constructing full Worker/Handshaker objects. // without constructing full Worker/Handshaker objects.
type connStatusInputs struct { type connStatusInputs struct {
forceRelay bool // NB_FORCE_RELAY or JS/WASM forceRelay bool // NB_FORCE_RELAY or JS/WASM
peerUsesRelay bool // remote peer advertises relay support AND local has relay peerUsesRelay bool // remote peer advertises relay support AND local has relay
relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay) relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay)
remoteSupportsICE bool // remote peer sent ICE credentials relayTransportConnected bool // the relay transport shared by all peers on that server is up
iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode) remoteSupportsICE bool // remote peer sent ICE credentials
iceStatusConnecting bool // statusICE is anything other than Disconnected iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode)
iceInProgress bool // a negotiation is currently in flight iceStatusConnected bool // statusICE reports Connected
iceInProgress bool // a negotiation is currently in flight
} }
// ConnStatus describe the status of a peer's connection // ConnStatus describe the status of a peer's connection
+71 -12
View File
@@ -30,6 +30,21 @@ func TestEvalConnStatus_ForceRelay(t *testing.T) {
}, },
want: guard.ConnStatusDisconnected, 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", name: "force relay, peer does NOT use relay - disconnected forever",
in: connStatusInputs{ in: connStatusInputs{
@@ -123,24 +138,28 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true in.peerUsesRelay = true
in.relayConnected = true in.relayConnected = true
in.iceStatusConnecting = true in.relayTransportConnected = true
in.iceStatusConnected = true
}, },
want: guard.ConnStatusConnected, 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) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false in.peerUsesRelay = false
in.relayConnected = 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, want: guard.ConnStatusConnected,
}, },
{ {
name: "ICE InProgress only, peer does NOT use relay", name: "ICE InProgress only, peer does NOT use relay",
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false in.peerUsesRelay = false
in.iceStatusConnecting = false in.iceStatusConnected = false
in.iceInProgress = true in.iceInProgress = true
}, },
want: guard.ConnStatusConnected, want: guard.ConnStatusConnected,
@@ -150,7 +169,8 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true in.peerUsesRelay = true
in.relayConnected = true in.relayConnected = true
in.iceStatusConnecting = false in.relayTransportConnected = true
in.iceStatusConnected = false
in.iceInProgress = false in.iceInProgress = false
}, },
want: guard.ConnStatusPartiallyConnected, want: guard.ConnStatusPartiallyConnected,
@@ -160,21 +180,60 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false in.peerUsesRelay = false
in.relayConnected = false in.relayConnected = false
in.iceStatusConnecting = false in.iceStatusConnected = false
in.iceInProgress = false in.iceInProgress = false
}, },
want: guard.ConnStatusDisconnected, 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) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = true in.peerUsesRelay = true
in.relayConnected = false 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, want: guard.ConnStatusDisconnected,
}, },
{ {
@@ -182,7 +241,7 @@ func TestEvalConnStatus_FullyAvailable(t *testing.T) {
mutator: func(in *connStatusInputs) { mutator: func(in *connStatusInputs) {
in.peerUsesRelay = false in.peerUsesRelay = false
in.relayConnected = true // not actually used since peer doesn't rely on it in.relayConnected = true // not actually used since peer doesn't rely on it
in.iceStatusConnecting = false in.iceStatusConnected = false
in.iceInProgress = false in.iceInProgress = false
}, },
want: guard.ConnStatusDisconnected, want: guard.ConnStatusDisconnected,
+5 -3
View File
@@ -14,7 +14,8 @@ type ConnStatus int
const ( const (
// ConnStatusDisconnected means neither ICE nor Relay is connected. // ConnStatusDisconnected means neither ICE nor Relay is connected.
ConnStatusDisconnected ConnStatus = iota 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 ConnStatusPartiallyConnected
// ConnStatusConnected means all required connections are established. // ConnStatusConnected means all required connections are established.
ConnStatusConnected ConnStatusConnected
@@ -87,8 +88,9 @@ func (g *Guard) SetICEConnDisconnected() {
// - Connected: no action, the peer is fully reachable. // - Connected: no action, the peer is fully reachable.
// - Disconnected (neither ICE nor Relay): retries aggressively with exponential backoff (800ms doubling // - 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. // 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 // - PartiallyConnected (one transport usable, the other not): retries up to 3 times
// to one attempt per hour. This limits signaling traffic when relay already provides connectivity. // 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 // 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. // counter and backoff ticker, giving ICE a fresh chance after network conditions change.
+52 -8
View File
@@ -12,6 +12,8 @@ type notifier struct {
serverStateLock sync.Mutex serverStateLock sync.Mutex
listenersLock sync.Mutex listenersLock sync.Mutex
listener Listener listener Listener
peerListWake chan struct{}
peerListStop chan struct{}
currentClientState bool currentClientState bool
lastNotification ClientState lastNotification ClientState
lastNumberOfPeers int lastNumberOfPeers int
@@ -62,7 +64,6 @@ func (n *notifier) setNetworkAvailable(available bool) {
func (n *notifier) setListener(listener Listener) { func (n *notifier) setListener(listener Listener) {
n.serverStateLock.Lock() n.serverStateLock.Lock()
lastNotification := n.effectiveState(n.lastNotification) lastNotification := n.effectiveState(n.lastNotification)
numOfPeers := n.lastNumberOfPeers
fqdnAddress := n.lastFqdnAddress fqdnAddress := n.lastFqdnAddress
address := n.lastIPAddress address := n.lastIPAddress
n.serverStateLock.Unlock() n.serverStateLock.Unlock()
@@ -70,17 +71,19 @@ func (n *notifier) setListener(listener Listener) {
n.listenersLock.Lock() n.listenersLock.Lock()
defer n.listenersLock.Unlock() defer n.listenersLock.Unlock()
n.stopPeerListDelivererLocked()
n.listener = listener n.listener = listener
listener.OnAddressChanged(fqdnAddress, address) listener.OnAddressChanged(fqdnAddress, address)
notifyListener(listener, lastNotification) notifyListener(listener, lastNotification)
// run on go routine to avoid on Java layer to call go functions on same thread n.startPeerListDelivererLocked(listener)
go listener.OnPeersListChanged(numOfPeers) n.wakePeerListDelivererLocked()
} }
func (n *notifier) removeListener() { func (n *notifier) removeListener() {
n.listenersLock.Lock() n.listenersLock.Lock()
defer n.listenersLock.Unlock() defer n.listenersLock.Unlock()
n.stopPeerListDelivererLocked()
n.listener = nil n.listener = nil
} }
@@ -178,15 +181,56 @@ func (n *notifier) peerListChanged(numOfPeers int) {
n.serverStateLock.Unlock() n.serverStateLock.Unlock()
n.listenersLock.Lock() n.listenersLock.Lock()
listener := n.listener defer n.listenersLock.Unlock()
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 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 func (n *notifier) wakePeerListDelivererLocked() {
go listener.OnPeersListChanged(numOfPeers) 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) { func (n *notifier) localAddressChanged(fqdn, address string) {
+155
View File
@@ -2,7 +2,9 @@ package peer
import ( import (
"sync" "sync"
"sync/atomic"
"testing" "testing"
"time"
) )
type mocListener struct { type mocListener struct {
@@ -115,3 +117,156 @@ func Test_notifier_RemoveListener(t *testing.T) {
t.Errorf("invalid state: %d", listener.peers) 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)
}
}
-35
View File
@@ -18,9 +18,7 @@ import (
"google.golang.org/protobuf/types/known/durationpb" "google.golang.org/protobuf/types/known/durationpb"
"google.golang.org/protobuf/types/known/timestamppb" "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/iface/configurer"
"github.com/netbirdio/netbird/client/internal/ingressgw"
"github.com/netbirdio/netbird/client/internal/relay" "github.com/netbirdio/netbird/client/internal/relay"
"github.com/netbirdio/netbird/client/proto" "github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/route"
@@ -161,7 +159,6 @@ type FullStatus struct {
RosenpassState RosenpassState RosenpassState RosenpassState
Relays []relay.ProbeResult Relays []relay.ProbeResult
NSGroupStates []NSGroupState NSGroupStates []NSGroupState
NumOfForwardingRules int
LazyConnectionEnabled bool LazyConnectionEnabled bool
Events []*proto.SystemEvent Events []*proto.SystemEvent
} }
@@ -247,8 +244,6 @@ type Status struct {
// read it without taking mux. // read it without taking mux.
networksRevision atomic.Uint64 networksRevision atomic.Uint64
ingressGwMgr *ingressgw.Manager
routeIDLookup routeIDLookup routeIDLookup routeIDLookup
wgIface WGIfaceStatus wgIface WGIfaceStatus
} }
@@ -276,12 +271,6 @@ func (d *Status) SetRelayMgr(manager *relayClient.Manager) {
d.relayMgr = manager d.relayMgr = manager
} }
func (d *Status) SetIngressGwMgr(ingressGwMgr *ingressgw.Manager) {
d.mux.Lock()
defer d.mux.Unlock()
d.ingressGwMgr = ingressGwMgr
}
// ReplaceOfflinePeers replaces // ReplaceOfflinePeers replaces
func (d *Status) ReplaceOfflinePeers(replacement []State) { func (d *Status) ReplaceOfflinePeers(replacement []State) {
d.mux.Lock() d.mux.Lock()
@@ -332,18 +321,6 @@ func (d *Status) GetPeer(peerPubKey string) (State, error) {
return state, nil 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. // PeerStateByIP returns the full peer State for the given tunnel IP.
// Matches against either the IPv4 (State.IP) or IPv6 (State.IPv6) tunnel // Matches against either the IPv4 (State.IP) or IPv6 (State.IPv6) tunnel
// address so dual-stack peers are reachable on either family. Only // address so dual-stack peers are reachable on either family. Only
@@ -1163,16 +1140,6 @@ func (d *Status) GetRelayStates() []relay.ProbeResult {
return relayStates 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 { func (d *Status) GetDNSStates() []NSGroupState {
d.mux.RLock() d.mux.RLock()
defer d.mux.RUnlock() defer d.mux.RUnlock()
@@ -1207,7 +1174,6 @@ func (d *Status) GetFullStatus() FullStatus {
Relays: d.GetRelayStates(), Relays: d.GetRelayStates(),
RosenpassState: d.GetRosenpassState(), RosenpassState: d.GetRosenpassState(),
NSGroupStates: d.GetDNSStates(), NSGroupStates: d.GetDNSStates(),
NumOfForwardingRules: len(d.ForwardingRules()),
LazyConnectionEnabled: d.GetLazyConnection(), LazyConnectionEnabled: d.GetLazyConnection(),
} }
@@ -1579,7 +1545,6 @@ func (fs FullStatus) ToProto() *proto.FullStatus {
pbFullStatus.LocalPeerState.WgPort = int32(fs.LocalPeerState.WgPort) pbFullStatus.LocalPeerState.WgPort = int32(fs.LocalPeerState.WgPort)
pbFullStatus.LocalPeerState.RosenpassPermissive = fs.RosenpassState.Permissive pbFullStatus.LocalPeerState.RosenpassPermissive = fs.RosenpassState.Permissive
pbFullStatus.LocalPeerState.RosenpassEnabled = fs.RosenpassState.Enabled pbFullStatus.LocalPeerState.RosenpassEnabled = fs.RosenpassState.Enabled
pbFullStatus.NumberOfForwardingRules = int32(fs.NumOfForwardingRules)
pbFullStatus.LazyConnectionEnabled = fs.LazyConnectionEnabled pbFullStatus.LazyConnectionEnabled = fs.LazyConnectionEnabled
pbFullStatus.LocalPeerState.Networks = maps.Keys(fs.LocalPeerState.Routes) pbFullStatus.LocalPeerState.Networks = maps.Keys(fs.LocalPeerState.Routes)
+4
View File
@@ -101,6 +101,10 @@ func (w *WorkerRelay) RelayIsSupportedLocally() bool {
return w.relayManager.HasRelayAddress() return w.relayManager.HasRelayAddress()
} }
func (w *WorkerRelay) IsTransportConnected() bool {
return w.relayManager.Ready()
}
func (w *WorkerRelay) CloseConn() { func (w *WorkerRelay) CloseConn() {
w.relayLock.Lock() w.relayLock.Lock()
conn := w.relayedConn conn := w.relayedConn
+377 -106
View File
@@ -10,7 +10,6 @@ import (
"os" "os"
"os/user" "os/user"
"path/filepath" "path/filepath"
"reflect"
"runtime" "runtime"
"slices" "slices"
"strings" "strings"
@@ -198,6 +197,11 @@ type Config struct {
MTU uint16 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 // policy is the MDM policy that produced the currently-set values
// for any MDM-enforced fields. Set by ApplyMDMPolicy on every // for any MDM-enforced fields. Set by ApplyMDMPolicy on every
// invocation. Never persisted to disk. Callers query enforcement // invocation. Never persisted to disk. Callers query enforcement
@@ -300,9 +304,11 @@ func fileExists(path string) (bool, error) {
return false, err return false, err
} }
// createNewConfig creates a new config generating a new Wireguard key and saving to file // newConfigSkeleton returns the field values a brand-new profile config starts
func createNewConfig(input ConfigInput) (*Config, error) { // from, before apply() fills in the rest. Shared with the dry-run baseline so
config := &Config{ // 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 // defaults to false only for new (post 0.26) configurations
ServerSSHAllowed: util.False(), ServerSSHAllowed: util.False(),
// Remote jobs are an explicit opt-in and default off, including for // 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(), RemoteJobsAllowed: util.False(),
WgPort: iface.DefaultWgPort, 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 { if _, err := config.apply(input); err != nil {
return nil, err return nil, err
@@ -318,6 +409,52 @@ func createNewConfig(input ConfigInput) (*Config, error) {
return config, nil 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) { func (config *Config) apply(input ConfigInput) (updated bool, err error) {
if config.Name != "" { if config.Name != "" {
sanitized, err := sanitizeDisplayName(config.Name) sanitized, err := sanitizeDisplayName(config.Name)
@@ -329,6 +466,13 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true 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 { if config.ManagementURL == nil {
log.Infof("using default Management URL %s", DefaultManagementURL) log.Infof("using default Management URL %s", DefaultManagementURL)
config.ManagementURL, err = parseURL("Management URL", 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 return false, err
} }
} }
if input.ManagementURL != "" && input.ManagementURL != config.ManagementURL.String() { // The comparison is on the endpoint the URL addresses, not on its
log.Infof("new Management URL provided, updated to %#v (old value %#v)", // spelling: the same endpoint can be written several ways (an implicit
input.ManagementURL, config.ManagementURL.String()) // :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) URL, err := parseURL("Management URL", input.ManagementURL)
if err != nil { if err != nil {
return false, err return false, err
} }
config.ManagementURL = URL if !SameServiceURL(URL, config.ManagementURL) {
updated = true log.Infof("new Management URL provided, updated to %#v (old value %#v)",
} else if config.ManagementURL == nil { URL.String(), config.ManagementURL.String())
log.Infof("using default Management URL %s", DefaultManagementURL) config.ManagementURL = URL
config.ManagementURL, err = parseURL("Management URL", DefaultManagementURL) updated = true
if err != nil {
return false, err
} }
} }
@@ -360,31 +505,20 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
return false, err return false, err
} }
} }
if input.AdminURL != "" && input.AdminURL != config.AdminURL.String() { // The admin panel is opened, not dialed, so unlike the Management URL its
log.Infof("new Admin Panel URL provided, updated to %#v (old value %#v)", // path is part of what identifies it: a panel served under /netbird is not
input.AdminURL, config.AdminURL.String()) // the one served at the root.
if input.AdminURL != "" {
newURL, err := parseURL("Admin Panel URL", input.AdminURL) newURL, err := parseURL("Admin Panel URL", input.AdminURL)
if err != nil { if err != nil {
return updated, err return updated, err
} }
config.AdminURL = newURL if !SameServiceURLIncludingPath(newURL, config.AdminURL) {
updated = true log.Infof("new Admin Panel URL provided, updated to %#v (old value %#v)",
} newURL.String(), config.AdminURL.String())
config.AdminURL = newURL
if config.PrivateKey == "" { updated = true
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
} }
config.SSHKey = string(pem)
updated = true
} }
if input.WireguardPort != nil && *input.WireguardPort != config.WgPort { if input.WireguardPort != nil && *input.WireguardPort != config.WgPort {
@@ -405,7 +539,14 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true 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 ])", log.Infof("updating NAT External IP [ %s ] (old value: [ %s ])",
strings.Join(input.NATExternalIPs, " "), strings.Join(input.NATExternalIPs, " "),
strings.Join(config.NATExternalIPs, " ")) strings.Join(config.NATExternalIPs, " "))
@@ -443,21 +584,12 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true 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) log.Infof("switching Network Monitor to %t", *input.NetworkMonitor)
config.NetworkMonitor = input.NetworkMonitor config.NetworkMonitor = input.NetworkMonitor
updated = true 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 { if input.CustomDNSAddress != nil && string(input.CustomDNSAddress) != config.CustomDNSAddress {
log.Infof("updating custom DNS address %#v (old value %#v)", log.Infof("updating custom DNS address %#v (old value %#v)",
string(input.CustomDNSAddress), config.CustomDNSAddress) string(input.CustomDNSAddress), config.CustomDNSAddress)
@@ -490,7 +622,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true updated = true
} }
if input.ServerSSHAllowed != nil && (config.ServerSSHAllowed == nil || *input.ServerSSHAllowed != *config.ServerSSHAllowed) { if input.ServerSSHAllowed != nil && *input.ServerSSHAllowed != *config.ServerSSHAllowed {
if *input.ServerSSHAllowed { if *input.ServerSSHAllowed {
log.Infof("enabling SSH server") log.Infof("enabling SSH server")
} else { } else {
@@ -498,20 +630,9 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
} }
config.ServerSSHAllowed = input.ServerSSHAllowed config.ServerSSHAllowed = input.ServerSSHAllowed
updated = true 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 { if *input.RemoteJobsAllowed {
log.Infof("enabling remote jobs") log.Infof("enabling remote jobs")
} else { } else {
@@ -519,14 +640,9 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
} }
config.RemoteJobsAllowed = input.RemoteJobsAllowed config.RemoteJobsAllowed = input.RemoteJobsAllowed
updated = true 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 { if *input.EnableSSHRoot {
log.Infof("enabling SSH root login") log.Infof("enabling SSH root login")
} else { } else {
@@ -536,7 +652,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true updated = true
} }
if input.EnableSSHSFTP != nil && (config.EnableSSHSFTP == nil || *input.EnableSSHSFTP != *config.EnableSSHSFTP) { if input.EnableSSHSFTP != nil && *input.EnableSSHSFTP != *config.EnableSSHSFTP {
if *input.EnableSSHSFTP { if *input.EnableSSHSFTP {
log.Infof("enabling SSH SFTP subsystem") log.Infof("enabling SSH SFTP subsystem")
} else { } else {
@@ -546,7 +662,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true updated = true
} }
if input.EnableSSHLocalPortForwarding != nil && (config.EnableSSHLocalPortForwarding == nil || *input.EnableSSHLocalPortForwarding != *config.EnableSSHLocalPortForwarding) { if input.EnableSSHLocalPortForwarding != nil && *input.EnableSSHLocalPortForwarding != *config.EnableSSHLocalPortForwarding {
if *input.EnableSSHLocalPortForwarding { if *input.EnableSSHLocalPortForwarding {
log.Infof("enabling SSH local port forwarding") log.Infof("enabling SSH local port forwarding")
} else { } else {
@@ -556,7 +672,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true updated = true
} }
if input.EnableSSHRemotePortForwarding != nil && (config.EnableSSHRemotePortForwarding == nil || *input.EnableSSHRemotePortForwarding != *config.EnableSSHRemotePortForwarding) { if input.EnableSSHRemotePortForwarding != nil && *input.EnableSSHRemotePortForwarding != *config.EnableSSHRemotePortForwarding {
if *input.EnableSSHRemotePortForwarding { if *input.EnableSSHRemotePortForwarding {
log.Infof("enabling SSH remote port forwarding") log.Infof("enabling SSH remote port forwarding")
} else { } else {
@@ -566,7 +682,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true updated = true
} }
if input.DisableSSHAuth != nil && (config.DisableSSHAuth == nil || *input.DisableSSHAuth != *config.DisableSSHAuth) { if input.DisableSSHAuth != nil && *input.DisableSSHAuth != *config.DisableSSHAuth {
if *input.DisableSSHAuth { if *input.DisableSSHAuth {
log.Infof("disabling SSH authentication") log.Infof("disabling SSH authentication")
} else { } else {
@@ -576,7 +692,7 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true 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) log.Infof("updating SSH JWT cache TTL to %d seconds", *input.SSHJWTCacheTTL)
config.SSHJWTCacheTTL = input.SSHJWTCacheTTL config.SSHJWTCacheTTL = input.SSHJWTCacheTTL
updated = true updated = true
@@ -659,13 +775,16 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true 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) log.Infof("setting SyncMessageVersion to %v", *input.SyncMessageVersion)
*config.SyncMessageVersion = *input.SyncMessageVersion config.SyncMessageVersion = input.SyncMessageVersion
updated = true updated = true
} }
if input.DisableNotifications != nil && (config.DisableNotifications == nil || *input.DisableNotifications != *config.DisableNotifications) { if input.DisableNotifications != nil && *input.DisableNotifications != *config.DisableNotifications {
if *input.DisableNotifications { if *input.DisableNotifications {
log.Infof("disabling notifications") log.Infof("disabling notifications")
} else { } else {
@@ -675,24 +794,24 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
updated = true updated = true
} }
if config.DisableNotifications == nil { // Compared, not just assigned: restating the path a config already holds
disabled := true // changes nothing, and reporting it as an update makes a caller that
config.DisableNotifications = &disabled // re-sends its own configuration look like one asking to change it.
log.Infof("setting notifications to disabled by default") if input.ClientCertKeyPath != "" && input.ClientCertKeyPath != config.ClientCertKeyPath {
updated = true
}
if input.ClientCertKeyPath != "" {
config.ClientCertKeyPath = input.ClientCertKeyPath config.ClientCertKeyPath = input.ClientCertKeyPath
updated = true updated = true
} }
if input.ClientCertPath != "" { if input.ClientCertPath != "" && input.ClientCertPath != config.ClientCertPath {
config.ClientCertPath = input.ClientCertPath config.ClientCertPath = input.ClientCertPath
updated = true 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) cert, err := tls.LoadX509KeyPair(config.ClientCertPath, config.ClientCertKeyPath)
if err != nil { if err != nil {
log.Error("Failed to load mTLS cert/key pair: ", err) 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) 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) { func parseURL(serviceName, serviceURL string) (*url.URL, error) {
parsedMgmtURL, err := url.ParseRequestURI(serviceURL) parsedMgmtURL, err := url.ParseRequestURI(serviceURL)
if err != nil { if err != nil {
@@ -930,6 +1092,84 @@ func isPreSharedKeyHidden(preSharedKey *string) bool {
return false 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 // UpdateConfig update existing configuration according to input configuration and return with the configuration
func UpdateConfig(input ConfigInput) (*Config, error) { func UpdateConfig(input ConfigInput) (*Config, error) {
configExists, err := fileExists(input.ConfigPath) 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) 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) return update(input)
} }
@@ -951,7 +1199,7 @@ func UpdateOrCreateConfig(input ConfigInput) (*Config, error) {
} }
if !configExists { if !configExists {
log.Infof("generating new config %s", input.ConfigPath) log.Infof("generating new config %s", input.ConfigPath)
cfg, err := createNewConfig(input) cfg, err := createProvisionedConfig(input)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -976,12 +1224,20 @@ func update(input ConfigInput) (*Config, error) {
return nil, err 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) updated, err := config.apply(input)
if err != nil { if err != nil {
return nil, err return nil, err
} }
if updated { if updated || identityGenerated {
if err := util.WriteJson(context.Background(), input.ConfigPath, config); err != nil { if err := util.WriteJson(context.Background(), input.ConfigPath, config); err != nil {
return nil, err return nil, err
} }
@@ -990,8 +1246,8 @@ func update(input ConfigInput) (*Config, error) {
return config, nil return config, nil
} }
// GetConfig read config file and return with Config and if it was created. Errors out if it does not exist // GetExistingConfig reads and returns the config if it exists on disk. Fails otherwise.
func GetConfig(configPath string) (*Config, error) { func GetExistingConfig(configPath string) (*Config, error) {
return readConfig(configPath, false) return readConfig(configPath, false)
} }
@@ -1074,17 +1330,27 @@ func UpdateOldManagementURL(ctx context.Context, config *Config, configPath stri
return newConfig, nil 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) { 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 // ReadConfigOrDefault reads the profile config at configPath, or resolves the
func ReadConfig(configPath string) (*Config, error) { // 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) 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) { func readConfig(configPath string, createIfMissing bool) (*Config, error) {
configExists, err := fileExists(configPath) configExists, err := fileExists(configPath)
if err != nil { if err != nil {
@@ -1102,12 +1368,8 @@ func readConfig(configPath string, createIfMissing bool) (*Config, error) {
return nil, err return nil, err
} }
// initialize through apply() without changes // initialize through apply() without changes
if changed, err := config.apply(ConfigInput{}); err != nil { if _, err := config.apply(ConfigInput{}); err != nil {
return nil, err return nil, err
} else if changed {
if err = WriteOutConfig(configPath, config); err != nil {
return nil, err
}
} }
return config, nil 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) return nil, fmt.Errorf("config file %s does not exist", configPath)
} }
cfg, err := createNewConfig(ConfigInput{ConfigPath: configPath}) return createNewConfig(ConfigInput{ConfigPath: configPath})
if err != nil {
return nil, err
}
err = WriteOutConfig(configPath, cfg)
return cfg, err
} }
// WriteOutConfig write put the prepared config to the given path // WriteOutConfig write put the prepared config to the given path
@@ -1144,7 +1400,7 @@ func DirectUpdateOrCreateConfig(input ConfigInput) (*Config, error) {
} }
if !configExists { if !configExists {
log.Infof("generating new config %s", input.ConfigPath) log.Infof("generating new config %s", input.ConfigPath)
cfg, err := createNewConfig(input) cfg, err := createProvisionedConfig(input)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -1171,12 +1427,18 @@ func directUpdate(input ConfigInput) (*Config, error) {
return nil, err 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) updated, err := config.apply(input)
if err != nil { if err != nil {
return nil, err return nil, err
} }
if updated { if updated || identityGenerated {
if err := util.DirectWriteJson(context.Background(), input.ConfigPath, config); err != nil { if err := util.DirectWriteJson(context.Background(), input.ConfigPath, config); err != nil {
return nil, err return nil, err
} }
@@ -1198,7 +1460,16 @@ func ConfigToJSON(config *Config) (string, error) {
// ConfigFromJSON deserializes a JSON string to a Config struct. // ConfigFromJSON deserializes a JSON string to a Config struct.
// This is useful for restoring config from alternative storage mechanisms. // 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) { func ConfigFromJSON(jsonStr string) (*Config, error) {
config := &Config{} config := &Config{}
err := json.Unmarshal([]byte(jsonStr), config) err := json.Unmarshal([]byte(jsonStr), config)
@@ -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")
}
@@ -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")
})
})
}
@@ -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")
})
}
@@ -196,7 +196,7 @@ func TestWireguardPortZeroExplicit(t *testing.T) {
assert.Equal(t, 0, config.WgPort, "WgPort should be 0 when explicitly set by user") assert.Equal(t, 0, config.WgPort, "WgPort should be 0 when explicitly set by user")
// Verify it persists // Verify it persists
readConfig, err := GetConfig(configPath) readConfig, err := GetExistingConfig(configPath)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, 0, readConfig.WgPort, "WgPort should remain 0 after reading from file") assert.Equal(t, 0, readConfig.WgPort, "WgPort should remain 0 after reading from file")
} }
@@ -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")
}
+25 -8
View File
@@ -313,7 +313,11 @@ func (s *ServiceManager) AddProfile(displayName, username string) (*Profile, err
} }
profPath := filepath.Join(configDir, id.String()+".json") 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 { if err != nil {
return nil, fmt.Errorf("failed to create new config: %w", err) 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 }, 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 { func (s *ServiceManager) RenameProfile(id ID, username string, newName string) error {
displayName, err := sanitizeDisplayName(newName) displayName, err := sanitizeDisplayName(newName)
if err != nil { if err != nil {
@@ -356,17 +373,17 @@ func (s *ServiceManager) RenameProfile(id ID, username string, newName string) e
return ErrProfileNotFound 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 { if err != nil {
return err return fmt.Errorf("read profile config: %w", err)
}
var cfg Config
if err := json.Unmarshal(data, &cfg); err != nil {
return err
} }
cfg.Name = displayName 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 fmt.Errorf("failed to write profile name: %w", err)
} }
return nil return nil
@@ -228,3 +228,27 @@ func TestRemoveProfile_DeletesStateFile(t *testing.T) {
assert.True(t, errors.Is(err, os.ErrNotExist), "state file should be removed") 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)
})
}
@@ -19,8 +19,7 @@ type IPForwardingState struct {
// routingV4/routingV6 track whether the routing path currently holds a // routingV4/routingV6 track whether the routing path currently holds a
// reference, so repeated EnableRouting calls (one per network-map update) // reference, so repeated EnableRouting calls (one per network-map update)
// hold at most one reference per family and an unpaired DisableRouting // hold at most one reference per family.
// can't release references held by DNAT rules.
routingV4 bool routingV4 bool
routingV6 bool routingV6 bool
@@ -95,31 +94,6 @@ func (f *IPForwardingState) ReleaseRouting() error {
return nil 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 { func (f *IPForwardingState) requestV4() error {
if f.v4Count == 0 { if f.v4Count == 0 {
if err := systemops.EnableV4IPForwarding(); err != nil { if err := systemops.EnableV4IPForwarding(); err != nil {
@@ -10,8 +10,7 @@ import (
) )
// TestRequestRoutingV6ToV4Transition verifies that a v4-only routing request // TestRequestRoutingV6ToV4Transition verifies that a v4-only routing request
// releases a previously held routing-owned v6 reference without touching // releases a previously held routing-owned v6 reference.
// references held by DNAT rules.
func TestRequestRoutingV6ToV4Transition(t *testing.T) { func TestRequestRoutingV6ToV4Transition(t *testing.T) {
f := NewIPForwardingState("wt-fwd-test") 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, 1, v4, "v4 reference kept")
assert.Equal(t, 0, v6, "routing-owned v6 reference released") 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") require.NoError(t, f.ReleaseRouting(), "release routing")
v4, v6 = f.Counts() v4, v6 = f.Counts()
assert.Equal(t, 0, v4, "all v4 references released") assert.Equal(t, 0, v4, "all v4 references released")
+4
View File
@@ -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). // SetConfigFromJSON stores the JSON config that later loads resolve instead of the config file (tvOS).
func (c *Client) SetConfigFromJSON(jsonStr string) error { 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 { if _, err := profilemanager.ConfigFromJSON(jsonStr); err != nil {
log.Errorf("SetConfigFromJSON: failed to parse config JSON: %v", err) log.Errorf("SetConfigFromJSON: failed to parse config JSON: %v", err)
return err return err
+29
View File
@@ -379,6 +379,35 @@ func (a *Auth) SetConfigFromJSON(jsonStr string) error {
} }
func (a *Auth) setBaseConfig(base *profilemanager.Config) 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) overlaid, err := copyConfig(base)
if err != nil { if err != nil {
return err return err
+7 -7
View File
@@ -49,7 +49,7 @@ func (p *Preferences) GetManagementURL() (string, error) {
return p.configInput.ManagementURL, nil return p.configInput.ManagementURL, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return "", err return "", err
} }
@@ -67,7 +67,7 @@ func (p *Preferences) GetAdminURL() (string, error) {
return p.configInput.AdminURL, nil return p.configInput.AdminURL, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return "", err return "", err
} }
@@ -89,7 +89,7 @@ func (p *Preferences) HasPreSharedKey() (bool, error) {
return *p.configInput.PreSharedKey != "", nil return *p.configInput.PreSharedKey != "", nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -115,7 +115,7 @@ func (p *Preferences) GetRosenpassEnabled() (bool, error) {
return *p.configInput.RosenpassEnabled, nil return *p.configInput.RosenpassEnabled, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -136,7 +136,7 @@ func (p *Preferences) GetRosenpassPermissive() (bool, error) {
return *p.configInput.RosenpassPermissive, nil return *p.configInput.RosenpassPermissive, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -149,7 +149,7 @@ func (p *Preferences) GetDisableIPv6() (bool, error) {
return *p.configInput.DisableIPv6, nil return *p.configInput.DisableIPv6, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -168,7 +168,7 @@ func (p *Preferences) GetRemoteJobsAllowed() (bool, error) {
return *p.configInput.RemoteJobsAllowed, nil return *p.configInput.RemoteJobsAllowed, nil
} }
cfg, err := profilemanager.ReadConfig(p.configInput.ConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(p.configInput.ConfigPath)
if err != nil { if err != nil {
return false, err return false, err
} }
+97
View File
@@ -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")
}
+4 -1
View File
@@ -192,7 +192,10 @@ func (pm *ProfileManager) LogoutProfile(id string) error {
return fmt.Errorf("profile %q does not exist", id) 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 { if err != nil {
return fmt.Errorf("read profile config: %w", err) return fmt.Errorf("read profile config: %w", err)
} }
+9
View File
@@ -76,6 +76,14 @@
<util:CloseApplication Id="CloseNetBird" CloseMessage="no" Target="netbird.exe" RebootPrompt="no" /> <util:CloseApplication Id="CloseNetBird" CloseMessage="no" Target="netbird.exe" RebootPrompt="no" />
<util:CloseApplication Id="CloseNetBirdUI" CloseMessage="no" Target="netbird-ui.exe" RebootPrompt="no" TerminateProcess="0" /> <util:CloseApplication Id="CloseNetBirdUI" CloseMessage="no" Target="netbird-ui.exe" RebootPrompt="no" TerminateProcess="0" />
<!-- Kill the UI before InstallValidate, otherwise its Restart Manager
check sees netbird-ui.exe in use and schedules the replacement for
the next reboot (3010). Keep it immediate and before InstallValidate.
CloseNetBirdUI above stays as a deferred LocalSystem fallback for UIs
in other sessions that this action cannot terminate. -->
<SetProperty Id="WixQuietExecCmdLine" Value="&quot;[System64Folder]taskkill.exe&quot; /F /IM netbird-ui.exe" Before="KillNetBirdUI" Sequence="execute" />
<CustomAction Id="KillNetBirdUI" BinaryRef="Wix4UtilCA_$(sys.BUILDARCHSHORT)" DllEntry="WixQuietExec" Execute="immediate" Return="ignore" />
<!-- WebView2 evergreen runtime detection. <!-- WebView2 evergreen runtime detection.
Probe both the per-machine and per-user EdgeUpdate keys; if either Probe both the per-machine and per-user EdgeUpdate keys; if either
reports a non-empty `pv` value the runtime is already installed reports a non-empty `pv` value the runtime is already installed
@@ -104,6 +112,7 @@
Return="check" /> Return="check" />
<InstallExecuteSequence> <InstallExecuteSequence>
<Custom Action="KillNetBirdUI" Before="InstallValidate" />
<Custom Action="InstallWebView2" Before="InstallFinalize" <Custom Action="InstallWebView2" Before="InstallFinalize"
Condition="NOT WEBVIEW2_VERSION_HKLM AND NOT WEBVIEW2_VERSION_HKCU AND NOT REMOVE" /> Condition="NOT WEBVIEW2_VERSION_HKLM AND NOT WEBVIEW2_VERSION_HKCU AND NOT REMOVE" />
</InstallExecuteSequence> </InstallExecuteSequence>
+32 -23
View File
@@ -2176,17 +2176,20 @@ func (x *SSHServerState) GetSessions() []*SSHSessionInfo {
// FullStatus contains the full state held by the Status instance // FullStatus contains the full state held by the Status instance
type FullStatus struct { type FullStatus struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
ManagementState *ManagementState `protobuf:"bytes,1,opt,name=managementState,proto3" json:"managementState,omitempty"` ManagementState *ManagementState `protobuf:"bytes,1,opt,name=managementState,proto3" json:"managementState,omitempty"`
SignalState *SignalState `protobuf:"bytes,2,opt,name=signalState,proto3" json:"signalState,omitempty"` SignalState *SignalState `protobuf:"bytes,2,opt,name=signalState,proto3" json:"signalState,omitempty"`
LocalPeerState *LocalPeerState `protobuf:"bytes,3,opt,name=localPeerState,proto3" json:"localPeerState,omitempty"` LocalPeerState *LocalPeerState `protobuf:"bytes,3,opt,name=localPeerState,proto3" json:"localPeerState,omitempty"`
Peers []*PeerState `protobuf:"bytes,4,rep,name=peers,proto3" json:"peers,omitempty"` Peers []*PeerState `protobuf:"bytes,4,rep,name=peers,proto3" json:"peers,omitempty"`
Relays []*RelayState `protobuf:"bytes,5,rep,name=relays,proto3" json:"relays,omitempty"` Relays []*RelayState `protobuf:"bytes,5,rep,name=relays,proto3" json:"relays,omitempty"`
DnsServers []*NSGroupState `protobuf:"bytes,6,rep,name=dns_servers,json=dnsServers,proto3" json:"dns_servers,omitempty"` DnsServers []*NSGroupState `protobuf:"bytes,6,rep,name=dns_servers,json=dnsServers,proto3" json:"dns_servers,omitempty"`
NumberOfForwardingRules int32 `protobuf:"varint,8,opt,name=NumberOfForwardingRules,proto3" json:"NumberOfForwardingRules,omitempty"` // Unused; the ingress port-forwarding feature was discontinued.
Events []*SystemEvent `protobuf:"bytes,7,rep,name=events,proto3" json:"events,omitempty"` //
LazyConnectionEnabled bool `protobuf:"varint,9,opt,name=lazyConnectionEnabled,proto3" json:"lazyConnectionEnabled,omitempty"` // Deprecated: Marked as deprecated in daemon.proto.
SshServerState *SSHServerState `protobuf:"bytes,10,opt,name=sshServerState,proto3" json:"sshServerState,omitempty"` NumberOfForwardingRules int32 `protobuf:"varint,8,opt,name=NumberOfForwardingRules,proto3" json:"NumberOfForwardingRules,omitempty"`
Events []*SystemEvent `protobuf:"bytes,7,rep,name=events,proto3" json:"events,omitempty"`
LazyConnectionEnabled bool `protobuf:"varint,9,opt,name=lazyConnectionEnabled,proto3" json:"lazyConnectionEnabled,omitempty"`
SshServerState *SSHServerState `protobuf:"bytes,10,opt,name=sshServerState,proto3" json:"sshServerState,omitempty"`
// networksRevision bumps whenever the set of routed networks (route and // networksRevision bumps whenever the set of routed networks (route and
// exit-node candidates) or their selected state changes. The UI fingerprints // exit-node candidates) or their selected state changes. The UI fingerprints
// on it to know when to re-fetch ListNetworks via the push stream, instead // on it to know when to re-fetch ListNetworks via the push stream, instead
@@ -2268,6 +2271,7 @@ func (x *FullStatus) GetDnsServers() []*NSGroupState {
return nil return nil
} }
// Deprecated: Marked as deprecated in daemon.proto.
func (x *FullStatus) GetNumberOfForwardingRules() int32 { func (x *FullStatus) GetNumberOfForwardingRules() int32 {
if x != nil { if x != nil {
return x.NumberOfForwardingRules return x.NumberOfForwardingRules
@@ -2600,7 +2604,10 @@ func (x *Network) GetResolvedIPs() map[string]*IPList {
return nil return nil
} }
// ForwardingRules // PortInfo, ForwardingRule and ForwardingRulesResponse are unused; the ingress
// port-forwarding feature was discontinued.
//
// Deprecated: Marked as deprecated in daemon.proto.
type PortInfo struct { type PortInfo struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
// Types that are valid to be assigned to PortSelection: // Types that are valid to be assigned to PortSelection:
@@ -2683,6 +2690,7 @@ func (*PortInfo_Port) isPortInfo_PortSelection() {}
func (*PortInfo_Range_) isPortInfo_PortSelection() {} func (*PortInfo_Range_) isPortInfo_PortSelection() {}
// Deprecated: Marked as deprecated in daemon.proto.
type ForwardingRule struct { type ForwardingRule struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
Protocol string `protobuf:"bytes,1,opt,name=protocol,proto3" json:"protocol,omitempty"` Protocol string `protobuf:"bytes,1,opt,name=protocol,proto3" json:"protocol,omitempty"`
@@ -2759,6 +2767,7 @@ func (x *ForwardingRule) GetTranslatedPort() *PortInfo {
return nil return nil
} }
// Deprecated: Marked as deprecated in daemon.proto.
type ForwardingRulesResponse struct { type ForwardingRulesResponse struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
Rules []*ForwardingRule `protobuf:"bytes,1,rep,name=rules,proto3" json:"rules,omitempty"` Rules []*ForwardingRule `protobuf:"bytes,1,rep,name=rules,proto3" json:"rules,omitempty"`
@@ -7303,7 +7312,7 @@ const file_daemon_proto_rawDesc = "" +
"\fportForwards\x18\x05 \x03(\tR\fportForwards\"^\n" + "\fportForwards\x18\x05 \x03(\tR\fportForwards\"^\n" +
"\x0eSSHServerState\x12\x18\n" + "\x0eSSHServerState\x12\x18\n" +
"\aenabled\x18\x01 \x01(\bR\aenabled\x122\n" + "\aenabled\x18\x01 \x01(\bR\aenabled\x122\n" +
"\bsessions\x18\x02 \x03(\v2\x16.daemon.SSHSessionInfoR\bsessions\"\xdb\x04\n" + "\bsessions\x18\x02 \x03(\v2\x16.daemon.SSHSessionInfoR\bsessions\"\xdf\x04\n" +
"\n" + "\n" +
"FullStatus\x12A\n" + "FullStatus\x12A\n" +
"\x0fmanagementState\x18\x01 \x01(\v2\x17.daemon.ManagementStateR\x0fmanagementState\x125\n" + "\x0fmanagementState\x18\x01 \x01(\v2\x17.daemon.ManagementStateR\x0fmanagementState\x125\n" +
@@ -7312,8 +7321,8 @@ const file_daemon_proto_rawDesc = "" +
"\x05peers\x18\x04 \x03(\v2\x11.daemon.PeerStateR\x05peers\x12*\n" + "\x05peers\x18\x04 \x03(\v2\x11.daemon.PeerStateR\x05peers\x12*\n" +
"\x06relays\x18\x05 \x03(\v2\x12.daemon.RelayStateR\x06relays\x125\n" + "\x06relays\x18\x05 \x03(\v2\x12.daemon.RelayStateR\x06relays\x125\n" +
"\vdns_servers\x18\x06 \x03(\v2\x14.daemon.NSGroupStateR\n" + "\vdns_servers\x18\x06 \x03(\v2\x14.daemon.NSGroupStateR\n" +
"dnsServers\x128\n" + "dnsServers\x12<\n" +
"\x17NumberOfForwardingRules\x18\b \x01(\x05R\x17NumberOfForwardingRules\x12+\n" + "\x17NumberOfForwardingRules\x18\b \x01(\x05B\x02\x18\x01R\x17NumberOfForwardingRules\x12+\n" +
"\x06events\x18\a \x03(\v2\x13.daemon.SystemEventR\x06events\x124\n" + "\x06events\x18\a \x03(\v2\x13.daemon.SystemEventR\x06events\x124\n" +
"\x15lazyConnectionEnabled\x18\t \x01(\bR\x15lazyConnectionEnabled\x12>\n" + "\x15lazyConnectionEnabled\x18\t \x01(\bR\x15lazyConnectionEnabled\x12>\n" +
"\x0esshServerState\x18\n" + "\x0esshServerState\x18\n" +
@@ -7339,22 +7348,22 @@ const file_daemon_proto_rawDesc = "" +
"\vresolvedIPs\x18\x05 \x03(\v2 .daemon.Network.ResolvedIPsEntryR\vresolvedIPs\x1aN\n" + "\vresolvedIPs\x18\x05 \x03(\v2 .daemon.Network.ResolvedIPsEntryR\vresolvedIPs\x1aN\n" +
"\x10ResolvedIPsEntry\x12\x10\n" + "\x10ResolvedIPsEntry\x12\x10\n" +
"\x03key\x18\x01 \x01(\tR\x03key\x12$\n" + "\x03key\x18\x01 \x01(\tR\x03key\x12$\n" +
"\x05value\x18\x02 \x01(\v2\x0e.daemon.IPListR\x05value:\x028\x01\"\x92\x01\n" + "\x05value\x18\x02 \x01(\v2\x0e.daemon.IPListR\x05value:\x028\x01\"\x96\x01\n" +
"\bPortInfo\x12\x14\n" + "\bPortInfo\x12\x14\n" +
"\x04port\x18\x01 \x01(\rH\x00R\x04port\x12.\n" + "\x04port\x18\x01 \x01(\rH\x00R\x04port\x12.\n" +
"\x05range\x18\x02 \x01(\v2\x16.daemon.PortInfo.RangeH\x00R\x05range\x1a/\n" + "\x05range\x18\x02 \x01(\v2\x16.daemon.PortInfo.RangeH\x00R\x05range\x1a/\n" +
"\x05Range\x12\x14\n" + "\x05Range\x12\x14\n" +
"\x05start\x18\x01 \x01(\rR\x05start\x12\x10\n" + "\x05start\x18\x01 \x01(\rR\x05start\x12\x10\n" +
"\x03end\x18\x02 \x01(\rR\x03endB\x0f\n" + "\x03end\x18\x02 \x01(\rR\x03end:\x02\x18\x01B\x0f\n" +
"\rportSelection\"\x80\x02\n" + "\rportSelection\"\x84\x02\n" +
"\x0eForwardingRule\x12\x1a\n" + "\x0eForwardingRule\x12\x1a\n" +
"\bprotocol\x18\x01 \x01(\tR\bprotocol\x12:\n" + "\bprotocol\x18\x01 \x01(\tR\bprotocol\x12:\n" +
"\x0fdestinationPort\x18\x02 \x01(\v2\x10.daemon.PortInfoR\x0fdestinationPort\x12,\n" + "\x0fdestinationPort\x18\x02 \x01(\v2\x10.daemon.PortInfoR\x0fdestinationPort\x12,\n" +
"\x11translatedAddress\x18\x03 \x01(\tR\x11translatedAddress\x12.\n" + "\x11translatedAddress\x18\x03 \x01(\tR\x11translatedAddress\x12.\n" +
"\x12translatedHostname\x18\x04 \x01(\tR\x12translatedHostname\x128\n" + "\x12translatedHostname\x18\x04 \x01(\tR\x12translatedHostname\x128\n" +
"\x0etranslatedPort\x18\x05 \x01(\v2\x10.daemon.PortInfoR\x0etranslatedPort\"G\n" + "\x0etranslatedPort\x18\x05 \x01(\v2\x10.daemon.PortInfoR\x0etranslatedPort:\x02\x18\x01\"K\n" +
"\x17ForwardingRulesResponse\x12,\n" + "\x17ForwardingRulesResponse\x12,\n" +
"\x05rules\x18\x01 \x03(\v2\x16.daemon.ForwardingRuleR\x05rules\"\x84\x02\n" + "\x05rules\x18\x01 \x03(\v2\x16.daemon.ForwardingRuleR\x05rules:\x02\x18\x01\"\x84\x02\n" +
"\x12DebugBundleRequest\x12\x1c\n" + "\x12DebugBundleRequest\x12\x1c\n" +
"\tanonymize\x18\x01 \x01(\bR\tanonymize\x12\x1e\n" + "\tanonymize\x18\x01 \x01(\bR\tanonymize\x12\x1e\n" +
"\n" + "\n" +
@@ -7705,7 +7714,7 @@ const file_daemon_proto_rawDesc = "" +
"\n" + "\n" +
"EXPOSE_UDP\x10\x03\x12\x0e\n" + "EXPOSE_UDP\x10\x03\x12\x0e\n" +
"\n" + "\n" +
"EXPOSE_TLS\x10\x042\xa3\x1c\n" + "EXPOSE_TLS\x10\x042\xa6\x1c\n" +
"\rDaemonService\x126\n" + "\rDaemonService\x126\n" +
"\x05Login\x12\x14.daemon.LoginRequest\x1a\x15.daemon.LoginResponse\"\x00\x12K\n" + "\x05Login\x12\x14.daemon.LoginRequest\x1a\x15.daemon.LoginResponse\"\x00\x12K\n" +
"\fWaitSSOLogin\x12\x1b.daemon.WaitSSOLoginRequest\x1a\x1c.daemon.WaitSSOLoginResponse\"\x00\x12-\n" + "\fWaitSSOLogin\x12\x1b.daemon.WaitSSOLoginRequest\x1a\x1c.daemon.WaitSSOLoginResponse\"\x00\x12-\n" +
@@ -7716,8 +7725,8 @@ const file_daemon_proto_rawDesc = "" +
"\tGetConfig\x12\x18.daemon.GetConfigRequest\x1a\x19.daemon.GetConfigResponse\"\x00\x12K\n" + "\tGetConfig\x12\x18.daemon.GetConfigRequest\x1a\x19.daemon.GetConfigResponse\"\x00\x12K\n" +
"\fListNetworks\x12\x1b.daemon.ListNetworksRequest\x1a\x1c.daemon.ListNetworksResponse\"\x00\x12Q\n" + "\fListNetworks\x12\x1b.daemon.ListNetworksRequest\x1a\x1c.daemon.ListNetworksResponse\"\x00\x12Q\n" +
"\x0eSelectNetworks\x12\x1d.daemon.SelectNetworksRequest\x1a\x1e.daemon.SelectNetworksResponse\"\x00\x12S\n" + "\x0eSelectNetworks\x12\x1d.daemon.SelectNetworksRequest\x1a\x1e.daemon.SelectNetworksResponse\"\x00\x12S\n" +
"\x10DeselectNetworks\x12\x1d.daemon.SelectNetworksRequest\x1a\x1e.daemon.SelectNetworksResponse\"\x00\x12J\n" + "\x10DeselectNetworks\x12\x1d.daemon.SelectNetworksRequest\x1a\x1e.daemon.SelectNetworksResponse\"\x00\x12M\n" +
"\x0fForwardingRules\x12\x14.daemon.EmptyRequest\x1a\x1f.daemon.ForwardingRulesResponse\"\x00\x12H\n" + "\x0fForwardingRules\x12\x14.daemon.EmptyRequest\x1a\x1f.daemon.ForwardingRulesResponse\"\x03\x88\x02\x01\x12H\n" +
"\vDebugBundle\x12\x1a.daemon.DebugBundleRequest\x1a\x1b.daemon.DebugBundleResponse\"\x00\x12H\n" + "\vDebugBundle\x12\x1a.daemon.DebugBundleRequest\x1a\x1b.daemon.DebugBundleResponse\"\x00\x12H\n" +
"\vGetLogLevel\x12\x1a.daemon.GetLogLevelRequest\x1a\x1b.daemon.GetLogLevelResponse\"\x00\x12H\n" + "\vGetLogLevel\x12\x1a.daemon.GetLogLevelRequest\x1a\x1b.daemon.GetLogLevelResponse\"\x00\x12H\n" +
"\vSetLogLevel\x12\x1a.daemon.SetLogLevelRequest\x1a\x1b.daemon.SetLogLevelResponse\"\x00\x12E\n" + "\vSetLogLevel\x12\x1a.daemon.SetLogLevelRequest\x1a\x1b.daemon.SetLogLevelResponse\"\x00\x12E\n" +
+14 -4
View File
@@ -45,7 +45,10 @@ service DaemonService {
// Deselect specific routes // Deselect specific routes
rpc DeselectNetworks(SelectNetworksRequest) returns (SelectNetworksResponse) {} rpc DeselectNetworks(SelectNetworksRequest) returns (SelectNetworksResponse) {}
rpc ForwardingRules(EmptyRequest) returns (ForwardingRulesResponse) {} // Unused; the ingress port-forwarding feature was discontinued.
rpc ForwardingRules(EmptyRequest) returns (ForwardingRulesResponse) {
option deprecated = true;
}
// DebugBundle creates a debug bundle // DebugBundle creates a debug bundle
rpc DebugBundle(DebugBundleRequest) returns (DebugBundleResponse) {} rpc DebugBundle(DebugBundleRequest) returns (DebugBundleResponse) {}
@@ -468,7 +471,8 @@ message FullStatus {
repeated PeerState peers = 4; repeated PeerState peers = 4;
repeated RelayState relays = 5; repeated RelayState relays = 5;
repeated NSGroupState dns_servers = 6; repeated NSGroupState dns_servers = 6;
int32 NumberOfForwardingRules = 8; // Unused; the ingress port-forwarding feature was discontinued.
int32 NumberOfForwardingRules = 8 [deprecated = true];
repeated SystemEvent events = 7; repeated SystemEvent events = 7;
@@ -511,8 +515,11 @@ message Network {
map<string, IPList> resolvedIPs = 5; map<string, IPList> resolvedIPs = 5;
} }
// ForwardingRules // PortInfo, ForwardingRule and ForwardingRulesResponse are unused; the ingress
// port-forwarding feature was discontinued.
message PortInfo { message PortInfo {
option deprecated = true;
oneof portSelection { oneof portSelection {
uint32 port = 1; uint32 port = 1;
Range range = 2; Range range = 2;
@@ -525,6 +532,8 @@ message PortInfo {
} }
message ForwardingRule { message ForwardingRule {
option deprecated = true;
string protocol = 1; string protocol = 1;
PortInfo destinationPort = 2; PortInfo destinationPort = 2;
string translatedAddress = 3; string translatedAddress = 3;
@@ -533,10 +542,11 @@ message ForwardingRule {
} }
message ForwardingRulesResponse { message ForwardingRulesResponse {
option deprecated = true;
repeated ForwardingRule rules = 1; repeated ForwardingRule rules = 1;
} }
// DebugBundler // DebugBundler
message DebugBundleRequest { message DebugBundleRequest {
bool anonymize = 1; bool anonymize = 1;
+5
View File
@@ -95,6 +95,8 @@ type DaemonServiceClient interface {
SelectNetworks(ctx context.Context, in *SelectNetworksRequest, opts ...grpc.CallOption) (*SelectNetworksResponse, error) SelectNetworks(ctx context.Context, in *SelectNetworksRequest, opts ...grpc.CallOption) (*SelectNetworksResponse, error)
// Deselect specific routes // Deselect specific routes
DeselectNetworks(ctx context.Context, in *SelectNetworksRequest, opts ...grpc.CallOption) (*SelectNetworksResponse, error) DeselectNetworks(ctx context.Context, in *SelectNetworksRequest, opts ...grpc.CallOption) (*SelectNetworksResponse, error)
// Deprecated: Do not use.
// Unused; the ingress port-forwarding feature was discontinued.
ForwardingRules(ctx context.Context, in *EmptyRequest, opts ...grpc.CallOption) (*ForwardingRulesResponse, error) ForwardingRules(ctx context.Context, in *EmptyRequest, opts ...grpc.CallOption) (*ForwardingRulesResponse, error)
// DebugBundle creates a debug bundle // DebugBundle creates a debug bundle
DebugBundle(ctx context.Context, in *DebugBundleRequest, opts ...grpc.CallOption) (*DebugBundleResponse, error) DebugBundle(ctx context.Context, in *DebugBundleRequest, opts ...grpc.CallOption) (*DebugBundleResponse, error)
@@ -290,6 +292,7 @@ func (c *daemonServiceClient) DeselectNetworks(ctx context.Context, in *SelectNe
return out, nil return out, nil
} }
// Deprecated: Do not use.
func (c *daemonServiceClient) ForwardingRules(ctx context.Context, in *EmptyRequest, opts ...grpc.CallOption) (*ForwardingRulesResponse, error) { func (c *daemonServiceClient) ForwardingRules(ctx context.Context, in *EmptyRequest, opts ...grpc.CallOption) (*ForwardingRulesResponse, error) {
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
out := new(ForwardingRulesResponse) out := new(ForwardingRulesResponse)
@@ -705,6 +708,8 @@ type DaemonServiceServer interface {
SelectNetworks(context.Context, *SelectNetworksRequest) (*SelectNetworksResponse, error) SelectNetworks(context.Context, *SelectNetworksRequest) (*SelectNetworksResponse, error)
// Deselect specific routes // Deselect specific routes
DeselectNetworks(context.Context, *SelectNetworksRequest) (*SelectNetworksResponse, error) DeselectNetworks(context.Context, *SelectNetworksRequest) (*SelectNetworksResponse, error)
// Deprecated: Do not use.
// Unused; the ingress port-forwarding feature was discontinued.
ForwardingRules(context.Context, *EmptyRequest) (*ForwardingRulesResponse, error) ForwardingRules(context.Context, *EmptyRequest) (*ForwardingRulesResponse, error)
// DebugBundle creates a debug bundle // DebugBundle creates a debug bundle
DebugBundle(context.Context, *DebugBundleRequest) (*DebugBundleResponse, error) DebugBundle(context.Context, *DebugBundleRequest) (*DebugBundleResponse, error)
-54
View File
@@ -1,54 +0,0 @@
package server
import (
"context"
firewall "github.com/netbirdio/netbird/client/firewall/manager"
"github.com/netbirdio/netbird/client/proto"
)
func (s *Server) ForwardingRules(context.Context, *proto.EmptyRequest) (*proto.ForwardingRulesResponse, error) {
s.mutex.Lock()
defer s.mutex.Unlock()
rules := s.statusRecorder.ForwardingRules()
responseRules := make([]*proto.ForwardingRule, 0, len(rules))
for _, rule := range rules {
respRule := &proto.ForwardingRule{
Protocol: string(rule.Protocol),
DestinationPort: portToProto(rule.DestinationPort),
TranslatedAddress: rule.TranslatedAddress.String(),
TranslatedHostname: s.hostNameByTranslateAddress(rule.TranslatedAddress.String()),
TranslatedPort: portToProto(rule.TranslatedPort),
}
responseRules = append(responseRules, respRule)
}
return &proto.ForwardingRulesResponse{Rules: responseRules}, nil
}
func (s *Server) hostNameByTranslateAddress(ip string) string {
hostName, ok := s.statusRecorder.PeerByIP(ip)
if !ok {
return ip
}
return hostName
}
func portToProto(port firewall.Port) *proto.PortInfo {
var portInfo proto.PortInfo
if !port.IsRange {
portInfo.PortSelection = &proto.PortInfo_Port{Port: uint32(port.Values[0])}
} else {
portInfo.PortSelection = &proto.PortInfo_Range_{
Range: &proto.PortInfo_Range{
Start: uint32(port.Values[0]),
End: uint32(port.Values[1]),
},
}
}
return &portInfo
}
+1 -1
View File
@@ -93,7 +93,7 @@ func TestLogin_ChangeThatBecomesPrivilegedMidRequestHasNoSideEffects(t *testing.
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, profilemanager.ID(activeProfile), active.ID, "the refused login switched the active profile anyway") require.Equal(t, profilemanager.ID(activeProfile), active.ID, "the refused login switched the active profile anyway")
stored, err := profilemanager.ReadConfig(targetPath) stored, err := profilemanager.GetExistingConfig(targetPath)
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, "https://api.netbird.io:443", stored.ManagementURL.String(), "the refused login moved the management URL") require.Equal(t, "https://api.netbird.io:443", stored.ManagementURL.String(), "the refused login moved the management URL")
} }
+6 -2
View File
@@ -7,6 +7,7 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/proto"
) )
func TestPersistLoginOverrides(t *testing.T) { func TestPersistLoginOverrides(t *testing.T) {
@@ -80,10 +81,13 @@ func TestPersistLoginOverrides(t *testing.T) {
require.NoError(t, err, "seed config") require.NoError(t, err, "seed config")
activeProf := &profilemanager.ActiveProfileState{ID: "default"} activeProf := &profilemanager.ActiveProfileState{ID: "default"}
err = persistLoginOverrides(activeProf, tt.newMgmtURL, tt.newPSK) err = persistLoginOverrides(activeProf, &proto.LoginRequest{
ManagementUrl: tt.newMgmtURL,
OptionalPreSharedKey: tt.newPSK,
})
require.NoError(t, err, "persistLoginOverrides") require.NoError(t, err, "persistLoginOverrides")
cfg, err := profilemanager.ReadConfig(profilemanager.DefaultConfigPath) cfg, err := profilemanager.ReadConfigOrDefault(profilemanager.DefaultConfigPath)
require.NoError(t, err, "read back config") require.NoError(t, err, "read back config")
require.Equal(t, tt.wantMgmtURL, cfg.ManagementURL.String(), "management URL") require.Equal(t, tt.wantMgmtURL, cfg.ManagementURL.String(), "management URL")
+1 -1
View File
@@ -129,7 +129,7 @@ func TestLogout_ForeignUserProfileDoesNotUseTheRunningConfig(t *testing.T) {
// refused with PermissionDenied. The namesake profile does not, so the // refused with PermissionDenied. The namesake profile does not, so the
// correct path gets as far as dialing its own unreachable management URL. // correct path gets as far as dialing its own unreachable management URL.
enableSSHOnProfile(t, cfgPath) enableSSHOnProfile(t, cfgPath)
running, err := profilemanager.GetConfig(cfgPath) running, err := profilemanager.GetExistingConfig(cfgPath)
require.NoError(t, err) require.NoError(t, err)
s.config = running s.config = running
s.connectClient = newDummyConnectClient(context.Background()) s.connectClient = newDummyConnectClient(context.Background())
-86
View File
@@ -180,92 +180,6 @@ func mdmManagedFieldConflicts(msg *proto.SetConfigRequest, policy *mdm.Policy) [
}) })
} }
// setConfigRequestHasConfigOverrides reports whether the SetConfigRequest
// carries ANY field that would actually mutate the persisted config.
// The CLI builds a SetConfigRequest unconditionally on every
// `netbird up` (see setupSetConfigReq in cmd/up.go) — a plain
// `netbird up` produces a request with every field at its zero value;
// the gate must skip such no-op invocations or it would always fire
// even when the user did not pass any --flag. Returns false on a nil
// msg; true when any management/admin URL, PSK, DNS/NAT list+clean
// flag, interface/port/MTU, or any optional bool/duration field is set.
func setConfigRequestHasConfigOverrides(msg *proto.SetConfigRequest) bool {
if msg == nil {
return false
}
return msg.ManagementUrl != "" ||
msg.AdminURL != "" ||
msg.OptionalPreSharedKey != nil ||
len(msg.CustomDNSAddress) > 0 ||
len(msg.NatExternalIPs) > 0 || msg.CleanNATExternalIPs ||
len(msg.ExtraIFaceBlacklist) > 0 ||
len(msg.DnsLabels) > 0 || msg.CleanDNSLabels ||
msg.DnsRouteInterval != nil ||
msg.RosenpassEnabled != nil ||
msg.RosenpassPermissive != nil ||
msg.InterfaceName != nil ||
msg.WireguardPort != nil ||
msg.Mtu != nil ||
msg.DisableAutoConnect != nil ||
msg.ServerSSHAllowed != nil ||
msg.RemoteJobsAllowed != nil ||
msg.NetworkMonitor != nil ||
msg.DisableClientRoutes != nil ||
msg.DisableServerRoutes != nil ||
msg.DisableDns != nil ||
msg.DisableFirewall != nil ||
msg.BlockLanAccess != nil ||
msg.DisableNotifications != nil ||
msg.BlockInbound != nil ||
msg.DisableIpv6 != nil ||
msg.EnableSSHRoot != nil ||
msg.EnableSSHSFTP != nil ||
msg.EnableSSHLocalPortForwarding != nil ||
msg.EnableSSHRemotePortForwarding != nil ||
msg.DisableSSHAuth != nil ||
msg.SshJWTCacheTTL != nil ||
msg.EnableLocalMetrics != nil ||
msg.LocalMetricsAddress != nil
}
// loginRequestHasConfigOverrides reports whether the LoginRequest
// carries ANY field that would mutate persisted daemon configuration
// (as opposed to pure-auth fields like setupKey, hostname, hint,
// profileName, username). Used by the Login handler to decide whether
// the `--disable-update-settings` / MDM gates must run: a re-auth that
// changes nothing about the configuration is always allowed.
func loginRequestHasConfigOverrides(msg *proto.LoginRequest) bool {
if msg == nil {
return false
}
return msg.ManagementUrl != "" ||
msg.AdminURL != "" ||
msg.PreSharedKey != "" || //nolint:staticcheck // SA1019: legacy proto field still accepted by Login
msg.OptionalPreSharedKey != nil ||
len(msg.CustomDNSAddress) > 0 ||
len(msg.NatExternalIPs) > 0 || msg.CleanNATExternalIPs ||
msg.RosenpassEnabled != nil ||
msg.InterfaceName != nil ||
msg.WireguardPort != nil ||
msg.DisableAutoConnect != nil ||
msg.ServerSSHAllowed != nil ||
msg.RemoteJobsAllowed != nil ||
msg.RosenpassPermissive != nil ||
len(msg.ExtraIFaceBlacklist) > 0 ||
msg.NetworkMonitor != nil ||
msg.DnsRouteInterval != nil ||
msg.DisableClientRoutes != nil ||
msg.DisableServerRoutes != nil ||
msg.DisableDns != nil ||
msg.DisableFirewall != nil ||
msg.BlockLanAccess != nil ||
msg.DisableNotifications != nil ||
len(msg.DnsLabels) > 0 || msg.CleanDNSLabels ||
msg.BlockInbound != nil ||
msg.EnableLocalMetrics != nil ||
msg.LocalMetricsAddress != nil
}
// loginRequestMDMConflicts mirrors mdmManagedFieldConflicts but for the // loginRequestMDMConflicts mirrors mdmManagedFieldConflicts but for the
// LoginRequest surface. Same value-aware semantics: a field set to the // LoginRequest surface. Same value-aware semantics: a field set to the
// MDM-enforced value is a no-op echo, not a conflict; only a divergent // MDM-enforced value is a no-op echo, not a conflict; only a divergent
+59
View File
@@ -0,0 +1,59 @@
package server
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/internal/profilemanager"
)
// The daemon provisions the peer's identity and persists it, because a key that
// stayed in memory would come back different on the next start and register a
// second peer. Provisioning is idempotent: a profile that already has an
// identity keeps the one on disk.
func TestProvisionProfileIdentity(t *testing.T) {
origDir := profilemanager.DefaultConfigPathDir
origPath := profilemanager.DefaultConfigPath
t.Cleanup(func() {
profilemanager.DefaultConfigPathDir = origDir
profilemanager.DefaultConfigPath = origPath
})
dir := t.TempDir()
profilemanager.DefaultConfigPathDir = dir
profilemanager.DefaultConfigPath = filepath.Join(dir, "default.json")
activeProf := &profilemanager.ActiveProfileState{ID: "default"}
t.Run("a profile with no file is provisioned and written", func(t *testing.T) {
_, err := os.Stat(profilemanager.DefaultConfigPath)
require.True(t, os.IsNotExist(err), "the fixture starts without a config file")
config, existed, err := provisionProfileIdentity(activeProf)
require.NoError(t, err)
require.False(t, existed, "the file was reported as pre-existing")
require.NotEmpty(t, config.PrivateKey)
stored, err := profilemanager.GetExistingConfig(profilemanager.DefaultConfigPath)
require.NoError(t, err, "provisioning did not write the config out")
require.Equal(t, config.PrivateKey, stored.PrivateKey, "the persisted identity is not the one returned")
require.NotEmpty(t, stored.SSHKey)
})
t.Run("a second call keeps the identity on disk", func(t *testing.T) {
before, err := profilemanager.GetExistingConfig(profilemanager.DefaultConfigPath)
require.NoError(t, err)
config, existed, err := provisionProfileIdentity(activeProf)
require.NoError(t, err)
require.True(t, existed)
require.Equal(t, before.PrivateKey, config.PrivateKey, "provisioning minted a second identity")
after, err := profilemanager.GetExistingConfig(profilemanager.DefaultConfigPath)
require.NoError(t, err)
require.Equal(t, before.PrivateKey, after.PrivateKey, "provisioning rewrote the stored identity")
})
}
+140 -64
View File
@@ -58,8 +58,13 @@ const (
// JWT token cache TTL for the client daemon (disabled by default) // JWT token cache TTL for the client daemon (disabled by default)
defaultJWTCacheTTL = 0 defaultJWTCacheTTL = 0
errRestoreResidualState = "failed to restore residual state: %v" errRestoreResidualState = "failed to restore residual state: %v"
errProfilesDisabled = "profiles are disabled, you cannot use this feature without profiles enabled" errProfilesDisabled = "profiles are disabled, you cannot use this feature without profiles enabled"
// errUpdateSettingsDisabled is returned with codes.FailedPrecondition, not
// codes.Unavailable: the daemon answered, and it refused. Unavailable means
// "the daemon cannot serve this", which is why the CLI downgrades it to a
// warning and the GUI reads it as an unreachable daemon — both wrong for a
// refusal the caller has to act on.
errUpdateSettingsDisabled = "update settings are disabled, you cannot use this feature without update settings enabled" errUpdateSettingsDisabled = "update settings are disabled, you cannot use this feature without update settings enabled"
errNetworksDisabled = "network selection is disabled by the administrator" errNetworksDisabled = "network selection is disabled by the administrator"
) )
@@ -510,16 +515,27 @@ func (s *Server) SetConfig(callerCtx context.Context, msg *proto.SetConfigReques
s.mutex.Lock() s.mutex.Lock()
defer s.mutex.Unlock() defer s.mutex.Unlock()
// Skip the update-settings gate when the request carries no actual stored, err := s.storedProfileConfig(msg.ProfileName, msg.Username)
// overrides: the CLI builds a SetConfigRequest unconditionally on if err != nil {
// every `netbird up` (setupSetConfigReq in cmd/up.go), so a plain return nil, err
// `netbird up` would otherwise always trip the gate and surface a }
// misleading "setConfig method is not available" warning, even when
// the user did not pass any config flag. config, err := s.setConfigInputFromRequest(msg)
if setConfigRequestHasConfigOverrides(msg) { if err != nil {
if s.checkUpdateSettingsDisabled() { return nil, err
return nil, gstatus.Errorf(codes.Unavailable, errUpdateSettingsDisabled) }
}
// Update-settings gate: refuse the request only when it would actually
// change a persisted setting. The CLI builds a SetConfigRequest
// unconditionally on every `netbird up` (setupSetConfigReq in
// cmd/up.go) and fills it from its flags and environment, so a service
// or container that restates the configuration it already runs with
// must pass the gate. Deciding this on field presence alone refused
// those callers, and — through the identical gate in Login — refused
// their login too, which left a client configured by environment
// (NB_MANAGEMENT_URL and friends) unable to come up at all.
if s.checkUpdateSettingsDisabled() && configChangeRequested(stored, config) {
return nil, gstatus.Errorf(codes.FailedPrecondition, errUpdateSettingsDisabled)
} }
// MDM gate: refuse the whole request if any of its fields is enforced // MDM gate: refuse the whole request if any of its fields is enforced
@@ -531,19 +547,10 @@ func (s *Server) SetConfig(callerCtx context.Context, msg *proto.SetConfigReques
return nil, err return nil, err
} }
stored, err := s.storedProfileConfig(msg.ProfileName, msg.Username)
if err != nil {
return nil, err
}
if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromSetConfig(msg)); err != nil { if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromSetConfig(msg)); err != nil {
return nil, err return nil, err
} }
config, err := s.setConfigInputFromRequest(msg)
if err != nil {
return nil, err
}
updatedConf, err := profilemanager.UpdateConfig(config) updatedConf, err := profilemanager.UpdateConfig(config)
if err != nil { if err != nil {
log.Errorf("failed to update profile config: %v", err) log.Errorf("failed to update profile config: %v", err)
@@ -659,37 +666,45 @@ func (s *Server) setConfigInputFromRequest(msg *proto.SetConfigRequest) (profile
// Login uses setup key to prepare configuration for the daemon. // Login uses setup key to prepare configuration for the daemon.
func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*proto.LoginResponse, error) { func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*proto.LoginResponse, error) {
activeProf, err := s.profileManager.GetActiveProfileState()
if err != nil {
return nil, fmt.Errorf("failed to get active profile state: %w", err)
}
// The stored config of the profile this request targets backs all three
// gates below. It is read before anything changes daemon state, so a
// refused login neither switches the profile nor cancels a login already
// in progress, and it is the profile the switch further down would
// activate.
stored, err := s.storedLoginConfig(activeProf, msg)
if err != nil {
return nil, err
}
// Config-override gates. LoginRequest carries the same surface as // Config-override gates. LoginRequest carries the same surface as
// SetConfigRequest (managementUrl, PSK, ssh/rosenpass/port toggles, // SetConfigRequest (managementUrl, PSK, ssh/rosenpass/port toggles,
// ...), so the same protections must apply. Without these the CLI // ...), so the same protections must apply. Without these the CLI
// command `netbird up --management-url=X` (which falls through to // command `netbird up --management-url=X` (which falls through to
// Login when SetConfig is rejected — see cmd/up.go) would silently // Login when SetConfig is rejected — see cmd/up.go) would silently
// bypass `--disable-update-settings` and any MDM policy. // bypass `--disable-update-settings` and any MDM policy.
if loginRequestHasConfigOverrides(msg) { //
if s.checkUpdateSettingsDisabled() { // The update-settings gate is value-aware, as in SetConfig: it looks at
return nil, gstatus.Errorf(codes.Unavailable, errUpdateSettingsDisabled) // what a login would actually persist (loginOverridesInput) and refuses
} // only a real divergence from the stored config. A login that restates
policy := s.mdmLoader.Load() // the values already on disk changes nothing, so it must go through —
if err := rejectMDMManagedFieldConflicts(loginRequestMDMConflicts(msg, policy)); err != nil { // that is what keeps a re-login, or a container restart carrying
return nil, err // NB_MANAGEMENT_URL, working with the kill switch on.
} if s.checkUpdateSettingsDisabled() && configChangeRequested(stored, loginOverridesInput(msg)) {
return nil, gstatus.Errorf(codes.FailedPrecondition, errUpdateSettingsDisabled)
} }
activeProf, err := s.profileManager.GetActiveProfileState() policy := s.mdmLoader.Load()
if err != nil { if err := rejectMDMManagedFieldConflicts(loginRequestMDMConflicts(msg, policy)); err != nil {
log.Errorf("failed to get active profile state: %v", err) return nil, err
return nil, fmt.Errorf("failed to get active profile state: %w", err)
} }
// Privilege gate: same restrictions as SetConfig, since LoginRequest can carry // Privilege gate: same restrictions as SetConfig, since LoginRequest can carry
// the same fields. It runs before anything here changes daemon state, so a // the same fields.
// refused login neither switches the profile nor cancels a login already in
// progress, and it reads the profile the request targets, which is the one the
// switch below would activate.
stored, err := s.storedLoginConfig(activeProf, msg)
if err != nil {
return nil, err
}
if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromLogin(msg)); err != nil { if err := requirePrivilegeForConfigChange(callerCtx, stored, privilegedChangeFromLogin(msg)); err != nil {
return nil, err return nil, err
} }
@@ -1281,6 +1296,10 @@ func (s *Server) storedLoginConfig(activeProf *profilemanager.ActiveProfileState
// storedConfigAtPath reads a profile config file, yielding nil when it does not // storedConfigAtPath reads a profile config file, yielding nil when it does not
// exist yet. // exist yet.
//
// Reading it has no side effect: profilemanager.GetExistingConfig does not
// write, so a request that the gates go on to refuse leaves the profile file as
// it found it.
func (s *Server) storedConfigAtPath(path string) (*profilemanager.Config, error) { func (s *Server) storedConfigAtPath(path string) (*profilemanager.Config, error) {
if _, err := os.Stat(path); err != nil { if _, err := os.Stat(path); err != nil {
if os.IsNotExist(err) { if os.IsNotExist(err) {
@@ -1289,7 +1308,7 @@ func (s *Server) storedConfigAtPath(path string) (*profilemanager.Config, error)
return nil, fmt.Errorf("stat profile config: %w", err) return nil, fmt.Errorf("stat profile config: %w", err)
} }
cfg, err := profilemanager.GetConfig(path) cfg, err := profilemanager.GetExistingConfig(path)
if err != nil { if err != nil {
return nil, fmt.Errorf("read profile config: %w", err) return nil, fmt.Errorf("read profile config: %w", err)
} }
@@ -1608,8 +1627,16 @@ func (s *Server) handleActiveProfileLogout(ctx context.Context) (*proto.LogoutRe
return &proto.LogoutResponse{}, nil return &proto.LogoutResponse{}, nil
} }
// getConfig reads config file and returns Config and whether the config file already existed. Errors out if it does not exist // provisionProfileIdentity resolves the active profile's config and puts the
func (s *Server) getConfig(activeProf *profilemanager.ActiveProfileState) (*profilemanager.Config, bool, error) { // keys that identify the peer on disk, reporting whether the config file
// already existed.
//
// This is the daemon's provisioning point: the config resolved here is the one
// the peer runs with, so it needs its identity, and that has to reach disk — a
// key that stays in memory would come back different on the next start and
// re-register the peer. Reads themselves are pure, so the write is here, in
// the open, instead of hiding inside the reader.
func provisionProfileIdentity(activeProf *profilemanager.ActiveProfileState) (*profilemanager.Config, bool, error) {
cfgPath, err := activeProf.FilePath() cfgPath, err := activeProf.FilePath()
if err != nil { if err != nil {
return nil, false, fmt.Errorf("failed to get active profile file path: %w", err) return nil, false, fmt.Errorf("failed to get active profile file path: %w", err)
@@ -1620,15 +1647,38 @@ func (s *Server) getConfig(activeProf *profilemanager.ActiveProfileState) (*prof
log.Infof("active profile config existed: %t, err %v", configExisted, err) log.Infof("active profile config existed: %t, err %v", configExisted, err)
config, err := profilemanager.ReadConfig(cfgPath) config, err := profilemanager.ReadConfigOrDefault(cfgPath)
if err != nil { if err != nil {
return nil, false, fmt.Errorf("failed to get config: %w", err) return nil, false, fmt.Errorf("failed to get config: %w", err)
} }
// Apply the daemon-owned MDM policy on top of the just-resolved generated, err := config.EnsureIdentity()
// Config. profilemanager's apply() initialises the policy to if err != nil {
// empty — the Loader lives outside Config, so this overlay step return nil, false, fmt.Errorf("ensure profile identity: %w", err)
// is driven externally here. }
if generated || !configExisted {
if err := profilemanager.WriteOutConfig(cfgPath, config); err != nil {
return nil, false, fmt.Errorf("write out profile config: %w", err)
}
}
return config, configExisted, nil
}
// getConfig resolves the active profile's config, provisions its identity and
// reports whether the config file already existed.
func (s *Server) getConfig(activeProf *profilemanager.ActiveProfileState) (*profilemanager.Config, bool, error) {
config, configExisted, err := provisionProfileIdentity(activeProf)
if err != nil {
return nil, false, err
}
// Apply the daemon-owned MDM policy on top of the just-resolved Config.
// profilemanager's apply() initialises the policy to empty — the Loader
// lives outside Config, so this overlay step is driven externally here.
// After the write above, on purpose: the overlay is runtime-only and
// re-derived on every load, so the file keeps the profile's own values.
config.ApplyMDMPolicy(s.mdmLoader.Load()) config.ApplyMDMPolicy(s.mdmLoader.Load())
return config, configExisted, nil return config, configExisted, nil
@@ -1683,7 +1733,7 @@ func (s *Server) logoutFromProfile(ctx context.Context, profile *profilemanager.
cfgPath = profilemanager.DefaultConfigPath cfgPath = profilemanager.DefaultConfigPath
} }
config, err := profilemanager.GetConfig(cfgPath) config, err := profilemanager.GetExistingConfig(cfgPath)
if err != nil { if err != nil {
return fmt.Errorf("profile '%s' not found", profile.ID) return fmt.Errorf("profile '%s' not found", profile.ID)
} }
@@ -1702,6 +1752,19 @@ func (s *Server) sendLogoutRequestWithConfig(ctx context.Context, config *profil
// Privilege gate: deregistering frees this machine's key to be registered // Privilege gate: deregistering frees this machine's key to be registered
// against another management server, which is only restricted while the SSH // against another management server, which is only restricted while the SSH
// server makes that a privilege handover. // server makes that a privilege handover.
// Ahead of the privilege gate on purpose. A profile with no identity was
// never registered — a logout clears the keys in place, so logging the same
// profile out twice lands here — so there is nothing to deregister and
// nothing for the gate to protect: what it guards against is handing this
// machine's registered key to another management server. Behind the gate,
// an unprivileged caller would be refused instead, and for a profile whose
// ServerSSHAllowed is unset that is every caller, since an absent value
// counts as SSH enabled.
if config.PrivateKey == "" {
log.Infof("profile carries no identity, nothing to deregister")
return nil
}
if err := requirePrivilegeForDeregistration(ctx, config); err != nil { if err := requirePrivilegeForDeregistration(ctx, config); err != nil {
return err return err
} }
@@ -2322,7 +2385,7 @@ func (s *Server) GetConfig(ctx context.Context, req *proto.GetConfigRequest) (*p
cfgPath = profilemanager.DefaultConfigPath cfgPath = profilemanager.DefaultConfigPath
} }
cfg, err := profilemanager.GetConfig(cfgPath) cfg, err := profilemanager.GetExistingConfig(cfgPath)
if err != nil { if err != nil {
log.Errorf("failed to get active profile config: %v", err) log.Errorf("failed to get active profile config: %v", err)
return nil, fmt.Errorf("failed to get active profile config: %w", err) return nil, fmt.Errorf("failed to get active profile config: %w", err)
@@ -2785,8 +2848,6 @@ func sendTerminalNotification() error {
return wallCmd.Wait() return wallCmd.Wait()
} }
// persistLoginOverrides writes management URL and pre-shared key from a LoginRequest to the
// active profile config so that subsequent reads pick them up. Empty/nil values are ignored.
// afterLoginPreCheck is a seam for tests to run a concurrent config change // afterLoginPreCheck is a seam for tests to run a concurrent config change
// between Login's first privilege check and the authoritative one. // between Login's first privilege check and the authoritative one.
var afterLoginPreCheck func() var afterLoginPreCheck func()
@@ -2817,6 +2878,15 @@ func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto.
return nil, nil, err return nil, nil, err
} }
// The update-settings decision is re-taken here for the same reason as the
// privilege one: Login's earlier check ran outside this lock, so the stored
// config it compared against could have moved since. This one is the
// authoritative check, and it is the last read before persistLoginOverrides
// writes.
if s.checkUpdateSettingsDisabled() && configChangeRequested(stored, loginOverridesInput(msg)) {
return nil, nil, gstatus.Errorf(codes.FailedPrecondition, errUpdateSettingsDisabled)
}
s.mutex.Lock() s.mutex.Lock()
if s.actCancel != nil { if s.actCancel != nil {
s.actCancel() s.actCancel()
@@ -2843,18 +2913,28 @@ func (s *Server) authorizeAndPrepareLogin(callerCtx context.Context, msg *proto.
return nil, nil, fmt.Errorf("active profile state: %w", err) return nil, nil, fmt.Errorf("active profile state: %w", err)
} }
if err := persistLoginOverrides(activeProf, msg.ManagementUrl, msg.OptionalPreSharedKey); err != nil { if err := persistLoginOverrides(activeProf, msg); err != nil {
return nil, nil, fmt.Errorf("persist login overrides: %w", err) return nil, nil, fmt.Errorf("persist login overrides: %w", err)
} }
// Provisioning under the same lock as the decision above, and next to the
// write it guards. getConfig would otherwise mint the identity and persist
// it once this returns: between its read and its write, a SetConfig that
// had already answered its caller would be overwritten by the config this
// login read before it landed.
if _, _, err := provisionProfileIdentity(activeProf); err != nil {
return nil, nil, err
}
return ctx, activeProf, nil return ctx, activeProf, nil
} }
func persistLoginOverrides(activeProf *profilemanager.ActiveProfileState, managementURL string, preSharedKey *string) error { // persistLoginOverrides writes the config fields a login request is allowed to
if preSharedKey != nil && *preSharedKey == "" { // carry into the active profile. It shares its input builder with the
preSharedKey = nil // update-settings gate, so the gate judges exactly the fields this writes.
} func persistLoginOverrides(activeProf *profilemanager.ActiveProfileState, msg *proto.LoginRequest) error {
if managementURL == "" && preSharedKey == nil { input := loginOverridesInput(msg)
if input.ManagementURL == "" && input.PreSharedKey == nil {
return nil return nil
} }
@@ -2863,11 +2943,7 @@ func persistLoginOverrides(activeProf *profilemanager.ActiveProfileState, manage
return fmt.Errorf("active profile file path: %w", err) return fmt.Errorf("active profile file path: %w", err)
} }
input := profilemanager.ConfigInput{ input.ConfigPath = cfgPath
ConfigPath: cfgPath,
ManagementURL: managementURL,
PreSharedKey: preSharedKey,
}
if _, err := profilemanager.UpdateOrCreateConfig(input); err != nil { if _, err := profilemanager.UpdateOrCreateConfig(input); err != nil {
return fmt.Errorf("update config: %w", err) return fmt.Errorf("update config: %w", err)
} }
+3 -4
View File
@@ -10,9 +10,9 @@ import (
"testing" "testing"
"time" "time"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.opentelemetry.io/otel" "go.opentelemetry.io/otel"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator" "github.com/netbirdio/netbird/management/server/integrations/integrated_validator/validator"
@@ -36,7 +36,6 @@ import (
"github.com/netbirdio/netbird/management/server" "github.com/netbirdio/netbird/management/server"
"github.com/netbirdio/netbird/management/server/activity" "github.com/netbirdio/netbird/management/server/activity"
nbcache "github.com/netbirdio/netbird/management/server/cache" nbcache "github.com/netbirdio/netbird/management/server/cache"
"github.com/netbirdio/netbird/management/server/integrations/port_forwarding"
"github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings" "github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/store"
@@ -200,8 +199,8 @@ func startManagement(t *testing.T, signalAddr string, counter *int) (*grpc.Serve
requestBuffer := server.NewAccountRequestBuffer(context.Background(), store) requestBuffer := server.NewAccountRequestBuffer(context.Background(), store)
peersUpdateManager := update_channel.NewPeersUpdateManager(metrics) peersUpdateManager := update_channel.NewPeersUpdateManager(metrics)
networkMapController := controller.NewController(context.Background(), store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config, nil) networkMapController := controller.NewController(context.Background(), store, metrics, peersUpdateManager, 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, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManagerMock, false, cacheStore) accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, settingsMockManager, permissionsManagerMock, false, cacheStore)
if err != nil { if err != nil {
return nil, "", err return nil, "", err
} }
+1 -1
View File
@@ -290,7 +290,7 @@ func TestSetConfig_MDMReject_AllOrNothing(t *testing.T) {
// Confirm RosenpassEnabled was NOT applied even though it was not // Confirm RosenpassEnabled was NOT applied even though it was not
// in the conflict list: the request was rejected as a whole. // in the conflict list: the request was rejected as a whole.
reloaded, err := profilemanager.GetConfig(cfgPath) reloaded, err := profilemanager.GetExistingConfig(cfgPath)
require.NoError(t, err) require.NoError(t, err)
assert.False(t, reloaded.RosenpassEnabled, "non-conflicting field must not be applied when request is rejected") assert.False(t, reloaded.RosenpassEnabled, "non-conflicting field must not be applied when request is rejected")
} }
+1 -1
View File
@@ -125,7 +125,7 @@ func TestSetConfig_AllFieldsSaved(t *testing.T) {
cfgPath, err := profState.FilePath() cfgPath, err := profState.FilePath()
require.NoError(t, err) require.NoError(t, err)
cfg, err := profilemanager.GetConfig(cfgPath) cfg, err := profilemanager.GetExistingConfig(cfgPath)
require.NoError(t, err) require.NoError(t, err)
require.Equal(t, "https://new-api.netbird.io:443", cfg.ManagementURL.String()) require.Equal(t, "https://new-api.netbird.io:443", cfg.ManagementURL.String())
+1 -17
View File
@@ -331,21 +331,5 @@ func sameManagementURL(stored *url.URL, requested string) bool {
return false return false
} }
return stored.Scheme == parsed.Scheme && return profilemanager.SameServiceURL(stored, parsed)
stored.Hostname() == parsed.Hostname() &&
effectivePort(stored) == effectivePort(parsed)
}
func effectivePort(u *url.URL) string {
if port := u.Port(); port != "" {
return port
}
switch u.Scheme {
case "https":
return "443"
case "http":
return "80"
default:
return ""
}
} }
+55
View File
@@ -0,0 +1,55 @@
package server
import (
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/proto"
)
// configChangeRequested reports whether applying input would move the target
// profile away from the configuration it already persists. It is the decision
// procedure of the update-settings kill switch (--disable-update-settings /
// NB_DISABLE_UPDATE_SETTINGS / the MDM DisableUpdateSettings key): that switch
// forbids *changing* settings, so a request that restates the stored values is
// not a change and must not be refused.
//
// This has to be judged on values, not on field presence. `netbird up` rebuilds
// the whole config surface of SetConfigRequest and LoginRequest from its flags
// and environment on every invocation, so a service or container configured by
// environment restates its own configuration on every start. A presence-based
// gate refused those requests, and because Login carries the same fields it
// refused the login too — leaving such a client unable to come up at all.
//
// A dry run that cannot be evaluated fails closed: the request counts as a
// change, so a malformed field can never open the gate. The error itself is
// reported to the caller by the real update path.
func configChangeRequested(stored *profilemanager.Config, input profilemanager.ConfigInput) bool {
changed, err := stored.WouldChange(input)
if err != nil {
log.Warnf("cannot evaluate the requested config change, treating it as a change: %v", err)
return true
}
return changed
}
// loginOverridesInput builds the ConfigInput a login request persists. The
// management URL and the pre-shared key are the only config fields the daemon
// applies from a LoginRequest; everything else on that message is either pure
// auth or ignored. An empty pre-shared key is dropped rather than written, so
// a login cannot clear the stored key by omission.
//
// Both the write (persistLoginOverrides) and the update-settings gate go
// through this builder, so the gate can neither refuse a field the write
// ignores nor miss one it applies.
func loginOverridesInput(msg *proto.LoginRequest) profilemanager.ConfigInput {
preSharedKey := msg.OptionalPreSharedKey
if preSharedKey != nil && *preSharedKey == "" {
preSharedKey = nil
}
return profilemanager.ConfigInput{
ManagementURL: msg.ManagementUrl,
PreSharedKey: preSharedKey,
}
}
+390
View File
@@ -0,0 +1,390 @@
package server
import (
"context"
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/require"
"google.golang.org/grpc/codes"
gstatus "google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
"github.com/netbirdio/netbird/client/proto"
)
// The seeded profile of setupServerWithProfile is created with this management
// URL, so a request carrying it restates what the profile already holds.
const storedManagementURL = "https://api.netbird.io:443"
// A client configured by environment re-sends its whole configuration on every
// `netbird up`: the CLI fills the request from its flags and env regardless of
// what changed. With the update-settings kill switch on, such a request must
// pass — nothing about the configuration moves.
func TestSetConfig_RestatingTheStoredConfigPassesTheGate(t *testing.T) {
s, ctx, profName, username, _ := setupServerWithProfile(t)
s.updateSettingsDisabled = true
_, err := s.SetConfig(ctx, &proto.SetConfigRequest{
ProfileName: profName,
Username: username,
ManagementUrl: storedManagementURL,
})
require.NoError(t, err, "restating the stored management URL is not a settings change")
}
// The same endpoint written without its default port is the same endpoint. A
// gate that compared raw strings refused NB_MANAGEMENT_URL=https://host, which
// is how the URL is normally spelled.
func TestSetConfig_EquivalentManagementURLPassesTheGate(t *testing.T) {
s, ctx, profName, username, _ := setupServerWithProfile(t)
s.updateSettingsDisabled = true
_, err := s.SetConfig(ctx, &proto.SetConfigRequest{
ProfileName: profName,
Username: username,
ManagementUrl: "https://api.netbird.io",
})
require.NoError(t, err, "an implicit :443 is the same management URL")
}
// The kill switch still has to do its job: a request that moves a setting is
// refused, and the profile keeps the value it had.
func TestSetConfig_ChangingASettingIsRefused(t *testing.T) {
s, ctx, profName, username, cfgPath := setupServerWithProfile(t)
s.updateSettingsDisabled = true
_, err := s.SetConfig(ctx, &proto.SetConfigRequest{
ProfileName: profName,
Username: username,
ManagementUrl: "https://mgmt.elsewhere.example:443",
})
require.Error(t, err, "moving the management URL is a settings change")
require.Equal(t, codes.FailedPrecondition, gstatus.Code(err), "want the update-settings refusal, got %v", err)
cfg, err := profilemanager.GetExistingConfig(cfgPath)
require.NoError(t, err)
require.Equal(t, storedManagementURL, cfg.ManagementURL.String(), "the refused request changed the config anyway")
}
// A field whose requested value differs from the stored one is a change even
// when the rest of the request restates the configuration.
func TestSetConfig_SingleDivergingFieldIsRefused(t *testing.T) {
s, ctx, profName, username, _ := setupServerWithProfile(t)
s.updateSettingsDisabled = true
rosenpass := true
_, err := s.SetConfig(ctx, &proto.SetConfigRequest{
ProfileName: profName,
Username: username,
ManagementUrl: storedManagementURL,
RosenpassEnabled: &rosenpass,
})
require.Error(t, err, "enabling Rosenpass is a settings change")
require.Equal(t, codes.FailedPrecondition, gstatus.Code(err), "want the update-settings refusal, got %v", err)
}
// With the switch off, the same diverging request goes through: the gate must
// not leak into a daemon that never enabled it.
func TestSetConfig_ChangeAllowedWhenTheSwitchIsOff(t *testing.T) {
s, ctx, profName, username, cfgPath := setupServerWithProfile(t)
_, err := s.SetConfig(ctx, &proto.SetConfigRequest{
ProfileName: profName,
Username: username,
ManagementUrl: "https://mgmt.elsewhere.example:443",
})
require.NoError(t, err)
cfg, err := profilemanager.GetExistingConfig(cfgPath)
require.NoError(t, err)
require.Equal(t, "https://mgmt.elsewhere.example:443", cfg.ManagementURL.String())
}
// Login carries the same config surface as SetConfig, so it is gated the same
// way: a login that would move a protected setting is refused before it can
// touch daemon state.
func TestLogin_ChangingTheManagementURLIsRefused(t *testing.T) {
s, _, profName, username, cfgPath := setupServerWithProfile(t)
s.updateSettingsDisabled = true
s.rootCtx = internal.CtxInitState(context.Background())
cancelled := false
s.actCancel = func() { cancelled = true }
_, err := s.Login(userCtx(), &proto.LoginRequest{
Username: &username,
ManagementUrl: "https://mgmt.elsewhere.example:443",
})
require.Error(t, err, "moving the management URL through Login is a settings change")
require.Equal(t, codes.FailedPrecondition, gstatus.Code(err), "want the update-settings refusal, got %v", err)
// "Refused before it can touch daemon state" is the contract, so check the
// state as well as the error.
cfg, err := profilemanager.GetExistingConfig(cfgPath)
require.NoError(t, err)
require.Equal(t, storedManagementURL, cfg.ManagementURL.String(), "the refused login moved the management URL")
require.False(t, cancelled, "the refused login cancelled the login already in progress")
active, err := s.profileManager.GetActiveProfileState()
require.NoError(t, err)
require.Equal(t, profilemanager.ID(profName), active.ID, "the refused login switched the active profile")
}
// seedProfileConfig writes a profile config carrying the given management URL
// and pre-shared key into a temp dir, and returns its path.
func seedProfileConfig(t *testing.T, managementURL, preSharedKey string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "seeded.json")
_, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{
ConfigPath: path,
ManagementURL: managementURL,
PreSharedKey: &preSharedKey,
})
require.NoError(t, err, "seed profile config")
return path
}
// The decision procedure itself, over the fields a login actually persists.
// A login that restates the stored values must not be refused: that is what
// keeps a re-login, or a container restart carrying NB_MANAGEMENT_URL, working
// with the kill switch on.
func TestLoginGateDecision(t *testing.T) {
stored, err := profilemanager.GetExistingConfig(seedProfileConfig(t, storedManagementURL, "stored-key"))
require.NoError(t, err)
redacted := mdm.PreSharedKeyRedactedSentinel
empty := ""
sameKey := "stored-key"
otherKey := "other-key"
tests := []struct {
name string
msg *proto.LoginRequest
wantChanged bool
}{
{
name: "pure auth carries no config",
msg: &proto.LoginRequest{SetupKey: "ABC"},
wantChanged: false,
},
{
name: "stored management URL restated",
msg: &proto.LoginRequest{ManagementUrl: storedManagementURL},
wantChanged: false,
},
{
name: "stored management URL without its default port",
msg: &proto.LoginRequest{ManagementUrl: "https://api.netbird.io"},
wantChanged: false,
},
{
name: "different management URL",
msg: &proto.LoginRequest{ManagementUrl: "https://mgmt.elsewhere.example:443"},
wantChanged: true,
},
{
name: "stored pre-shared key restated",
msg: &proto.LoginRequest{OptionalPreSharedKey: &sameKey},
wantChanged: false,
},
{
name: "redacted pre-shared key echoed back",
msg: &proto.LoginRequest{OptionalPreSharedKey: &redacted},
wantChanged: false,
},
{
name: "empty pre-shared key is not a request to clear it",
msg: &proto.LoginRequest{OptionalPreSharedKey: &empty},
wantChanged: false,
},
{
name: "different pre-shared key",
msg: &proto.LoginRequest{OptionalPreSharedKey: &otherKey},
wantChanged: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.wantChanged, configChangeRequested(stored, loginOverridesInput(tt.msg)))
})
}
}
// A profile with no config on disk yet is judged against the config the daemon
// would create for it, so a first login that asks for the defaults is not a
// change while one that asks for a different management URL is.
func TestGateDecisionWithoutStoredConfig(t *testing.T) {
require.False(t, configChangeRequested(nil, profilemanager.ConfigInput{}),
"a request carrying nothing cannot change anything")
require.False(t, configChangeRequested(nil, profilemanager.ConfigInput{ManagementURL: profilemanager.DefaultManagementURL}),
"asking for the default management URL is what the daemon would write anyway")
require.True(t, configChangeRequested(nil, profilemanager.ConfigInput{ManagementURL: "https://mgmt.elsewhere.example:443"}),
"asking for a non-default management URL is a change")
}
// A dry run that cannot be evaluated must fail closed, or a malformed field
// would open the gate.
func TestGateDecisionFailsClosedOnAnInvalidRequest(t *testing.T) {
require.True(t, configChangeRequested(nil, profilemanager.ConfigInput{ManagementURL: "not-a-url"}),
"an unevaluable request must count as a change")
}
// The gate reads the stored config to decide, and reading it must not write it:
// a refused request has to leave the profile file byte-for-byte as it was.
// A config file missing a field the config layer fills in (MTU, here) is what
// makes the normalization write fire.
func TestSetConfig_RefusedRequestLeavesTheConfigFileUntouched(t *testing.T) {
s, ctx, profName, username, cfgPath := setupServerWithProfile(t)
s.updateSettingsDisabled = true
require.NoError(t, os.WriteFile(cfgPath, []byte(`{"WgIface":"wt0"}`), 0o600))
before, err := os.ReadFile(cfgPath)
require.NoError(t, err)
_, err = s.SetConfig(ctx, &proto.SetConfigRequest{
ProfileName: profName,
Username: username,
ManagementUrl: "https://mgmt.elsewhere.example:443",
})
require.Error(t, err)
require.Equal(t, codes.FailedPrecondition, gstatus.Code(err), "want the update-settings refusal, got %v", err)
after, err := os.ReadFile(cfgPath)
require.NoError(t, err)
require.Equal(t, string(before), string(after), "the refused request rewrote the profile config")
}
// The container case that the string comparison still broke: the management URL
// supplied through the environment is the stored one, written with a trailing
// slash.
func TestSetConfig_ManagementURLSpellingsPassTheGate(t *testing.T) {
for _, spelling := range []string{
"https://api.netbird.io",
"https://api.netbird.io/",
"https://api.netbird.io:443/",
"https://API.netbird.io:443",
} {
t.Run(spelling, func(t *testing.T) {
s, ctx, profName, username, _ := setupServerWithProfile(t)
s.updateSettingsDisabled = true
_, err := s.SetConfig(ctx, &proto.SetConfigRequest{
ProfileName: profName,
Username: username,
ManagementUrl: spelling,
})
require.NoError(t, err, "%q is the stored management URL written differently", spelling)
})
}
}
// The RPC the whole fix hangs on. Login is retried by the CLI in a backoff
// loop, so a login that restates the stored configuration — which is what a
// container configured by environment sends on every start — must get past the
// gate, or the client never comes up at all.
//
// Past the gate the handler goes on to do real work this test does not stand
// up, so the assertion is only that the refusal did not happen.
func TestLogin_RestatingTheStoredConfigPassesTheGate(t *testing.T) {
s, _, _, username, _ := setupServerWithProfile(t)
s.updateSettingsDisabled = true
s.rootCtx = internal.CtxInitState(context.Background())
// Stand in for the management round trip the handler makes once the gate
// lets it through, so this test exercises the gate and not the network:
// without it the profile's management URL is dialed for real.
s.isLoginRequiredFn = func(context.Context) (bool, error) { return false, nil }
_, err := s.Login(userCtx(), &proto.LoginRequest{
Username: &username,
ManagementUrl: storedManagementURL,
})
if err != nil {
require.NotEqual(t, codes.FailedPrecondition, gstatus.Code(err),
"the gate refused a login that changes nothing: %v", err)
require.NotContains(t, err.Error(), "update settings are disabled",
"the gate refused a login that changes nothing: %v", err)
}
}
// The value-aware decision has the same synchronization problem as the
// privileged-change one: Login's first check runs outside guardedConfigMu, so
// the stored config it compared against can move before the write. A login that
// was a no-op when it was checked must not be written once it has become a
// change.
func TestLogin_ChangeThatAppearsMidRequestIsRefused(t *testing.T) {
s, _, _, username, _ := setupServerWithProfile(t)
s.updateSettingsDisabled = true
s.rootCtx = internal.CtxInitState(context.Background())
target := "moved-under-us"
targetPath := filepath.Join(profilemanager.DefaultConfigPathDir, target+".json")
_, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{
ConfigPath: targetPath,
ManagementURL: storedManagementURL,
})
require.NoError(t, err)
cancelled := false
s.actCancel = func() { cancelled = true }
// Stand in for a concurrent writer that repoints the profile between the two
// checks, which is the interleaving the lock has to make safe. The login
// restates the URL the profile held when it was checked, so the first check
// sees a no-op and lets it through.
afterLoginPreCheck = func() {
_, err := profilemanager.UpdateOrCreateConfig(profilemanager.ConfigInput{
ConfigPath: targetPath,
ManagementURL: "https://mgmt.elsewhere.example:443",
})
require.NoError(t, err)
}
t.Cleanup(func() { afterLoginPreCheck = nil })
_, err = s.Login(userCtx(), &proto.LoginRequest{
ProfileName: &target,
Username: &username,
ManagementUrl: storedManagementURL,
})
require.Error(t, err, "the login became a settings change before it was written")
require.Equal(t, codes.FailedPrecondition, gstatus.Code(err), "want the update-settings refusal, got %v", err)
require.False(t, cancelled, "the refused login cancelled the login already in progress")
stored, err := profilemanager.GetExistingConfig(targetPath)
require.NoError(t, err)
require.Equal(t, "https://mgmt.elsewhere.example:443", stored.ManagementURL.String(),
"the refused login wrote the management URL it was asked for")
}
// Logging out a profile that was already logged out must not fail: the logout
// clears the keys in place, so the second attempt finds a profile with no
// identity, which was never registered and has nothing to deregister.
func TestLogout_ProfileWithoutAnIdentityIsANoOp(t *testing.T) {
s, _, _, _, cfgPath := setupServerWithProfile(t)
loggedOut, err := profilemanager.GetExistingConfig(cfgPath)
require.NoError(t, err)
loggedOut.PrivateKey = ""
loggedOut.SSHKey = ""
require.NoError(t, profilemanager.WriteOutConfig(cfgPath, loggedOut))
stored, err := profilemanager.GetExistingConfig(cfgPath)
require.NoError(t, err)
require.NoError(t, s.sendLogoutRequestWithConfig(privilegedTestCtx(), stored),
"logging out an identity-less profile must not fail")
// And for an unprivileged caller too: the deregistration privilege gate
// guards the handover of a registered key, so with no key there is nothing
// to guard. An unset SSH setting is what arms that gate — sshServerEnabled
// reads an absent value as enabled — so this stands in for every legacy
// profile, where behind the gate the caller would be refused.
stored.ServerSSHAllowed = nil
require.NoError(t, s.sendLogoutRequestWithConfig(userCtx(), stored),
"an unprivileged caller could not log out a profile with nothing to deregister")
}
+42 -52
View File
@@ -140,28 +140,27 @@ type SSHServerStateOutput struct {
} }
type OutputOverview struct { type OutputOverview struct {
Peers PeersStateOutput `json:"peers" yaml:"peers"` Peers PeersStateOutput `json:"peers" yaml:"peers"`
CliVersion string `json:"cliVersion" yaml:"cliVersion"` CliVersion string `json:"cliVersion" yaml:"cliVersion"`
DaemonVersion string `json:"daemonVersion" yaml:"daemonVersion"` DaemonVersion string `json:"daemonVersion" yaml:"daemonVersion"`
DaemonStatus DaemonStatus `json:"daemonStatus" yaml:"daemonStatus"` DaemonStatus DaemonStatus `json:"daemonStatus" yaml:"daemonStatus"`
ManagementState ManagementStateOutput `json:"management" yaml:"management"` ManagementState ManagementStateOutput `json:"management" yaml:"management"`
SignalState SignalStateOutput `json:"signal" yaml:"signal"` SignalState SignalStateOutput `json:"signal" yaml:"signal"`
Relays RelayStateOutput `json:"relays" yaml:"relays"` Relays RelayStateOutput `json:"relays" yaml:"relays"`
IP string `json:"netbirdIp" yaml:"netbirdIp"` IP string `json:"netbirdIp" yaml:"netbirdIp"`
IPv6 string `json:"netbirdIpv6,omitempty" yaml:"netbirdIpv6,omitempty"` IPv6 string `json:"netbirdIpv6,omitempty" yaml:"netbirdIpv6,omitempty"`
PubKey string `json:"publicKey" yaml:"publicKey"` PubKey string `json:"publicKey" yaml:"publicKey"`
KernelInterface bool `json:"usesKernelInterface" yaml:"usesKernelInterface"` KernelInterface bool `json:"usesKernelInterface" yaml:"usesKernelInterface"`
WgPort int `json:"wireguardPort" yaml:"wireguardPort"` WgPort int `json:"wireguardPort" yaml:"wireguardPort"`
FQDN string `json:"fqdn" yaml:"fqdn"` FQDN string `json:"fqdn" yaml:"fqdn"`
RosenpassEnabled bool `json:"quantumResistance" yaml:"quantumResistance"` RosenpassEnabled bool `json:"quantumResistance" yaml:"quantumResistance"`
RosenpassPermissive bool `json:"quantumResistancePermissive" yaml:"quantumResistancePermissive"` RosenpassPermissive bool `json:"quantumResistancePermissive" yaml:"quantumResistancePermissive"`
Networks []string `json:"networks" yaml:"networks"` Networks []string `json:"networks" yaml:"networks"`
NumberOfForwardingRules int `json:"forwardingRules" yaml:"forwardingRules"` NSServerGroups []NsServerGroupStateOutput `json:"dnsServers" yaml:"dnsServers"`
NSServerGroups []NsServerGroupStateOutput `json:"dnsServers" yaml:"dnsServers"` Events []SystemEventOutput `json:"events" yaml:"events"`
Events []SystemEventOutput `json:"events" yaml:"events"` LazyConnectionEnabled bool `json:"lazyConnectionEnabled" yaml:"lazyConnectionEnabled"`
LazyConnectionEnabled bool `json:"lazyConnectionEnabled" yaml:"lazyConnectionEnabled"` ProfileName string `json:"profileName" yaml:"profileName"`
ProfileName string `json:"profileName" yaml:"profileName"` SSHServerState SSHServerStateOutput `json:"sshServer" yaml:"sshServer"`
SSHServerState SSHServerStateOutput `json:"sshServer" yaml:"sshServer"`
// SessionExpiresAt is the absolute UTC instant at which the peer's SSO // SessionExpiresAt is the absolute UTC instant at which the peer's SSO
// session expires. nil when the peer is not SSO-tracked or login // session expires. nil when the peer is not SSO-tracked or login
// expiration is disabled. Pointer (rather than zero-value time.Time) so // expiration is disabled. Pointer (rather than zero-value time.Time) so
@@ -190,28 +189,27 @@ func ConvertToStatusOutputOverview(pbFullStatus *proto.FullStatus, opts ConvertO
peersOverview := mapPeers(pbFullStatus.GetPeers(), opts.StatusFilter, opts.PrefixNamesFilter, opts.PrefixNamesFilterMap, opts.IPsFilter, opts.ConnectionTypeFilter) peersOverview := mapPeers(pbFullStatus.GetPeers(), opts.StatusFilter, opts.PrefixNamesFilter, opts.PrefixNamesFilterMap, opts.IPsFilter, opts.ConnectionTypeFilter)
overview := OutputOverview{ overview := OutputOverview{
Peers: peersOverview, Peers: peersOverview,
CliVersion: version.NetbirdVersion(), CliVersion: version.NetbirdVersion(),
DaemonVersion: opts.DaemonVersion, DaemonVersion: opts.DaemonVersion,
DaemonStatus: opts.DaemonStatus, DaemonStatus: opts.DaemonStatus,
ManagementState: managementOverview, ManagementState: managementOverview,
SignalState: signalOverview, SignalState: signalOverview,
Relays: relayOverview, Relays: relayOverview,
IP: pbFullStatus.GetLocalPeerState().GetIP(), IP: pbFullStatus.GetLocalPeerState().GetIP(),
IPv6: pbFullStatus.GetLocalPeerState().GetIpv6(), IPv6: pbFullStatus.GetLocalPeerState().GetIpv6(),
PubKey: pbFullStatus.GetLocalPeerState().GetPubKey(), PubKey: pbFullStatus.GetLocalPeerState().GetPubKey(),
KernelInterface: pbFullStatus.GetLocalPeerState().GetKernelInterface(), KernelInterface: pbFullStatus.GetLocalPeerState().GetKernelInterface(),
WgPort: int(pbFullStatus.GetLocalPeerState().GetWgPort()), WgPort: int(pbFullStatus.GetLocalPeerState().GetWgPort()),
FQDN: pbFullStatus.GetLocalPeerState().GetFqdn(), FQDN: pbFullStatus.GetLocalPeerState().GetFqdn(),
RosenpassEnabled: pbFullStatus.GetLocalPeerState().GetRosenpassEnabled(), RosenpassEnabled: pbFullStatus.GetLocalPeerState().GetRosenpassEnabled(),
RosenpassPermissive: pbFullStatus.GetLocalPeerState().GetRosenpassPermissive(), RosenpassPermissive: pbFullStatus.GetLocalPeerState().GetRosenpassPermissive(),
Networks: pbFullStatus.GetLocalPeerState().GetNetworks(), Networks: pbFullStatus.GetLocalPeerState().GetNetworks(),
NumberOfForwardingRules: int(pbFullStatus.GetNumberOfForwardingRules()), NSServerGroups: mapNSGroups(pbFullStatus.GetDnsServers()),
NSServerGroups: mapNSGroups(pbFullStatus.GetDnsServers()), Events: mapEvents(pbFullStatus.GetEvents()),
Events: mapEvents(pbFullStatus.GetEvents()), LazyConnectionEnabled: pbFullStatus.GetLazyConnectionEnabled(),
LazyConnectionEnabled: pbFullStatus.GetLazyConnectionEnabled(), ProfileName: opts.ProfileName,
ProfileName: opts.ProfileName, SSHServerState: sshServerOverview,
SSHServerState: sshServerOverview,
} }
if !opts.SessionExpiresAt.IsZero() { if !opts.SessionExpiresAt.IsZero() {
t := opts.SessionExpiresAt t := opts.SessionExpiresAt
@@ -573,11 +571,6 @@ func (o *OutputOverview) GeneralSummary(showURL bool, showRelays bool, showNameS
) )
} }
var forwardingRulesString string
if o.NumberOfForwardingRules > 0 {
forwardingRulesString = fmt.Sprintf("Forwarding rules: %d\n", o.NumberOfForwardingRules)
}
goos := runtime.GOOS goos := runtime.GOOS
goarch := runtime.GOARCH goarch := runtime.GOARCH
goarm := "" goarm := ""
@@ -619,7 +612,6 @@ func (o *OutputOverview) GeneralSummary(showURL bool, showRelays bool, showNameS
"SSH Server: %s\n"+ "SSH Server: %s\n"+
"Networks: %s\n"+ "Networks: %s\n"+
"%s"+ "%s"+
"%s"+
"Peers count: %s\n", "Peers count: %s\n",
fmt.Sprintf("%s/%s%s", goos, goarch, goarm), fmt.Sprintf("%s/%s%s", goos, goarch, goarm),
daemonVersion, daemonVersion,
@@ -638,7 +630,6 @@ func (o *OutputOverview) GeneralSummary(showURL bool, showRelays bool, showNameS
lazyConnectionEnabledStatus, lazyConnectionEnabledStatus,
sshServerStatus, sshServerStatus,
networks, networks,
forwardingRulesString,
sessionExpiryString, sessionExpiryString,
peersCountString, peersCountString,
) )
@@ -691,7 +682,6 @@ func ToProtoFullStatus(fullStatus peer.FullStatus) *proto.FullStatus {
pbFullStatus.LocalPeerState.RosenpassPermissive = fullStatus.RosenpassState.Permissive pbFullStatus.LocalPeerState.RosenpassPermissive = fullStatus.RosenpassState.Permissive
pbFullStatus.LocalPeerState.RosenpassEnabled = fullStatus.RosenpassState.Enabled pbFullStatus.LocalPeerState.RosenpassEnabled = fullStatus.RosenpassState.Enabled
pbFullStatus.LocalPeerState.Networks = maps.Keys(fullStatus.LocalPeerState.Routes) pbFullStatus.LocalPeerState.Networks = maps.Keys(fullStatus.LocalPeerState.Routes)
pbFullStatus.NumberOfForwardingRules = int32(fullStatus.NumOfForwardingRules)
pbFullStatus.LazyConnectionEnabled = fullStatus.LazyConnectionEnabled pbFullStatus.LazyConnectionEnabled = fullStatus.LazyConnectionEnabled
for _, peerState := range fullStatus.Peers { for _, peerState := range fullStatus.Peers {
-2
View File
@@ -378,7 +378,6 @@ func TestParsingToJSON(t *testing.T) {
"networks": [ "networks": [
"10.10.0.0/24" "10.10.0.0/24"
], ],
"forwardingRules": 0,
"dnsServers": [ "dnsServers": [
{ {
"servers": [ "servers": [
@@ -496,7 +495,6 @@ quantumResistance: false
quantumResistancePermissive: false quantumResistancePermissive: false
networks: networks:
- 10.10.0.0/24 - 10.10.0.0/24
forwardingRules: 0
dnsServers: dnsServers:
- servers: - servers:
- 8.8.8.8:53 - 8.8.8.8:53
+1 -16
View File
@@ -10,7 +10,7 @@ Every method returns `$CancellablePromise<T>` (a Wails3 wrapper around `Promise`
// Services // Services
import { import {
Connection, Peers, ProfileSwitcher, Profiles, Connection, Peers, ProfileSwitcher, Profiles,
Settings, Networks, Forwarding, Debug, Update, WindowManager, Settings, Networks, Debug, Update, WindowManager,
I18n, Preferences, I18n, Preferences,
} from "@bindings/services"; } from "@bindings/services";
@@ -20,7 +20,6 @@ import type {
Profile, ProfileRef, ActiveProfile, Profile, ProfileRef, ActiveProfile,
Config, ConfigParams, SetConfigParams, Features, Config, ConfigParams, SetConfigParams, Features,
Network, SelectNetworksParams, Network, SelectNetworksParams,
ForwardingRule, PortInfo, PortRange,
LoginParams, LoginResult, LogoutParams, WaitSSOParams, UpParams, LoginParams, LoginResult, LogoutParams, WaitSSOParams, UpParams,
DebugBundleParams, DebugBundleResult, LogLevel, DebugBundleParams, DebugBundleResult, LogLevel,
UpdateResult, UpdateAvailable, UpdateProgress, UpdateResult, UpdateAvailable, UpdateProgress,
@@ -129,14 +128,6 @@ Networks.Deselect(p: SelectNetworksParams): Promise<void>
Exit-node filter: `range === "0.0.0.0/0" || range === "::/0"`. Domain network: `domains.length > 0`. CIDR overlap check is client-side. Exit-node filter: `range === "0.0.0.0/0" || range === "::/0"`. Domain network: `domains.length > 0`. CIDR overlap check is client-side.
## `Forwarding`
```ts
Forwarding.List(): Promise<ForwardingRule[]>
```
`PortInfo` is a daemon-side oneof — exactly one of `port?: number` or `range?: PortRange` is populated. `protocol` is the lowercase daemon string (`"tcp"` / `"udp"`).
## `Debug` ## `Debug`
```ts ```ts
@@ -269,12 +260,6 @@ The tray also reads a tray-only synthetic `"Error"` for icon purposes; the front
`Network`: `{ id, range: string; selected: boolean; domains: string[]; resolvedIps: Record<string, string[]> }`. `Network`: `{ id, range: string; selected: boolean; domains: string[]; resolvedIps: Record<string, string[]> }`.
`ForwardingRule`: `{ protocol: string; destinationPort: PortInfo; translatedAddress, translatedHostname: string; translatedPort: PortInfo }`.
`PortInfo`: `{ port?: number | null; range?: PortRange | null }` (exactly one populated).
`PortRange`: `{ start, end: number }` (inclusive).
`LoginParams`: `{ profileName, username, managementUrl, setupKey, preSharedKey, hostname, hint: string }`. `LoginParams`: `{ profileName, username, managementUrl, setupKey, preSharedKey, hostname, hint: string }`.
`LoginResult`: `{ needsSsoLogin: boolean; userCode, verificationUri, verificationUriComplete: string }`. `LoginResult`: `{ needsSsoLogin: boolean; userCode, verificationUri, verificationUriComplete: string }`.
@@ -80,7 +80,7 @@ export const CopyToClipboard = ({
aria-label={resolvedLabel} aria-label={resolvedLabel}
aria-live={"polite"} aria-live={"polite"}
className={cn( className={cn(
"group/copy wails-no-draggable pointer-events-auto inline-flex cursor-default items-center gap-2 rounded-sm text-left outline-none", "group/copy wails-no-draggable pointer-events-auto inline-flex cursor-default items-center gap-2 rounded-sm text-start outline-none",
"focus-visible:ring-2 focus-visible:ring-nb-gray-50/60 focus-visible:ring-offset-2 focus-visible:ring-offset-nb-gray-940", "focus-visible:ring-2 focus-visible:ring-nb-gray-50/60 focus-visible:ring-offset-2 focus-visible:ring-offset-nb-gray-940",
className, className,
)} )}
@@ -97,14 +97,14 @@ export const CopyToClipboard = ({
<span <span
aria-hidden={"true"} aria-hidden={"true"}
className={ className={
"pointer-events-none absolute bottom-0 left-0 right-0 border-b border-dashed border-transparent group-hover/copy:border-nb-gray-500" "pointer-events-none absolute inset-x-0 bottom-0 border-b border-dashed border-transparent group-hover/copy:border-nb-gray-500"
} }
/> />
</span> </span>
<span <span
aria-hidden={"true"} aria-hidden={"true"}
className={cn( className={cn(
"relative right-[1px] top-[2px] inline-flex shrink-0", "relative end-[1px] top-[2px] inline-flex shrink-0",
iconAlignment === "left" ? "order-first" : "order-last", iconAlignment === "left" ? "order-first" : "order-last",
iconClassName, iconClassName,
)} )}
@@ -3,8 +3,12 @@ import { cva } from "class-variance-authority";
import { Check, ChevronRight } from "lucide-react"; import { Check, ChevronRight } from "lucide-react";
import * as React from "react"; import * as React from "react";
import { cn } from "@/lib/cn"; import { cn } from "@/lib/cn";
import { useDirection } from "@/hooks/useDirection";
const DropdownMenu = DropdownMenuPrimitive.Root; const DropdownMenu = (props: React.ComponentProps<typeof DropdownMenuPrimitive.Root>) => {
const dir = useDirection();
return <DropdownMenuPrimitive.Root dir={dir} {...props} />;
};
const DropdownMenuTrigger = DropdownMenuPrimitive.Trigger; const DropdownMenuTrigger = DropdownMenuPrimitive.Trigger;
const DropdownMenuGroup = DropdownMenuPrimitive.Group; const DropdownMenuGroup = DropdownMenuPrimitive.Group;
const DropdownMenuPortal = DropdownMenuPrimitive.Portal; const DropdownMenuPortal = DropdownMenuPrimitive.Portal;
@@ -32,16 +36,16 @@ const DropdownMenuSubTrigger = React.forwardRef<
<DropdownMenuPrimitive.SubTrigger <DropdownMenuPrimitive.SubTrigger
ref={ref} ref={ref}
className={cn( className={cn(
"relative flex cursor-default select-none items-center rounded-md py-1.5 pl-3 pr-2 text-sm outline-none", "relative flex cursor-default select-none items-center rounded-md py-1.5 pe-2 ps-3 text-sm outline-none",
"transition-colors data-[disabled]:pointer-events-none data-[disabled]:opacity-50", "transition-colors data-[disabled]:pointer-events-none data-[disabled]:opacity-50",
inset && "pl-8", inset && "ps-8",
menuItemVariants({ variant }), menuItemVariants({ variant }),
className, className,
)} )}
{...props} {...props}
> >
{children} {children}
<ChevronRight className={"ml-auto h-4 w-4"} aria-hidden={"true"} /> <ChevronRight className={"ms-auto h-4 w-4 rtl:-scale-x-100"} aria-hidden={"true"} />
</DropdownMenuPrimitive.SubTrigger> </DropdownMenuPrimitive.SubTrigger>
)); ));
DropdownMenuSubTrigger.displayName = DropdownMenuPrimitive.SubTrigger.displayName; DropdownMenuSubTrigger.displayName = DropdownMenuPrimitive.SubTrigger.displayName;
@@ -102,9 +106,9 @@ const DropdownMenuItem = React.forwardRef<
<DropdownMenuPrimitive.Item <DropdownMenuPrimitive.Item
ref={ref} ref={ref}
className={cn( className={cn(
"relative flex cursor-default select-none items-center rounded-md py-1.5 pl-2 pr-2 text-sm outline-none", "relative flex cursor-default select-none items-center rounded-md px-2 py-1.5 text-sm outline-none",
"transition-colors data-[disabled]:pointer-events-none data-[disabled]:opacity-50", "transition-colors data-[disabled]:pointer-events-none data-[disabled]:opacity-50",
inset && "pl-8", inset && "ps-8",
menuItemVariants({ variant }), menuItemVariants({ variant }),
className, className,
)} )}
@@ -134,7 +138,7 @@ const DropdownMenuCheckboxItem = React.forwardRef<
<DropdownMenuPrimitive.CheckboxItem <DropdownMenuPrimitive.CheckboxItem
ref={ref} ref={ref}
className={cn( className={cn(
"relative flex cursor-default select-none items-center rounded-sm py-1.5 pl-8 pr-2 text-sm outline-none", "relative flex cursor-default select-none items-center rounded-sm py-1.5 pe-2 ps-8 text-sm outline-none",
"text-nb-gray-200 transition-colors hover:bg-nb-gray-900 hover:text-nb-gray-50 focus-visible:bg-nb-gray-900 focus-visible:text-nb-gray-50", "text-nb-gray-200 transition-colors hover:bg-nb-gray-900 hover:text-nb-gray-50 focus-visible:bg-nb-gray-900 focus-visible:text-nb-gray-50",
"data-[disabled]:pointer-events-none data-[disabled]:opacity-50", "data-[disabled]:pointer-events-none data-[disabled]:opacity-50",
className, className,
@@ -142,7 +146,7 @@ const DropdownMenuCheckboxItem = React.forwardRef<
checked={checked} checked={checked}
{...props} {...props}
> >
<span className={"absolute left-2 flex h-3.5 w-3.5 items-center justify-center"}> <span className={"absolute start-2 flex h-3.5 w-3.5 items-center justify-center"}>
<DropdownMenuPrimitive.ItemIndicator> <DropdownMenuPrimitive.ItemIndicator>
<Check className={"h-4 w-4"} /> <Check className={"h-4 w-4"} />
</DropdownMenuPrimitive.ItemIndicator> </DropdownMenuPrimitive.ItemIndicator>
@@ -170,7 +174,7 @@ const DropdownMenuRadioItem = React.forwardRef<
{children} {children}
<span <span
aria-hidden={"true"} aria-hidden={"true"}
className={"ml-auto flex w-4 shrink-0 items-center justify-center"} className={"ms-auto flex w-4 shrink-0 items-center justify-center"}
> >
<DropdownMenuPrimitive.ItemIndicator> <DropdownMenuPrimitive.ItemIndicator>
<Check size={14} className={"text-netbird"} /> <Check size={14} className={"text-netbird"} />
@@ -190,7 +194,7 @@ const DropdownMenuLabel = React.forwardRef<
ref={ref} ref={ref}
className={cn( className={cn(
"px-2 py-1.5 text-sm font-semibold text-nb-gray-200", "px-2 py-1.5 text-sm font-semibold text-nb-gray-200",
inset && "pl-8", inset && "ps-8",
className, className,
)} )}
{...props} {...props}
@@ -212,7 +216,7 @@ DropdownMenuSeparator.displayName = DropdownMenuPrimitive.Separator.displayName;
const DropdownMenuShortcut = ({ className, ...props }: React.HTMLAttributes<HTMLSpanElement>) => ( const DropdownMenuShortcut = ({ className, ...props }: React.HTMLAttributes<HTMLSpanElement>) => (
<span <span
className={cn("ml-auto text-xs tracking-widest text-nb-gray-100 opacity-60", className)} className={cn("ms-auto text-xs tracking-widest text-nb-gray-100 opacity-60", className)}
{...props} {...props}
/> />
); );
@@ -12,6 +12,7 @@ import { useFocusVisible } from "@/hooks/useFocusVisible";
import { loadLanguages } from "@/lib/i18n"; import { loadLanguages } from "@/lib/i18n";
import { cn } from "@/lib/cn"; import { cn } from "@/lib/cn";
import { errorDialog, formatErrorMessage } from "@/lib/errors"; import { errorDialog, formatErrorMessage } from "@/lib/errors";
import { useDirection } from "@/hooks/useDirection";
// No flag icons: flags represent countries, not languages. https://www.flagsarenotlanguages.com/blog/ // No flag icons: flags represent countries, not languages. https://www.flagsarenotlanguages.com/blog/
@@ -21,6 +22,7 @@ const labelFor = (lang: Language): string =>
: lang.displayName; : lang.displayName;
export function LanguagePicker() { export function LanguagePicker() {
const dir = useDirection();
const { t, i18n } = useTranslation(); const { t, i18n } = useTranslation();
const [languages, setLanguages] = useState<Language[]>([]); const [languages, setLanguages] = useState<Language[]>([]);
const [open, setOpen] = useState(false); const [open, setOpen] = useState(false);
@@ -112,7 +114,7 @@ export function LanguagePicker() {
aria-hidden={"true"} aria-hidden={"true"}
className={"shrink-0 text-nb-gray-200"} className={"shrink-0 text-nb-gray-200"}
/> />
<span className={"flex-1 truncate text-left"}> <span className={"flex-1 truncate text-start"}>
{current ? labelFor(current) : "—"} {current ? labelFor(current) : "—"}
</span> </span>
<ChevronDown <ChevronDown
@@ -168,7 +170,11 @@ export function LanguagePicker() {
</div> </div>
</div> </div>
<ScrollArea.Root type={"auto"} className={"-mx-1 overflow-hidden"}> <ScrollArea.Root
dir={dir}
type={"auto"}
className={"-mx-1 overflow-hidden"}
>
<ScrollArea.Viewport className={"max-h-64 px-1"}> <ScrollArea.Viewport className={"max-h-64 px-1"}>
<Command.List> <Command.List>
<Command.Empty> <Command.Empty>
+10 -1
View File
@@ -1,6 +1,7 @@
import { type ReactNode, useEffect, useRef, useState } from "react"; import { type ReactNode, useEffect, useRef, useState } from "react";
import * as RTooltip from "@radix-ui/react-tooltip"; import * as RTooltip from "@radix-ui/react-tooltip";
import { cn } from "@/lib/cn"; import { cn } from "@/lib/cn";
import { useDirection } from "@/hooks/useDirection";
type Props = { type Props = {
content: ReactNode; content: ReactNode;
@@ -29,6 +30,8 @@ export const Tooltip = ({
contentClassName, contentClassName,
closeDelay = 0, closeDelay = 0,
}: Props) => { }: Props) => {
const dir = useDirection();
const physicalSide = dir === "rtl" ? mirrorSide(side) : side;
const [open, setOpen] = useState(false); const [open, setOpen] = useState(false);
const hoveringRef = useRef(false); const hoveringRef = useRef(false);
const closeTimer = useRef<ReturnType<typeof setTimeout> | null>(null); const closeTimer = useRef<ReturnType<typeof setTimeout> | null>(null);
@@ -73,7 +76,7 @@ export const Tooltip = ({
</RTooltip.Trigger> </RTooltip.Trigger>
<RTooltip.Portal> <RTooltip.Portal>
<RTooltip.Content <RTooltip.Content
side={side} side={physicalSide}
align={align} align={align}
sideOffset={sideOffset} sideOffset={sideOffset}
alignOffset={alignOffset} alignOffset={alignOffset}
@@ -96,3 +99,9 @@ export const Tooltip = ({
</RTooltip.Provider> </RTooltip.Provider>
); );
}; };
function mirrorSide(side: Props["side"]): Props["side"] {
if (side === "left") return "right";
if (side === "right") return "left";
return side;
}
@@ -6,9 +6,16 @@ type Props = {
className?: string; className?: string;
tooltipContent?: ReactNode; tooltipContent?: ReactNode;
delayDuration?: number; delayDuration?: number;
dir?: "ltr" | "rtl" | "auto";
}; };
export const TruncatedText = ({ text, className, tooltipContent, delayDuration = 600 }: Props) => { export const TruncatedText = ({
text,
className,
tooltipContent,
delayDuration = 600,
dir = "auto",
}: Props) => {
const ref = useRef<HTMLSpanElement>(null); const ref = useRef<HTMLSpanElement>(null);
const [overflowing, setOverflowing] = useState(false); const [overflowing, setOverflowing] = useState(false);
@@ -19,7 +26,7 @@ export const TruncatedText = ({ text, className, tooltipContent, delayDuration =
}, [text]); }, [text]);
const span = ( const span = (
<span ref={ref} className={className}> <span ref={ref} dir={dir} className={className}>
{text} {text}
</span> </span>
); );
@@ -3,12 +3,15 @@ import * as Tabs from "@radix-ui/react-tabs";
import { type LucideProps } from "lucide-react"; import { type LucideProps } from "lucide-react";
import { cn } from "@/lib/cn"; import { cn } from "@/lib/cn";
import { useFocusVisible } from "@/hooks/useFocusVisible"; import { useFocusVisible } from "@/hooks/useFocusVisible";
import { useDirection } from "@/hooks/useDirection";
const Root = forwardRef<HTMLDivElement, Omit<Tabs.TabsProps, "orientation">>( const Root = forwardRef<HTMLDivElement, Omit<Tabs.TabsProps, "orientation">>(
function VerticalTabsRoot({ className, ...props }, ref) { function VerticalTabsRoot({ className, ...props }, ref) {
const dir = useDirection();
return ( return (
<Tabs.Root <Tabs.Root
ref={ref} ref={ref}
dir={dir}
orientation={"vertical"} orientation={"vertical"}
className={cn("flex min-h-0 flex-1", className)} className={cn("flex min-h-0 flex-1", className)}
{...props} {...props}
@@ -24,7 +27,7 @@ const List = forwardRef<HTMLDivElement, Tabs.TabsListProps>(function VerticalTab
return ( return (
<Tabs.List <Tabs.List
ref={ref} ref={ref}
className={cn("flex w-full flex-col gap-1 p-5 pr-0", className)} className={cn("flex w-full flex-col gap-1 p-5 pe-0", className)}
{...props} {...props}
/> />
); );
@@ -46,7 +49,7 @@ const Trigger = forwardRef<HTMLButtonElement, TriggerProps>(function VerticalTab
<Tabs.Trigger <Tabs.Trigger
ref={ref} ref={ref}
className={cn( className={cn(
"group flex w-full cursor-default items-center gap-3 rounded-md border border-transparent px-2 py-2.5 text-left outline-none dark:border-0", "group flex w-full cursor-default items-center gap-3 rounded-md border border-transparent px-2 py-2.5 text-start outline-none dark:border-0",
"transition-colors duration-150", "transition-colors duration-150",
"data-[state=active]:border-nb-gray-800 data-[state=active]:bg-white dark:data-[state=active]:bg-nb-gray-930", "data-[state=active]:border-nb-gray-800 data-[state=active]:bg-white dark:data-[state=active]:bg-nb-gray-930",
"data-[state=inactive]:hover:bg-nb-gray-850 dark:data-[state=inactive]:hover:bg-nb-gray-935", "data-[state=inactive]:hover:bg-nb-gray-850 dark:data-[state=inactive]:hover:bg-nb-gray-935",
@@ -60,7 +63,7 @@ const Trigger = forwardRef<HTMLButtonElement, TriggerProps>(function VerticalTab
size={iconSize} size={iconSize}
aria-hidden={"true"} aria-hidden={"true"}
className={cn( className={cn(
"ml-2 shrink-0 transition-colors duration-150", "ms-2 shrink-0 transition-colors duration-150",
"text-nb-gray-350 dark:text-nb-gray-400", "text-nb-gray-350 dark:text-nb-gray-400",
"group-data-[state=active]:text-nb-gray-100", "group-data-[state=active]:text-nb-gray-100",
)} )}
@@ -75,7 +78,7 @@ const Trigger = forwardRef<HTMLButtonElement, TriggerProps>(function VerticalTab
{title} {title}
</span> </span>
{adornment && ( {adornment && (
<div aria-hidden={"true"} className={"ml-auto mr-2 shrink-0"}> <div aria-hidden={"true"} className={"me-2 ms-auto shrink-0"}>
{adornment} {adornment}
</div> </div>
)} )}
@@ -54,9 +54,9 @@ export const ConfirmModal = ({
onOpenAutoFocus={(e) => e.preventDefault()} onOpenAutoFocus={(e) => e.preventDefault()}
> >
<div className={"flex flex-col gap-5 px-5"}> <div className={"flex flex-col gap-5 px-5"}>
<div className={"flex flex-col gap-1 pl-1"}> <div className={"flex flex-col gap-1 ps-1"}>
<DialogHeading align={"left"}>{title}</DialogHeading> <DialogHeading align={"start"}>{title}</DialogHeading>
<DialogDescription align={"left"} className={"whitespace-pre-line"}> <DialogDescription align={"start"} className={"whitespace-pre-line"}>
{description} {description}
</DialogDescription> </DialogDescription>
</div> </div>
@@ -95,7 +95,7 @@ export const Content = forwardRef<ElementRef<typeof DialogPrimitive.Content>, Co
{showClose && ( {showClose && (
<DialogPrimitive.Close <DialogPrimitive.Close
className={cn( className={cn(
"absolute right-3 top-3 z-10 rounded-md p-3 transition-colors", "absolute end-3 top-3 z-10 rounded-md p-3 transition-colors",
"text-nb-gray-300 hover:text-nb-gray-100", "text-nb-gray-300 hover:text-nb-gray-100",
"focus:outline-none disabled:pointer-events-none", "focus:outline-none disabled:pointer-events-none",
)} )}
@@ -1,12 +1,12 @@
import { type ReactNode } from "react"; import { type ReactNode } from "react";
import { cn } from "@/lib/cn"; import { cn } from "@/lib/cn";
type DialogAlign = "left" | "center" | "right"; type DialogAlign = "start" | "center" | "end";
const alignClass: Record<DialogAlign, string> = { const alignClass: Record<DialogAlign, string> = {
left: "text-left", start: "text-start",
center: "text-center", center: "text-center",
right: "text-right", end: "text-end",
}; };
type DialogDescriptionProps = { type DialogDescriptionProps = {
@@ -1,12 +1,12 @@
import { type ReactNode } from "react"; import { type ReactNode } from "react";
import { cn } from "@/lib/cn"; import { cn } from "@/lib/cn";
type DialogAlign = "left" | "center" | "right"; type DialogAlign = "start" | "center" | "end";
const alignClass: Record<DialogAlign, string> = { const alignClass: Record<DialogAlign, string> = {
left: "text-left", start: "text-start",
center: "text-center", center: "text-center",
right: "text-right", end: "text-end",
}; };
type DialogHeadingProps = { type DialogHeadingProps = {
@@ -65,7 +65,7 @@ export const DaemonOutdatedOverlay = () => {
{clientVersion === "development" ? ( {clientVersion === "development" ? (
<span> <span>
{t("settings.about.clientName")}{" "} {t("settings.about.clientName")}{" "}
<span className={"font-mono text-yellow-400"}> <span dir={"auto"} className={"font-mono text-yellow-400"}>
{t("settings.about.development")} {t("settings.about.development")}
</span> </span>
</span> </span>
@@ -77,7 +77,7 @@ export const DaemonOutdatedOverlay = () => {
{guiVersion === "development" ? ( {guiVersion === "development" ? (
<span> <span>
{t("settings.about.guiName")}{" "} {t("settings.about.guiName")}{" "}
<span className={"font-mono text-yellow-400"}> <span dir={"auto"} className={"font-mono text-yellow-400"}>
{t("settings.about.development")} {t("settings.about.development")}
</span> </span>
</span> </span>
@@ -86,13 +86,13 @@ function buildInputClassName(
"file:border-0 file:bg-transparent file:text-sm file:font-medium", "file:border-0 file:bg-transparent file:text-sm file:font-medium",
"focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-offset-2", "focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-offset-2",
"disabled:cursor-not-allowed disabled:opacity-40", "disabled:cursor-not-allowed disabled:opacity-40",
opts.hasCustomPrefix && "!rounded-l-none !border-l-0", opts.hasCustomPrefix && "!rounded-s-none !border-s-0",
opts.hasSuffix && "!pr-9", opts.hasSuffix && "!pe-9",
opts.hasIcon && "!pl-10", opts.hasIcon && "!ps-10",
"border", "border",
opts.readOnly && "!border-nb-gray-800 !bg-nb-gray-910 text-nb-gray-350", opts.readOnly && "!border-nb-gray-800 !bg-nb-gray-910 text-nb-gray-350",
opts.showStepper && opts.showStepper &&
"!rounded-r-none [-moz-appearance:textfield] [&::-webkit-inner-spin-button]:appearance-none [&::-webkit-outer-spin-button]:appearance-none", "!rounded-e-none [-moz-appearance:textfield] [&::-webkit-inner-spin-button]:appearance-none [&::-webkit-outer-spin-button]:appearance-none",
opts.className, opts.className,
); );
} }
@@ -107,7 +107,7 @@ function InputAffix({
<div <div
className={cn( className={cn(
inputVariants({ prefixSuffixVariant: error ? "error" : "default" }), inputVariants({ prefixSuffixVariant: error ? "error" : "default" }),
"flex h-[40px] w-auto rounded-l-md bg-white px-3 py-2 text-sm", "flex h-[40px] w-auto rounded-s-md bg-white px-3 py-2 text-sm",
"items-center whitespace-nowrap border", "items-center whitespace-nowrap border",
disabled && "opacity-40", disabled && "opacity-40",
className, className,
@@ -122,7 +122,7 @@ function InputIconSlot({ icon, disabled }: Readonly<{ icon: ReactNode; disabled?
return ( return (
<div <div
className={cn( className={cn(
"absolute left-0 top-0 flex h-full items-center pl-3 text-xs leading-[0] dark:text-nb-gray-300", "absolute start-0 top-0 flex h-full items-center ps-3 text-xs leading-[0] dark:text-nb-gray-300",
disabled && "opacity-40", disabled && "opacity-40",
)} )}
> >
@@ -138,7 +138,7 @@ function InputSuffixSlot({
return ( return (
<div <div
className={cn( className={cn(
"pointer-events-none absolute right-0 top-0 flex h-full select-none items-center pr-3 text-xs leading-[0] dark:text-nb-gray-300", "pointer-events-none absolute end-0 top-0 flex h-full select-none items-center pe-3 text-xs leading-[0] dark:text-nb-gray-300",
disabled && "opacity-30", disabled && "opacity-30",
)} )}
> >
@@ -157,7 +157,7 @@ function NumberStepper({
<div <div
className={cn( className={cn(
"flex h-[40px] shrink-0 flex-col overflow-hidden", "flex h-[40px] shrink-0 flex-col overflow-hidden",
"rounded-r-md border border-l-0", "rounded-e-md border border-s-0",
"border-neutral-200 bg-white dark:border-nb-gray-700 dark:bg-nb-gray-900", "border-neutral-200 bg-white dark:border-nb-gray-700 dark:bg-nb-gray-900",
error && "dark:border-red-500", error && "dark:border-red-500",
disabled && "pointer-events-none opacity-40", disabled && "pointer-events-none opacity-40",
@@ -226,6 +226,7 @@ export const Input = forwardRef<HTMLInputElement, InputProps>(function Input(
showPasswordToggle = false, showPasswordToggle = false,
copy = false, copy = false,
id, id,
dir,
...props ...props
}, },
ref, ref,
@@ -336,7 +337,7 @@ export const Input = forwardRef<HTMLInputElement, InputProps>(function Input(
return ( return (
<div className={"flex w-full min-w-0 flex-col"}> <div className={"flex w-full min-w-0 flex-col"}>
{label && <Label htmlFor={inputId}>{label}</Label>} {label && <Label htmlFor={inputId}>{label}</Label>}
<div className={cn("relative flex h-[40px] w-full", maxWidthClass)}> <div dir={dir} className={cn("relative flex h-[40px] w-full", maxWidthClass)}>
{customPrefix && ( {customPrefix && (
<InputAffix <InputAffix
content={customPrefix} content={customPrefix}

Some files were not shown because too many files have changed in this diff Show More