Merge branch 'main' into file-share

# Conflicts:
#	client/android/client.go
#	client/android/login.go
#	client/android/profile_prefs.go
#	client/internal/connect.go
#	client/internal/engine.go
#	client/mobile/profile_state.go
#	client/ui/i18n/locales/de/common.json
#	client/ui/i18n/locales/en/common.json
#	client/ui/i18n/locales/es/common.json
#	client/ui/i18n/locales/fr/common.json
#	client/ui/i18n/locales/hu/common.json
#	client/ui/i18n/locales/it/common.json
#	client/ui/i18n/locales/ja/common.json
#	client/ui/i18n/locales/pt/common.json
#	client/ui/i18n/locales/ru/common.json
#	client/ui/i18n/locales/zh-CN/common.json
#	client/ui/main.go
This commit is contained in:
Zoltán Papp
2026-08-27 11:13:19 +02:00
344 changed files with 19335 additions and 4922 deletions
+12 -1
View File
@@ -12,6 +12,13 @@ on:
AWS issues it. Leave empty for the Sonnet 4.6 default. AWS issues it. Leave empty for the Sonnet 4.6 default.
required: false required: false
default: "" default: ""
test_pattern:
description: >-
Package pattern to run. Defaults to the whole suite; narrow it to one
package (e.g. ./e2e/agentnetwork/...) when a run only needs that
package's answer and not the sixteen minutes the container suite costs.
required: false
default: "./e2e/..."
concurrency: concurrency:
group: ${{ github.workflow }}-${{ github.ref }} group: ${{ github.workflow }}-${{ github.ref }}
@@ -77,4 +84,8 @@ jobs:
GOOGLE_VERTEX_PROJECT: ${{ secrets.E2E_GOOGLE_VERTEX_PROJECT }} GOOGLE_VERTEX_PROJECT: ${{ secrets.E2E_GOOGLE_VERTEX_PROJECT }}
GOOGLE_VERTEX_REGION: ${{ secrets.E2E_GOOGLE_VERTEX_REGION }} GOOGLE_VERTEX_REGION: ${{ secrets.E2E_GOOGLE_VERTEX_REGION }}
GOOGLE_VERTEX_MODEL: ${{ secrets.E2E_GOOGLE_VERTEX_MODEL }} GOOGLE_VERTEX_MODEL: ${{ secrets.E2E_GOOGLE_VERTEX_MODEL }}
run: go test -tags e2e -timeout 40m -v ./e2e/... # Read through an env var rather than interpolated into the run
# script: a dispatch input reaching a shell command directly is a
# script-injection seam, however trusted the dispatcher.
TEST_PATTERN: ${{ inputs.test_pattern || './e2e/...' }}
run: go test -tags e2e -timeout 40m -v "$TEST_PATTERN"
+33
View File
@@ -0,0 +1,33 @@
name: protobuf checks
on:
push:
branches:
- main
- "release-*"
pull_request:
paths:
- ".github/workflows/buf.yml"
- "**/buf.yaml"
- "**/buf.lock"
- "**/buf.gen.yaml"
- "**.proto"
permissions:
contents: read
pull-requests: read
jobs:
buf:
runs-on: ubuntu-latest
steps:
- name: Checkout code
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- uses: bufbuild/buf-action@8c6a16e16f12ba20b6470afa9c2ba9b5ba8c97c3 # v1.5.0
with:
push: false
archive: false
pr_comment: false
build: false
lint: false
format: false
breaking: true
@@ -1,72 +0,0 @@
name: Mobile
on:
push:
branches:
- main
- "release-*"
pull_request:
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }}
cancel-in-progress: true
jobs:
android_build:
name: "Android / Build"
runs-on: ubuntu-latest
steps:
- name: Checkout repository
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: "go.mod"
- name: Setup Android SDK
uses: android-actions/setup-android@40fd30fb8d7440372e1316f5d1809ec01dcd3699 # v4.0.1
with:
cmdline-tools-version: 8512546
- name: Setup Java
uses: actions/setup-java@1bcf9fb12cf4aa7d266a90ae39939e61372fe520
with:
java-version: "11"
distribution: "adopt"
- name: NDK Cache
id: ndk-cache
uses: actions/cache@2c8a9bd7457de244a408f35966fab2fb45fda9c8 # v6.0.0
with:
path: /usr/local/lib/android/sdk/ndk
key: ndk-cache-23.1.7779620
- name: Setup NDK
run: /usr/local/lib/android/sdk/cmdline-tools/7.0/bin/sdkmanager --install "ndk;23.1.7779620"
- name: install gomobile
run: go install golang.org/x/mobile/cmd/gomobile@v0.0.0-20251113184115-a159579294ab
- name: gomobile init
run: gomobile init
- name: build android netbird lib
run: PATH=$PATH:$(go env GOPATH) gomobile bind -o $GITHUB_WORKSPACE/netbird.aar -javapkg=io.netbird.gomobile -ldflags="-checklinkname=0 -X golang.zx2c4.com/wireguard/ipc.socketDirectory=/data/data/io.netbird.client/cache/wireguard -X github.com/netbirdio/netbird/version.version=buildtest" $GITHUB_WORKSPACE/client/android
env:
CGO_ENABLED: 0
ANDROID_NDK_HOME: /usr/local/lib/android/sdk/ndk/23.1.7779620
ios_build:
name: "iOS / Build"
runs-on: macos-latest
steps:
- name: Checkout repository
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: "go.mod"
- name: install gomobile
run: go install golang.org/x/mobile/cmd/gomobile@v0.0.0-20251113184115-a159579294ab
- name: gomobile init
run: gomobile init
- name: build iOS netbird lib
run: PATH=$PATH:$(go env GOPATH) gomobile bind -target=ios -bundleid=io.netbird.framework -ldflags="-X github.com/netbirdio/netbird/version.version=buildtest" -o ./NetBirdSDK.xcframework ./client/ios/NetBirdSDK
env:
CGO_ENABLED: 0
+78
View File
@@ -0,0 +1,78 @@
name: No New Replace Directives
on:
pull_request:
paths:
- "go.mod"
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }}
cancel-in-progress: true
jobs:
check-replace-directives:
name: check-replace-directives
runs-on: ubuntu-latest
steps:
- name: Checkout code
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
fetch-depth: 0
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: go.mod
- name: Compare replace directives against the base branch
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
set -euo pipefail
# A replace directive only applies when this module is the main
# module. Anything importing netbird as a library, the embedded
# clients among them, resolves the replaced path upstream instead and
# fails to build against whatever the replacement provides. Requiring
# a fork under its own module path avoids that; a replace does not.
#
# go.mod is parsed rather than diffed so that reordering, comments and
# single-line versus block syntax do not register as changes.
#
# Versions are part of the key because a replace can be scoped to one
# version of a module. Keyed on paths alone, retargeting such a
# directive at a different version would read as unchanged.
list_replaces() {
go mod edit -json "$1" \
| jq -r '
def ref: .Path + (if (.Version // "") == "" then "" else " " + .Version end);
(.Replace // [])[] | "\(.Old | ref) => \(.New | ref)"
' \
| sort
}
git show "${BASE_SHA}:go.mod" > /tmp/base-go.mod
list_replaces /tmp/base-go.mod > /tmp/base-replaces
list_replaces go.mod > /tmp/head-replaces
added=$(comm -13 /tmp/base-replaces /tmp/head-replaces)
if [ -n "$added" ]; then
echo "::error::This PR adds a replace directive to go.mod:"
echo "$added" | sed 's/^/ /'
echo ""
echo "A replace directive applies only to the main module, so it does not"
echo "reach anything that imports netbird as a library. Require the module"
echo "under a path you control instead, as done for github.com/netbirdio/go-nat."
exit 1
fi
removed=$(comm -23 /tmp/base-replaces /tmp/head-replaces)
if [ -n "$removed" ]; then
echo "This PR removes replace directives:"
echo "$removed" | sed 's/^/ /'
fi
echo "No new replace directives."
+13
View File
@@ -37,3 +37,16 @@ jobs:
repo: netbirdio/ios-client repo: netbirdio/ios-client
token: ${{ secrets.NC_GITHUB_TOKEN }} token: ${{ secrets.NC_GITHUB_TOKEN }}
inputs: '{ "tag": "${{ github.ref_name }}" }' inputs: '{ "tag": "${{ github.ref_name }}" }'
trigger_dashboard_bump:
runs-on: ubuntu-latest
if: github.event.created && !github.event.deleted && startsWith(github.ref, 'refs/tags/v') && !contains(github.ref_name, '-')
steps:
- name: Trigger dashboard wasm client bump
uses: benc-uk/workflow-dispatch@31e2b3319479a63f0ab15bf800eff9e913504e26 # v1.3.2
with:
workflow: bump-netbird.yml
ref: main
repo: netbirdio/dashboard
token: ${{ secrets.NC_GITHUB_TOKEN }}
inputs: '{ "tag": "${{ github.ref_name }}" }'
+10
View File
@@ -92,6 +92,11 @@ nfpms:
dst: /usr/share/applications/org.wails.netbird.desktop dst: /usr/share/applications/org.wails.netbird.desktop
- src: client/ui/build/appicon.png - src: client/ui/build/appicon.png
dst: /usr/share/pixmaps/netbird.png dst: /usr/share/pixmaps/netbird.png
# Names the polkit action for the elevation prompt the app raises when an
# unprivileged user changes a privileged setting; without it the dialog
# shows a raw command line.
- src: client/ui/build/linux/polkit/io.netbird.settings.policy
dst: /usr/share/polkit-1/actions/io.netbird.settings.policy
dependencies: dependencies:
- netbird (>= 0.75.0) - netbird (>= 0.75.0)
- libgtk-4-1 (>= 4.14) - libgtk-4-1 (>= 4.14)
@@ -116,6 +121,11 @@ nfpms:
dst: /usr/share/applications/org.wails.netbird.desktop dst: /usr/share/applications/org.wails.netbird.desktop
- src: client/ui/build/appicon.png - src: client/ui/build/appicon.png
dst: /usr/share/pixmaps/netbird.png dst: /usr/share/pixmaps/netbird.png
# Names the polkit action for the elevation prompt the app raises when an
# unprivileged user changes a privileged setting; without it the dialog
# shows a raw command line.
- src: client/ui/build/linux/polkit/io.netbird.settings.policy
dst: /usr/share/polkit-1/actions/io.netbird.settings.policy
dependencies: dependencies:
- netbird >= 0.75.0 - netbird >= 0.75.0
- (gtk4 >= 4.14 or libgtk-4-1 >= 4.14) - (gtk4 >= 4.14 or libgtk-4-1 >= 4.14)
+1 -1
View File
@@ -1,6 +1,6 @@
# NetBird Agent Guidelines # NetBird Agent Guidelines
**NetBird** is an open-source connectivity platform: a WireGuard®-based overlay **NetBird** is an open source connectivity platform: a WireGuard®-based overlay
network with a control plane. The **agent** (`client/`) runs on user machines as network with a control plane. The **agent** (`client/`) runs on user machines as
a privileged daemon and manages the WireGuard interface, routing, firewall, and a privileged daemon and manages the WireGuard interface, routing, firewall, and
DNS. **Management** (`management/`) is the control plane and REST/gRPC API, DNS. **Management** (`management/`) is the control plane and REST/gRPC API,
+1 -1
View File
@@ -479,7 +479,7 @@ go test -race ./client/internal/dns/...
## Checklist before submitting a PR ## Checklist before submitting a PR
As a critical network service and open-source project, we must enforce a few As a critical network service and open source project, we must enforce a few
things before submitting a pull request. The things before submitting a pull request. The
[pull request template](/.github/pull_request_template.md) mirrors this list — [pull request template](/.github/pull_request_template.md) mirrors this list —
fill it in rather than deleting it. fill it in rather than deleting it.
+1 -1
View File
@@ -130,7 +130,7 @@ In November 2022, NetBird joined the [StartUpSecure program](https://www.forschu
![CISPA_Logo_BLACK_EN_RZ_RGB (1)](https://user-images.githubusercontent.com/700848/203091324-c6d311a0-22b5-4b05-a288-91cbc6cdcc46.png) ![CISPA_Logo_BLACK_EN_RZ_RGB (1)](https://user-images.githubusercontent.com/700848/203091324-c6d311a0-22b5-4b05-a288-91cbc6cdcc46.png)
### Acknowledgements ### Acknowledgements
We build on open-source technologies like [WireGuard®](https://www.wireguard.com/), [Pion ICE](https://github.com/pion/ice), and [Rosenpass](https://rosenpass.eu). We greatly appreciate the work these projects are doing, and we'd love it if you could support them too (e.g., by starring or contributing). We build on open source technologies like [WireGuard®](https://www.wireguard.com/), [Pion ICE](https://github.com/pion/ice), and [Rosenpass](https://rosenpass.eu). We greatly appreciate the work these projects are doing, and we'd love it if you could support them too (e.g., by starring or contributing).
### Legal ### Legal
This repository is licensed under the BSD-3-Clause license, which applies to all parts of the repository except for the directories management/, signal/ and relay/. This repository is licensed under the BSD-3-Clause license, which applies to all parts of the repository except for the directories management/, signal/ and relay/.
+1 -1
View File
@@ -14,7 +14,7 @@ Report security issues one of these two ways:
on this repository. This is the preferred route: it keeps the discussion, the draft advisory, and the credit in one place. on this repository. This is the preferred route: it keeps the discussion, the draft advisory, and the credit in one place.
- **Email** — `security@netbird.io`. - **Email** — `security@netbird.io`.
If the finding affects NetBird Cloud or our hosted infrastructure rather than the open-source code, email us rather than If the finding affects NetBird Cloud or our hosted infrastructure rather than the open source code, email us rather than
filing a repository report. filing a repository report.
### What to include ### What to include
+29
View File
@@ -40,6 +40,35 @@ You can then use this private endpoint to configure your AI agents, whether that
Full step-by-step setup: Full step-by-step setup:
**https://docs.netbird.io/agent-network/quickstart** **https://docs.netbird.io/agent-network/quickstart**
## Client settings that don't follow the endpoint
Most of an agent's traffic follows the base URL you hand it, but a few
client-side checks call their vendor directly and never reach the proxy. On a
network that blocks direct egress they fail even though inference works, so
they are worth setting once when you roll the endpoint out.
For Claude Code:
- **Fast mode** checks availability against `api.anthropic.com` rather than the
configured base URL. Set `CLAUDE_CODE_SKIP_FAST_MODE_ORG_CHECK=1` when the
agent authenticates with `ANTHROPIC_AUTH_TOKEN` alone (the usual shape when
the proxy injects the real provider key) or when a TLS-inspecting proxy
answers the check itself. Set
`CLAUDE_CODE_SKIP_FAST_MODE_NETWORK_ERRORS=1` when the network refuses the
connection outright. Fast mode is an Anthropic-API feature, so it is
unavailable on a Bedrock- or Vertex-backed endpoint whatever you set.
- **Model discovery** is off by default. Set
`CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY=1` for the picker to list the
models your policies authorise; the proxy filters the response to that set.
The client gives discovery a three-second budget and treats any redirect as
a failure, so the endpoint must serve `/v1/models` directly.
- **The WebFetch domain safety check** also calls `api.anthropic.com` directly
and is unaffected by the variables above.
Allowing direct egress to `api.anthropic.com` covers the network cases but not
the credential one, where the check reaches Anthropic and is rejected because
the agent presents a proxy-issued key.
## Architecture ## Architecture
Agent Network is built on two existing NetBird capabilities: Agent Network is built on two existing NetBird capabilities:
+39 -9
View File
@@ -26,6 +26,7 @@ import (
"github.com/netbirdio/netbird/client/internal/routemanager" "github.com/netbirdio/netbird/client/internal/routemanager"
"github.com/netbirdio/netbird/client/internal/stdnet" "github.com/netbirdio/netbird/client/internal/stdnet"
"github.com/netbirdio/netbird/client/net" "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/netevents"
"github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/client/system"
"github.com/netbirdio/netbird/formatter" "github.com/netbirdio/netbird/formatter"
"github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/route"
@@ -40,11 +41,6 @@ const (
AnonymizeLevelStrict = nbAnonymize.LevelStrictString AnonymizeLevelStrict = nbAnonymize.LevelStrictString
) )
// ConnectionListener export internal Listener for mobile
type ConnectionListener interface {
peer.Listener
}
// TunAdapter export internal TunAdapter for mobile // TunAdapter export internal TunAdapter for mobile
type TunAdapter interface { type TunAdapter interface {
device.TunAdapter device.TunAdapter
@@ -85,6 +81,10 @@ type Client struct {
deviceName string deviceName string
uiVersion string uiVersion string
networkChangeListener listener.NetworkChangeListener networkChangeListener listener.NetworkChangeListener
// netMgr outlives engine restarts: it mirrors the OS connectivity, not
// the engine lifecycle. Run and RunWithoutLogin inject its state and
// sweeper into each new ConnectClient.
netMgr *netevents.Manager
stateMu sync.RWMutex stateMu sync.RWMutex
connectClient *internal.ConnectClient connectClient *internal.ConnectClient
@@ -154,14 +154,17 @@ func NewClient(androidSDKVersion int, deviceName string, uiVersion string, tunAd
execWorkaround(androidSDKVersion) execWorkaround(androidSDKVersion)
net.SetAndroidProtectSocketFn(tunAdapter.ProtectSocket) net.SetAndroidProtectSocketFn(tunAdapter.ProtectSocket)
system.SetIFaceDiscover(iFaceDiscover)
recorder := peer.NewRecorder("")
return &Client{ return &Client{
deviceName: deviceName, deviceName: deviceName,
uiVersion: uiVersion, uiVersion: uiVersion,
tunAdapter: tunAdapter, tunAdapter: tunAdapter,
iFaceDiscover: iFaceDiscover, iFaceDiscover: iFaceDiscover,
recorder: peer.NewRecorder(""), recorder: recorder,
ctxCancelLock: &sync.Mutex{}, ctxCancelLock: &sync.Mutex{},
networkChangeListener: networkChangeListener, networkChangeListener: networkChangeListener,
netMgr: netevents.NewManager(recorder),
} }
} }
@@ -202,7 +205,9 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
} }
// todo do not throw error in case of cancelled context // todo do not throw error in case of cancelled context
ctx = internal.CtxInitState(ctx) ctx = internal.CtxInitState(ctx)
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder)
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetEvents(c.netMgr))
c.attachFileDrop(connectClient, cfgFile) c.attachFileDrop(connectClient, cfgFile)
c.setState(cfg, cacheDir, cfgFile, connectClient) c.setState(cfg, cacheDir, cfgFile, connectClient)
// This path runs the interactive SSO flow, so reaching here means the peer // This path runs the interactive SSO flow, so reaching here means the peer
@@ -244,7 +249,8 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
// todo do not throw error in case of cancelled context // todo do not throw error in case of cancelled context
ctx = internal.CtxInitState(ctx) ctx = internal.CtxInitState(ctx)
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder) connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetEvents(c.netMgr))
c.attachFileDrop(connectClient, cfgFile) c.attachFileDrop(connectClient, cfgFile)
c.setState(cfg, cacheDir, cfgFile, connectClient) c.setState(cfg, cacheDir, cfgFile, connectClient)
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir) return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
@@ -293,6 +299,26 @@ func (c *Client) GetTunSettings() (*TunSettings, error) {
}, nil }, nil
} }
// SetNetworkAvailable feeds OS-reported network availability into the client.
// While unavailable, the internal reconnect loops suspend their attempts and
// the connection listener reports NoNetwork instead of Connecting; when
// availability returns, the loops resume immediately with a fresh backoff.
// Losing the last network also sweeps the registered connections: nothing can
// redial while offline, so the stale sockets would otherwise stay silently
// "connected" until their own timeouts and the client would keep reporting
// Connected with no network at all.
func (c *Client) SetNetworkAvailable(available bool) {
c.netMgr.SetNetworkAvailable(available)
}
// NotifyNetworkChange marks the management, signal and relay connections
// stale after the OS switched networks and schedules a sweep that cuts
// whatever has not redialed on the new network by then. The engine and the
// TUN device stay untouched.
func (c *Client) NotifyNetworkChange() {
c.netMgr.NotifyNetworkChange()
}
// DebugBundle generates a debug bundle, uploads it, and returns the upload key. // DebugBundle generates a debug bundle, uploads it, and returns the upload key.
// It works both with and without a running engine. anonymizeLevel is "default" // It works both with and without a running engine. anonymizeLevel is "default"
// or "strict"; strict also anonymizes internal IP ranges, peer names, and // or "strict"; strict also anonymizes internal IP ranges, peer names, and
@@ -533,7 +559,11 @@ func (c *Client) OnUpdatedHostDNS(list *DNSList) error {
// SetConnectionListener set the network connection listener // SetConnectionListener set the network connection listener
func (c *Client) SetConnectionListener(listener ConnectionListener) { func (c *Client) SetConnectionListener(listener ConnectionListener) {
c.recorder.SetConnectionListener(listener) if listener == nil {
c.recorder.RemoveConnectionListener()
return
}
c.recorder.SetConnectionListener(connectionListenerAdapter{listener})
} }
// RemoveConnectionListener remove connection listener // RemoveConnectionListener remove connection listener
+2 -1
View File
@@ -8,6 +8,7 @@ import (
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal" "github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/mobile"
) )
// SetFileDropSink installs the platform sink incoming payloads are staged // SetFileDropSink installs the platform sink incoming payloads are staged
@@ -72,7 +73,7 @@ func (c *Client) fileDropFor(configDir, profileID string) (*FileDrop, error) {
// pair one profile's engine with another's transfers. A failure is not fatal: // pair one profile's engine with another's transfers. A failure is not fatal:
// the tunnel is worth more than the feature, so the engine runs on without it. // the tunnel is worth more than the feature, so the engine runs on without it.
func (c *Client) attachFileDrop(cc *internal.ConnectClient, cfgFile string) { func (c *Client) attachFileDrop(cc *internal.ConnectClient, cfgFile string) {
configDir, profileID, err := profileLocationFor(cfgFile) configDir, profileID, err := mobile.ProfileLocationFor(cfgFile)
if err != nil { if err != nil {
log.Warnf("file drop is unavailable: %v", err) log.Warnf("file drop is unavailable: %v", err)
return return
+41
View File
@@ -0,0 +1,41 @@
//go:build android
package android
import (
"github.com/netbirdio/netbird/client/internal/peer"
)
// Client state values delivered via ConnectionListener.OnStateChanged,
// re-exported as basic constants so gomobile emits them into the generated
// Java bindings. They mirror peer.ClientState*: append-only, never reorder.
const (
ClientStateDisconnected = int(peer.ClientStateDisconnected)
ClientStateConnected = int(peer.ClientStateConnected)
ClientStateConnecting = int(peer.ClientStateConnecting)
ClientStateDisconnecting = int(peer.ClientStateDisconnecting)
ClientStateNoNetwork = int(peer.ClientStateNoNetwork)
)
// ConnectionListener export internal Listener for mobile. It mirrors
// peer.Listener with OnStateChanged taking a plain int (one of the
// ClientState* constants), because gomobile cannot bind named types.
type ConnectionListener interface {
OnStateChanged(state int)
OnConnected()
OnDisconnected()
OnConnecting()
OnDisconnecting()
OnAddressChanged(string, string)
OnPeersListChanged(int)
}
// connectionListenerAdapter adapts the gomobile-facing ConnectionListener to
// peer.Listener, converting the typed state to the int the binding carries.
type connectionListenerAdapter struct {
ConnectionListener
}
func (a connectionListenerAdapter) OnStateChanged(state peer.ClientState) {
a.ConnectionListener.OnStateChanged(int(state))
}
+1 -30
View File
@@ -10,7 +10,6 @@ import (
"testing" "testing"
"github.com/netbirdio/netbird/client/internal/filedrop" "github.com/netbirdio/netbird/client/internal/filedrop"
"github.com/netbirdio/netbird/client/internal/profilemanager"
) )
type stubStream struct { type stubStream struct {
@@ -224,38 +223,10 @@ func TestFileDropSendWithoutTunnelFails(t *testing.T) {
} }
} }
func TestProfileLocationForSplitsConfigPath(t *testing.T) {
root := t.TempDir()
dir, id, err := profileLocationFor(filepath.Join(root, defaultConfigFilename))
if err != nil {
t.Fatalf("default profile: %v", err)
}
if dir != root || id != profilemanager.DefaultProfileName {
t.Fatalf("default profile = (%q, %q), want (%q, %q)", dir, id, root, profilemanager.DefaultProfileName)
}
named := filepath.Join(root, profilesSubdir, "aaaaaaaabbbbbbbbccccccccdddddddd.json")
dir, id, err = profileLocationFor(named)
if err != nil {
t.Fatalf("named profile: %v", err)
}
if dir != root || id != "aaaaaaaabbbbbbbbccccccccdddddddd" {
t.Fatalf("named profile = (%q, %q), want (%q, %q)", dir, id, root, "aaaaaaaabbbbbbbbccccccccdddddddd")
}
for _, path := range []string{"", filepath.Join(root, "stray.json"), filepath.Join(root, profilesSubdir, "not-an-id!.json")} {
if _, _, err := profileLocationFor(path); err == nil {
t.Fatalf("expected an error for %q", path)
}
}
}
func writeTestProfile(t *testing.T, configDir, id string) { func writeTestProfile(t *testing.T, configDir, id string) {
t.Helper() t.Helper()
pm := NewProfileManager(configDir) if _, err := NewProfileManager(configDir).impl.ProfilePrefs(id); err != nil {
if _, err := pm.serviceMgr.ProfilePrefs(profilemanager.ID(id), androidUsername); err != nil {
t.Fatalf("resolve prefs for %s: %v", id, err) t.Fatalf("resolve prefs for %s: %v", id, err)
} }
} }
+3 -2
View File
@@ -8,6 +8,7 @@ import (
"github.com/netbirdio/netbird/client/internal/auth" "github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/profilemanager" "github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mobile"
"github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/client/system"
) )
@@ -181,7 +182,7 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
// Stored after Login, not before: a rejected token must not leave a hint // Stored after Login, not before: a rejected token must not leave a hint
// pointing at an account that cannot be used. // pointing at an account that cannot be used.
if email != "" && a.cfgPath != "" { if email != "" && a.cfgPath != "" {
if err := writeProfileEmail(a.cfgPath, email); err != nil { if err := mobile.WriteProfileEmail(a.cfgPath, email); err != nil {
log.Warnf("failed to store profile account email: %v", err) log.Warnf("failed to store profile account email: %v", err)
} }
} }
@@ -208,7 +209,7 @@ func profileLoginHint(cfgPath string) string {
if cfgPath == "" { if cfgPath == "" {
return "" return ""
} }
return readProfileEmail(cfgPath) return mobile.ReadProfileEmail(cfgPath)
} }
// runOAuthFlow drives an already acquired OAuth flow to a token: requests the // runOAuthFlow drives an already acquired OAuth flow to a token: requests the
+62 -228
View File
@@ -3,42 +3,37 @@
package android package android
import ( import (
"fmt" "github.com/netbirdio/netbird/client/mobile"
"os"
"path/filepath"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/profilemanager"
) )
const ( const (
// Android uses a single user context per app (non-empty username required by ServiceManager) // Android uses a single user context per app.
androidUsername = "android" androidUsername = "android"
) )
// Profile represents a profile for gomobile // Profile represents a profile for gomobile.
type Profile struct { type Profile struct {
ID string ID string
Name string Name string
// Email is the account this profile last logged in with, "" if it never // Email is the account this profile last logged in with, "" if it never
// completed an SSO login. Kept across logouts; cleared when the profile is // completed an SSO login. Kept across logouts; cleared when the profile is
// removed. See profile_state.go. // removed. See client/mobile/profile_state.go.
Email string Email string
IsActive bool IsActive bool
} }
// ProfileArray wraps profiles for gomobile compatibility // ProfileArray wraps profiles for gomobile compatibility (gomobile cannot
// bind Go slices directly).
type ProfileArray struct { type ProfileArray struct {
items []*Profile items []*Profile
} }
// Length returns the number of profiles // Length returns the number of profiles.
func (p *ProfileArray) Length() int { func (p *ProfileArray) Length() int {
return len(p.items) return len(p.items)
} }
// Get returns the profile at index i // Get returns the profile at index i, or nil if out of range.
func (p *ProfileArray) Get(i int) *Profile { func (p *ProfileArray) Get(i int) *Profile {
if i < 0 || i >= len(p.items) { if i < 0 || i >= len(p.items) {
return nil return nil
@@ -46,259 +41,98 @@ func (p *ProfileArray) Get(i int) *Profile {
return p.items[i] return p.items[i]
} }
/* // ProfileManager adapts the shared mobile profile manager (client/mobile) to
// gomobile-friendly types. See that package for the on-disk layout and
/data/data/io.netbird.client/files/ ← configDir parameter // semantics.
├── netbird.cfg ← Default profile config
├── state.json ← Default profile state
├── active_profile.json ← Active profile tracker (JSON with Name + Username)
└── profiles/ ← Subdirectory for non-default profiles
├── work.json ← Legacy work profile config
├── work.state.json ← Legacy work profile state
├── 4c5f5c8198c3989cffb5b5394f5a7ae0.json ← ID profile config
├── 4c5f5c8198c3989cffb5b5394f5a7ae0.state.json ← ID profile state
*/
// ProfileManager manages profiles for Android
// It wraps the internal profilemanager to provide Android-specific behavior
type ProfileManager struct { type ProfileManager struct {
configDir string impl *mobile.ProfileManager
serviceMgr *profilemanager.ServiceManager
} }
// NewProfileManager creates a new profile manager for Android // NewProfileManager creates a new profile manager for Android. configDir is
// the app's files directory.
func NewProfileManager(configDir string) *ProfileManager { func NewProfileManager(configDir string) *ProfileManager {
// Set the default config path for Android (stored in root configDir, not profiles/) return &ProfileManager{impl: mobile.NewProfileManager(configDir, androidUsername)}
defaultConfigPath := filepath.Join(configDir, defaultConfigFilename)
// Set global paths for Android
profilemanager.DefaultConfigPathDir = configDir
profilemanager.DefaultConfigPath = defaultConfigPath
profilemanager.ActiveProfileStatePath = filepath.Join(configDir, "active_profile.json")
// Create ServiceManager with profiles/ subdirectory
// This avoids modifying the global ConfigDirOverride for profile listing
profilesDir := filepath.Join(configDir, profilesSubdir)
serviceMgr := profilemanager.NewServiceManagerWithProfilesDir(defaultConfigPath, profilesDir)
return &ProfileManager{
configDir: configDir,
serviceMgr: serviceMgr,
}
} }
// ListProfiles returns all available profiles // ListProfiles returns all available profiles, including the default profile,
// with their active status set.
func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) { func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) {
// Use ServiceManager (looks in profiles/ directory, checks active_profile.json for IsActive) profiles, err := pm.impl.ListProfiles()
internalProfiles, err := pm.serviceMgr.ListProfiles(androidUsername)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to list profiles: %w", err) return nil, err
} }
// Convert internal profiles to Android Profile type items := make([]*Profile, 0, len(profiles))
var profiles []*Profile for i := range profiles {
for _, p := range internalProfiles { items = append(items, fromMobileProfile(&profiles[i]))
profiles = append(profiles, &Profile{
ID: p.ID.String(),
Name: p.Name,
Email: pm.profileEmail(p.ID.String()),
IsActive: p.IsActive,
})
} }
return &ProfileArray{items: items}, nil
return &ProfileArray{items: profiles}, nil
} }
// GetActiveProfile returns the currently active profile name // GetActiveProfile returns the currently active profile.
func (pm *ProfileManager) GetActiveProfile() (*Profile, error) { func (pm *ProfileManager) GetActiveProfile() (*Profile, error) {
// Use ServiceManager to stay consistent with ListProfiles p, err := pm.impl.GetActiveProfile()
// ServiceManager uses active_profile.json
activeState, err := pm.serviceMgr.GetActiveProfileState()
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to get active profile: %w", err) return nil, err
} }
return fromMobileProfile(p), nil
// ActiveProfileState only stores the ID (and username), not the display
// name. Resolve the ID to the full profile so callers get the real Name.
prof, err := pm.serviceMgr.ResolveProfile(activeState.ID.String(), androidUsername)
if err != nil {
return nil, fmt.Errorf("failed to resolve active profile %q: %w", activeState.ID, err)
}
return &Profile{
ID: prof.ID.String(),
Name: prof.Name,
Email: pm.profileEmail(prof.ID.String()),
IsActive: true,
}, nil
} }
// profileEmail returns the account email recorded for a profile. Display-only, so // SwitchProfile records the given profile ID as the active profile. The caller
// an unresolvable path degrades to "" rather than an error. // must stop the VPN tunnel before switching.
func (pm *ProfileManager) profileEmail(id string) string {
configPath, err := pm.getProfileConfigPath(id)
if err != nil {
return ""
}
return readProfileEmail(configPath)
}
// SwitchProfile switches to a different profile
func (pm *ProfileManager) SwitchProfile(id string) error { func (pm *ProfileManager) SwitchProfile(id string) error {
// Use ServiceManager to stay consistent with ListProfiles return pm.impl.SwitchProfile(id)
// ServiceManager uses active_profile.json
err := pm.serviceMgr.SetActiveProfileState(&profilemanager.ActiveProfileState{
ID: profilemanager.ID(id),
Username: androidUsername,
})
if err != nil {
return fmt.Errorf("failed to switch profile: %w", err)
}
log.Infof("switched to profile: %s", id)
return nil
} }
// AddProfile creates a new profile // AddProfile creates a new profile with the given display name and a
// generated ID.
func (pm *ProfileManager) AddProfile(profileName string) error { func (pm *ProfileManager) AddProfile(profileName string) error {
// Use ServiceManager (creates profile in profiles/ directory) _, err := pm.impl.AddProfile(profileName)
profile, err := pm.serviceMgr.AddProfile(profileName, androidUsername) return err
if err != nil {
return fmt.Errorf("failed to add profile: %w", err)
}
log.Infof("created new profile: %s", profile.ID)
return nil
} }
// LogoutProfile logs out from a profile (clears authentication) // RenameProfile changes the display name of the profile identified by id. The
func (pm *ProfileManager) LogoutProfile(id string) error { // on-disk filename (the ID) is left unchanged.
configPath, err := pm.getProfileConfigPath(id)
if err != nil {
return err
}
if !profilemanager.IsValidProfileFilenameStem(profilemanager.ID(id)) {
return fmt.Errorf("id '%s' is not valid", id)
}
// Check if profile exists
if _, err := os.Stat(configPath); os.IsNotExist(err) {
return fmt.Errorf("profile '%s' does not exist", id)
}
// Read current config using internal profilemanager
config, err := profilemanager.ReadConfig(configPath)
if err != nil {
return fmt.Errorf("failed to read profile config: %w", err)
}
// Clear authentication by removing private key and SSH key
config.PrivateKey = ""
config.SSHKey = ""
// Save config using internal profilemanager
if err := profilemanager.WriteOutConfig(configPath, config); err != nil {
return fmt.Errorf("failed to save config: %w", err)
}
// The stored account email is kept on purpose, matching the desktop and CLI
// logout semantics: the next login passes it as the login_hint so the IdP
// preselects the account. Removing the profile is what deletes it.
log.Infof("logged out from profile: %s", id)
return nil
}
// RenameProfile changes a profile's display name. The profile ID, and therefore
// its on-disk filename, is left untouched: only the "name" field of the config
// is rewritten. This works for the default profile too, whose config lives in
// netbird.cfg rather than under profiles/.
func (pm *ProfileManager) RenameProfile(id string, newName string) error { func (pm *ProfileManager) RenameProfile(id string, newName string) error {
if err := pm.serviceMgr.RenameProfile(profilemanager.ID(id), androidUsername, newName); err != nil { return pm.impl.RenameProfile(id, newName)
return fmt.Errorf("failed to rename profile: %w", err)
}
log.Infof("renamed profile %s to: %s", id, newName)
return nil
} }
// RemoveProfile deletes a profile // LogoutProfile clears authentication data for a profile, forcing a re-login.
// The management URL and other settings are preserved.
func (pm *ProfileManager) LogoutProfile(id string) error {
return pm.impl.LogoutProfile(id)
}
// RemoveProfile deletes a profile. The default profile and the active profile
// cannot be removed.
func (pm *ProfileManager) RemoveProfile(id string) error { func (pm *ProfileManager) RemoveProfile(id string) error {
configPath, err := pm.getProfileConfigPath(id) return pm.impl.RemoveProfile(id)
if err != nil {
return err
}
// Use ServiceManager (removes profile from profiles/ directory)
if err := pm.serviceMgr.RemoveProfile(profilemanager.ID(id), androidUsername); err != nil {
return fmt.Errorf("failed to remove profile: %w", err)
}
// The account file is this package's, not the ServiceManager's, so it must
// go here. The default profile has a fixed filename, so a recreated one
// would otherwise inherit the deleted profile's email as its login_hint.
// Not fatal: the profile itself is gone.
if err := removeProfileEmail(configPath); err != nil {
log.Warnf("failed to remove stored account email for profile %s: %v", id, err)
}
log.Infof("removed profile: %s", id)
return nil
} }
// getProfileConfigPath returns the config file path for a profile // GetConfigPath returns the config file path for the given profile ID. Java
// This is needed for Android-specific path handling (netbird.cfg for default profile) // should call this instead of constructing paths with Preferences.configFile().
func (pm *ProfileManager) getProfileConfigPath(id string) (string, error) {
if !profilemanager.IsValidProfileFilenameStem(profilemanager.ID(id)) {
return "", fmt.Errorf("id %q is not valid", id)
}
if id == profilemanager.DefaultProfileName {
// Android uses netbird.cfg for default profile instead of default.json
// Default profile is stored in root configDir, not in profiles/
return filepath.Join(pm.configDir, defaultConfigFilename), nil
}
profilesDir := filepath.Join(pm.configDir, profilesSubdir)
return filepath.Join(profilesDir, id+".json"), nil
}
// GetConfigPath returns the config file path for a given profile id
// Java should call this instead of constructing paths with Preferences.configFile()
func (pm *ProfileManager) GetConfigPath(id string) (string, error) { func (pm *ProfileManager) GetConfigPath(id string) (string, error) {
return pm.getProfileConfigPath(id) return pm.impl.GetConfigPath(id)
} }
// GetStateFilePath returns the state file path for a given profile // GetStateFilePath returns the state file path for the given profile ID. Java
// Java should call this instead of constructing paths with Preferences.stateFile() // should call this instead of constructing paths with Preferences.stateFile().
func (pm *ProfileManager) GetStateFilePath(id string) (string, error) { func (pm *ProfileManager) GetStateFilePath(id string) (string, error) {
if id == "" || id == profilemanager.DefaultProfileName { return pm.impl.GetStateFilePath(id)
return filepath.Join(pm.configDir, "state.json"), nil
}
if !profilemanager.IsValidProfileFilenameStem(profilemanager.ID(id)) {
return "", fmt.Errorf("id %q is not valid", id)
}
profilesDir := filepath.Join(pm.configDir, profilesSubdir)
return filepath.Join(profilesDir, id+".state.json"), nil
} }
// GetActiveConfigPath returns the config file path for the currently active profile // GetActiveConfigPath returns the config file path for the currently active
// Java should call this instead of Preferences.getActiveProfileName() + Preferences.configFile() // profile.
func (pm *ProfileManager) GetActiveConfigPath() (string, error) { func (pm *ProfileManager) GetActiveConfigPath() (string, error) {
activeProfile, err := pm.GetActiveProfile() return pm.impl.GetActiveConfigPath()
if err != nil {
return "", fmt.Errorf("failed to get active profile: %w", err)
}
return pm.GetConfigPath(activeProfile.ID)
} }
// GetActiveStateFilePath returns the state file path for the currently active profile // GetActiveStateFilePath returns the state file path for the currently active
// Java should call this instead of Preferences.getActiveProfileName() + Preferences.stateFile() // profile.
func (pm *ProfileManager) GetActiveStateFilePath() (string, error) { func (pm *ProfileManager) GetActiveStateFilePath() (string, error) {
activeProfile, err := pm.GetActiveProfile() return pm.impl.GetActiveStateFilePath()
if err != nil { }
return "", fmt.Errorf("failed to get active profile: %w", err)
} func fromMobileProfile(p *mobile.Profile) *Profile {
return pm.GetStateFilePath(activeProfile.ID) return &Profile{ID: p.ID, Name: p.Name, Email: p.Email, IsActive: p.IsActive}
} }
+2 -3
View File
@@ -21,10 +21,9 @@ func newProfilePrefs(configDir, profileID string) (*profilePrefs, error) {
if configDir == "" || profileID == "" { if configDir == "" || profileID == "" {
return nil, fmt.Errorf("profile prefs require a config dir and profile ID") return nil, fmt.Errorf("profile prefs require a config dir and profile ID")
} }
pm := NewProfileManager(configDir) prefs, err := NewProfileManager(configDir).impl.ProfilePrefs(profileID)
prefs, err := pm.serviceMgr.ProfilePrefs(profilemanager.ID(profileID), androidUsername)
if err != nil { if err != nil {
return nil, fmt.Errorf("resolve profile prefs: %w", err) return nil, err
} }
return &profilePrefs{prefs: prefs}, nil return &profilePrefs{prefs: prefs}, nil
} }
+2 -2
View File
@@ -45,8 +45,8 @@ func daemonServerOptions(network string) []grpc.ServerOption {
return nil return nil
} }
creds := ipcauth.NewTransportCredentials() creds := ipcauth.NewTransportCredentials() //nolint:staticcheck
if creds == nil { if creds == nil { //nolint:staticcheck // nil only on platforms without a peer-identity primitive
log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied", runtime.GOOS) log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied", runtime.GOOS)
return nil return nil
} }
+2 -2
View File
@@ -27,8 +27,8 @@ func listenOnAddress(addr string) (*socketListener, error) {
} }
if network == "npipe" { if network == "npipe" {
listener, path, err := listenNamedPipe(address) listener, path, err := listenNamedPipe(address) //nolint:staticcheck
if err != nil { if err != nil { //nolint:staticcheck // always errors on non-Windows builds
return nil, err return nil, err
} }
return &socketListener{Listener: listener, network: network, address: path}, nil return &socketListener{Listener: listener, network: network, address: path}, nil
+1 -1
View File
@@ -6,7 +6,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/golang/mock/gomock" "go.uber.org/mock/gomock"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.opentelemetry.io/otel" "go.opentelemetry.io/otel"
"google.golang.org/grpc" "google.golang.org/grpc"
+23
View File
@@ -85,12 +85,24 @@ type Options struct {
DisableIPv6 bool DisableIPv6 bool
// BlockInbound blocks all inbound connections from peers // BlockInbound blocks all inbound connections from peers
BlockInbound bool BlockInbound bool
// EnableRosenpass enables the Rosenpass post-quantum key exchange.
EnableRosenpass bool
// RosenpassPermissive lets a Rosenpass-enabled peer still connect to peers
// that do not run Rosenpass (falling back to the plain WireGuard PSK).
RosenpassPermissive bool
// BlockLANAccess blocks the embedded peer from reaching the host's // BlockLANAccess blocks the embedded peer from reaching the host's
// LAN (RFC 1918, link-local, loopback) when it's used as a routing // LAN (RFC 1918, link-local, loopback) when it's used as a routing
// peer. Mirrors profilemanager.ConfigInput.BlockLANAccess. Useful // peer. Mirrors profilemanager.ConfigInput.BlockLANAccess. Useful
// when the embedded client must never act as a stepping stone into // when the embedded client must never act as a stepping stone into
// the host's local network (e.g. the proxy's overlay peer). // the host's local network (e.g. the proxy's overlay peer).
BlockLANAccess bool BlockLANAccess bool
// LazyConnectionEnabled is a tri-state local override for lazy connections,
// mirroring the NB_LAZY_CONN env var. Nil defers to the management feature
// flag; a set value overrides it in both directions. A short-lived client
// that reaches only a few known peers can set this to false, so its peers
// connect eagerly and the first request does not wait for the connection to
// be established.
LazyConnectionEnabled *bool
// WireguardPort is the port for the tunnel interface. Use 0 for a random port. // WireguardPort is the port for the tunnel interface. Use 0 for a random port.
WireguardPort *int WireguardPort *int
// MTU is the MTU for the tunnel interface. // MTU is the MTU for the tunnel interface.
@@ -203,6 +215,8 @@ func New(opts Options) (*Client, error) {
DisableIPv6: &opts.DisableIPv6, DisableIPv6: &opts.DisableIPv6,
BlockInbound: &opts.BlockInbound, BlockInbound: &opts.BlockInbound,
BlockLANAccess: &opts.BlockLANAccess, BlockLANAccess: &opts.BlockLANAccess,
RosenpassEnabled: &opts.EnableRosenpass,
RosenpassPermissive: &opts.RosenpassPermissive,
WireguardPort: opts.WireguardPort, WireguardPort: opts.WireguardPort,
MTU: opts.MTU, MTU: opts.MTU,
DNSLabels: parsedLabels, DNSLabels: parsedLabels,
@@ -220,6 +234,15 @@ func New(opts Options) (*Client, error) {
config.PrivateKey = opts.PrivateKey config.PrivateKey = opts.PrivateKey
} }
if opts.LazyConnectionEnabled != nil {
// Runtime-only override, read back through lazyconn.ParseState; a set value
// wins over the management feature flag in both directions.
config.LazyConnection = "off"
if *opts.LazyConnectionEnabled {
config.LazyConnection = "on"
}
}
if opts.Performance.PreallocatedBuffersPerPool != nil { if opts.Performance.PreallocatedBuffersPerPool != nil {
wgdevice.SetPreallocatedBuffersPerPool(*opts.Performance.PreallocatedBuffersPerPool) wgdevice.SetPreallocatedBuffersPerPool(*opts.Performance.PreallocatedBuffersPerPool)
} }
+1 -1
View File
@@ -6,7 +6,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/golang/mock/gomock" "go.uber.org/mock/gomock"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"google.golang.org/grpc" "google.golang.org/grpc"
+1 -1
View File
@@ -763,7 +763,7 @@ func (r *router) addNatRule(pair firewall.RouterPair) error {
exprs = append(exprs, sourceExp...) exprs = append(exprs, sourceExp...)
exprs = append(exprs, destExp...) exprs = append(exprs, destExp...)
var markValue uint32 = nbnet.PreroutingFwmarkMasquerade markValue := nbnet.PreroutingFwmarkMasquerade
if pair.Inverse { if pair.Inverse {
markValue = nbnet.PreroutingFwmarkMasqueradeReturn markValue = nbnet.PreroutingFwmarkMasqueradeReturn
} }
@@ -5,7 +5,7 @@ import (
"net/netip" "net/netip"
"testing" "testing"
"github.com/golang/mock/gomock" "go.uber.org/mock/gomock"
"github.com/google/gopacket" "github.com/google/gopacket"
"github.com/google/gopacket/layers" "github.com/google/gopacket/layers"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -4,7 +4,7 @@ import (
"net/netip" "net/netip"
"testing" "testing"
"github.com/golang/mock/gomock" "go.uber.org/mock/gomock"
"github.com/google/gopacket/layers" "github.com/google/gopacket/layers"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
+40 -16
View File
@@ -16,28 +16,52 @@ import (
"google.golang.org/grpc" "google.golang.org/grpc"
nbnet "github.com/netbirdio/netbird/client/net" nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/netevents/sweep"
) )
// Sweeper registers in-flight dials for the network change sweep.
type Sweeper interface {
StartDial(ctx context.Context) *sweep.Dial
}
func WithCustomDialer(_ bool, _ string) grpc.DialOption { func WithCustomDialer(_ bool, _ string) grpc.DialOption {
return grpc.WithContextDialer(dialContext)
}
// WithSweeper dials like WithCustomDialer but registers connections and
// dials with the sweeper. Append it after WithCustomDialer: gRPC applies
// dial options in order, so the later context dialer wins.
func WithSweeper(sweeper Sweeper) grpc.DialOption {
return grpc.WithContextDialer(func(ctx context.Context, addr string) (net.Conn, error) { return grpc.WithContextDialer(func(ctx context.Context, addr string) (net.Conn, error) {
if runtime.GOOS == "linux" { dial := sweeper.StartDial(ctx)
currentUser, err := user.Current() defer dial.Release()
if err != nil {
return nil, status.Errorf(codes.FailedPrecondition, "failed to get current user: %v", err)
}
// the custom dialer requires root permissions which are not required for use cases run as non-root conn, err := dialContext(dial.Ctx(), addr)
if currentUser.Uid != "0" {
log.Debug("Not running as root, using standard dialer")
dialer := &net.Dialer{}
return dialer.DialContext(ctx, "tcp", addr)
}
}
conn, err := nbnet.NewDialer().DialContext(ctx, "tcp", addr)
if err != nil { if err != nil {
return nil, fmt.Errorf("nbnet.NewDialer().DialContext: %w", err) return nil, err
} }
return conn, nil return dial.WrapConn(conn)
}) })
} }
func dialContext(ctx context.Context, addr string) (net.Conn, error) {
if runtime.GOOS == "linux" {
currentUser, err := user.Current()
if err != nil {
return nil, status.Errorf(codes.FailedPrecondition, "failed to get current user: %v", err)
}
// the custom dialer requires root permissions which are not required for use cases run as non-root
if currentUser.Uid != "0" {
log.Debug("Not running as root, using standard dialer")
dialer := &net.Dialer{}
return dialer.DialContext(ctx, "tcp", addr)
}
}
conn, err := nbnet.NewDialer().DialContext(ctx, "tcp", addr)
if err != nil {
return nil, fmt.Errorf("nbnet.NewDialer().DialContext: %w", err)
}
return conn, nil
}
+13
View File
@@ -1,13 +1,26 @@
package grpc package grpc
import ( import (
"context"
"google.golang.org/grpc" "google.golang.org/grpc"
"github.com/netbirdio/netbird/client/netevents/sweep"
"github.com/netbirdio/netbird/util/wsproxy/client" "github.com/netbirdio/netbird/util/wsproxy/client"
) )
// Sweeper registers in-flight dials for the network change sweep.
type Sweeper interface {
StartDial(ctx context.Context) *sweep.Dial
}
// WithCustomDialer returns a gRPC dial option that uses WebSocket transport for WASM/JS environments. // WithCustomDialer returns a gRPC dial option that uses WebSocket transport for WASM/JS environments.
// The component parameter specifies the WebSocket proxy component path (e.g., "/management", "/signal"). // The component parameter specifies the WebSocket proxy component path (e.g., "/management", "/signal").
func WithCustomDialer(tlsEnabled bool, component string) grpc.DialOption { func WithCustomDialer(tlsEnabled bool, component string) grpc.DialOption {
return client.WithWebSocketDialer(tlsEnabled, component) return client.WithWebSocketDialer(tlsEnabled, component)
} }
// WithSweeper is a no-op on WASM/JS: there is no network change signal.
func WithSweeper(_ Sweeper) grpc.DialOption {
return grpc.EmptyDialOption{}
}
+56
View File
@@ -0,0 +1,56 @@
package grpc
import (
"context"
"errors"
"time"
"github.com/cenkalti/backoff/v4"
)
// ChangeWatcher exposes OS network availability transitions.
type ChangeWatcher interface {
Changed() <-chan struct{}
}
// Retry mirrors backoff.Retry, but the sleep between attempts also wakes on
// OS network availability transitions: an operation cut down by a network
// change retries the moment the network settles instead of sleeping through
// the recovery. A nil watcher never fires, leaving plain backoff.Retry
// behavior.
func Retry(ctx context.Context, operation backoff.Operation, bo backoff.BackOff, watcher ChangeWatcher) error {
bo.Reset()
for {
err := operation()
if err == nil {
return nil
}
var permanent *backoff.PermanentError
if errors.As(err, &permanent) {
return permanent.Err
}
next := bo.NextBackOff()
if next == backoff.Stop {
if cerr := ctx.Err(); cerr != nil {
return cerr
}
return err
}
var changed <-chan struct{}
if watcher != nil {
changed = watcher.Changed()
}
timer := time.NewTimer(next)
select {
case <-timer.C:
case <-changed:
timer.Stop()
case <-ctx.Done():
timer.Stop()
return ctx.Err()
}
}
}
+91
View File
@@ -0,0 +1,91 @@
package grpc
import (
"context"
"errors"
"testing"
"time"
"github.com/cenkalti/backoff/v4"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/netevents/netstate"
)
func TestRetryWakesOnNetworkChange(t *testing.T) {
ns := netstate.New()
attempts := 0
operation := func() error {
attempts++
if attempts == 1 {
return errors.New("cut by network change")
}
return nil
}
go func() {
time.Sleep(20 * time.Millisecond)
ns.Set(false)
}()
start := time.Now()
err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Minute), ns)
require.NoError(t, err)
assert.Equal(t, 2, attempts, "network change must cause one immediate retry")
assert.Less(t, time.Since(start), time.Second, "the transition must cut the minute-long sleep short")
}
func TestRetryPermanentError(t *testing.T) {
sentinel := errors.New("permission denied")
operation := func() error {
return backoff.Permanent(sentinel)
}
err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Millisecond), nil)
assert.ErrorIs(t, err, sentinel, "permanent errors must stop retries")
}
func TestRetryNilNetState(t *testing.T) {
attempts := 0
operation := func() error {
attempts++
if attempts < 3 {
return errors.New("transient")
}
return nil
}
err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Millisecond), nil)
require.NoError(t, err)
assert.Equal(t, 3, attempts, "nil network state must preserve timed retries")
}
func TestRetryStops(t *testing.T) {
failure := errors.New("still failing")
operation := func() error {
return failure
}
err := Retry(context.Background(), operation, &backoff.StopBackOff{}, nil)
assert.ErrorIs(t, err, failure, "stop backoff must return the operation error")
}
func TestRetryCtxCancelDuringSleep(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
operation := func() error {
return errors.New("failing")
}
go func() {
time.Sleep(20 * time.Millisecond)
cancel()
}()
start := time.Now()
err := Retry(ctx, operation, backoff.NewConstantBackOff(time.Minute), netstate.New())
assert.ErrorIs(t, err, context.Canceled, "context cancellation must stop the retry loop")
assert.Less(t, time.Since(start), time.Second, "context cancellation must interrupt backoff sleep")
}
+1 -1
View File
@@ -502,7 +502,7 @@ func toBytes(s string) (int64, error) {
func getFwmark() int { func getFwmark() int {
if nbnet.AdvancedRouting() && runtime.GOOS == "linux" { if nbnet.AdvancedRouting() && runtime.GOOS == "linux" {
return nbnet.ControlPlaneMark return int(nbnet.ControlPlaneMark)
} }
return 0 return 0
} }
+1 -1
View File
@@ -4,7 +4,7 @@ import (
"net" "net"
"testing" "testing"
"github.com/golang/mock/gomock" "go.uber.org/mock/gomock"
"github.com/google/gopacket" "github.com/google/gopacket"
"github.com/google/gopacket/layers" "github.com/google/gopacket/layers"
+1 -1
View File
@@ -8,7 +8,7 @@ import (
"net/netip" "net/netip"
reflect "reflect" reflect "reflect"
gomock "github.com/golang/mock/gomock" gomock "go.uber.org/mock/gomock"
) )
// MockPacketFilter is a mock of PacketFilter interface. // MockPacketFilter is a mock of PacketFilter interface.
+1 -1
View File
@@ -8,7 +8,7 @@ import (
os "os" os "os"
reflect "reflect" reflect "reflect"
gomock "github.com/golang/mock/gomock" gomock "go.uber.org/mock/gomock"
tun "golang.zx2c4.com/wireguard/tun" tun "golang.zx2c4.com/wireguard/tun"
) )
+3 -3
View File
@@ -53,15 +53,15 @@ func NewProxyBind(bind Bind, mtu uint16) *ProxyBind {
return p return p
} }
// AddTurnConn adds a new connection to the bind. // AddRelayedConn adds a new connection to the bind.
// endpoint is the NetBird address of the remote peer. The SetEndpoint return with the address what will be used in the // endpoint is the NetBird address of the remote peer. The SetEndpoint return with the address what will be used in the
// WireGuard configuration. // WireGuard configuration.
// //
// Parameters: // Parameters:
// - ctx: Context is used for proxyToLocal to avoid unnecessary error messages // - ctx: Context is used for proxyToLocal to avoid unnecessary error messages
// - nbAddr: The NetBird UDP address of the remote peer, it required to generate fake address // - nbAddr: The NetBird UDP address of the remote peer, it required to generate fake address
// - remoteConn: The established TURN connection to the remote peer // - remoteConn: The established relayed connection to the remote peer
func (p *ProxyBind) AddTurnConn(ctx context.Context, nbAddr *net.UDPAddr, remoteConn net.Conn) error { func (p *ProxyBind) AddRelayedConn(ctx context.Context, nbAddr *net.UDPAddr, remoteConn net.Conn) error {
fakeNetIP, err := fakeAddress(nbAddr) fakeNetIP, err := fakeAddress(nbAddr)
if err != nil { if err != nil {
return err return err
+26 -26
View File
@@ -30,9 +30,9 @@ type WGEBPFProxy struct {
proxyPort int proxyPort int
mtu uint16 mtu uint16
ebpfManager ebpfMgr.Manager ebpfManager ebpfMgr.Manager
turnConnStore map[uint16]net.Conn relayedConnStore map[uint16]net.Conn
turnConnMutex sync.Mutex relayedConnMutex sync.Mutex
lastUsedPort uint16 lastUsedPort uint16
rawConnIPv4 net.PacketConn rawConnIPv4 net.PacketConn
@@ -50,7 +50,7 @@ func NewWGEBPFProxy(wgPort int, mtu uint16) *WGEBPFProxy {
localWGListenPort: wgPort, localWGListenPort: wgPort,
mtu: mtu, mtu: mtu,
ebpfManager: ebpf.GetEbpfManagerInstance(), ebpfManager: ebpf.GetEbpfManagerInstance(),
turnConnStore: make(map[uint16]net.Conn), relayedConnStore: make(map[uint16]net.Conn),
} }
return wgProxy return wgProxy
} }
@@ -110,14 +110,14 @@ func (p *WGEBPFProxy) Listen() error {
return nil return nil
} }
// AddTurnConn add new turn connection for the proxy // AddRelayedConn add new relayed connection for the proxy
func (p *WGEBPFProxy) AddTurnConn(turnConn net.Conn) (*net.UDPAddr, error) { func (p *WGEBPFProxy) AddRelayedConn(relayedConn net.Conn) (*net.UDPAddr, error) {
wgEndpointPort, err := p.storeTurnConn(turnConn) wgEndpointPort, err := p.storeRelayedConn(relayedConn)
if err != nil { if err != nil {
return nil, err return nil, err
} }
log.Infof("turn conn added to wg proxy store: %s, endpoint port: :%d", turnConn.RemoteAddr(), wgEndpointPort) log.Infof("relayed conn added to wg proxy store: %s, endpoint port: :%d", relayedConn.RemoteAddr(), wgEndpointPort)
wgEndpoint := &net.UDPAddr{ wgEndpoint := &net.UDPAddr{
IP: net.ParseIP(loopbackAddr), IP: net.ParseIP(loopbackAddr),
@@ -186,48 +186,48 @@ func (p *WGEBPFProxy) readAndForwardPacket(buf []byte) error {
return fmt.Errorf("failed to read UDP packet from WG: %w", err) return fmt.Errorf("failed to read UDP packet from WG: %w", err)
} }
p.turnConnMutex.Lock() p.relayedConnMutex.Lock()
conn, ok := p.turnConnStore[uint16(addr.Port)] conn, ok := p.relayedConnStore[uint16(addr.Port)]
p.turnConnMutex.Unlock() p.relayedConnMutex.Unlock()
if !ok { if !ok {
if p.ctx.Err() == nil { if p.ctx.Err() == nil {
log.Debugf("turn conn not found by port because conn already has been closed: %d", addr.Port) log.Debugf("relayed conn not found by port because conn already has been closed: %d", addr.Port)
} }
return nil return nil
} }
if _, err := conn.Write(buf[:n]); err != nil { if _, err := conn.Write(buf[:n]); err != nil {
return fmt.Errorf("failed to forward local WG packet (%d) to remote turn conn: %w", addr.Port, err) return fmt.Errorf("forward local WG packet (%d) to remote relayed conn: %w", addr.Port, err)
} }
return nil return nil
} }
func (p *WGEBPFProxy) storeTurnConn(turnConn net.Conn) (uint16, error) { func (p *WGEBPFProxy) storeRelayedConn(relayedConn net.Conn) (uint16, error) {
p.turnConnMutex.Lock() p.relayedConnMutex.Lock()
defer p.turnConnMutex.Unlock() defer p.relayedConnMutex.Unlock()
np, err := p.nextFreePort() np, err := p.nextFreePort()
if err != nil { if err != nil {
return np, err return np, err
} }
p.turnConnStore[np] = turnConn p.relayedConnStore[np] = relayedConn
return np, nil return np, nil
} }
func (p *WGEBPFProxy) removeTurnConn(turnConnID uint16) { func (p *WGEBPFProxy) removeRelayedConn(relayedConnID uint16) {
p.turnConnMutex.Lock() p.relayedConnMutex.Lock()
defer p.turnConnMutex.Unlock() defer p.relayedConnMutex.Unlock()
_, ok := p.turnConnStore[turnConnID] _, ok := p.relayedConnStore[relayedConnID]
if ok { if ok {
log.Debugf("remove turn conn from store by port: %d", turnConnID) log.Debugf("remove relayed conn from store by port: %d", relayedConnID)
} }
delete(p.turnConnStore, turnConnID) delete(p.relayedConnStore, relayedConnID)
} }
func (p *WGEBPFProxy) nextFreePort() (uint16, error) { func (p *WGEBPFProxy) nextFreePort() (uint16, error) {
if len(p.turnConnStore) == 65535 { if len(p.relayedConnStore) == 65535 {
return 0, fmt.Errorf("reached maximum turn connection numbers") return 0, fmt.Errorf("reached maximum relayed connection numbers")
} }
generatePort: generatePort:
if p.lastUsedPort == 65535 { if p.lastUsedPort == 65535 {
@@ -236,7 +236,7 @@ generatePort:
p.lastUsedPort++ p.lastUsedPort++
} }
if _, ok := p.turnConnStore[p.lastUsedPort]; ok { if _, ok := p.relayedConnStore[p.lastUsedPort]; ok {
goto generatePort goto generatePort
} }
return p.lastUsedPort, nil return p.lastUsedPort, nil
+11 -11
View File
@@ -9,32 +9,32 @@ import (
func TestWGEBPFProxy_connStore(t *testing.T) { func TestWGEBPFProxy_connStore(t *testing.T) {
wgProxy := NewWGEBPFProxy(1, 1280) wgProxy := NewWGEBPFProxy(1, 1280)
p, _ := wgProxy.storeTurnConn(nil) p, _ := wgProxy.storeRelayedConn(nil)
if p != 1 { if p != 1 {
t.Errorf("invalid initial port: %d", wgProxy.lastUsedPort) t.Errorf("invalid initial port: %d", wgProxy.lastUsedPort)
} }
numOfConns := 10 numOfConns := 10
for i := 0; i < numOfConns; i++ { for i := 0; i < numOfConns; i++ {
p, _ = wgProxy.storeTurnConn(nil) p, _ = wgProxy.storeRelayedConn(nil)
} }
if p != uint16(numOfConns)+1 { if p != uint16(numOfConns)+1 {
t.Errorf("invalid last used port: %d, expected: %d", p, numOfConns+1) t.Errorf("invalid last used port: %d, expected: %d", p, numOfConns+1)
} }
if len(wgProxy.turnConnStore) != numOfConns+1 { if len(wgProxy.relayedConnStore) != numOfConns+1 {
t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.turnConnStore), numOfConns+1) t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), numOfConns+1)
} }
} }
func TestWGEBPFProxy_portCalculation_overflow(t *testing.T) { func TestWGEBPFProxy_portCalculation_overflow(t *testing.T) {
wgProxy := NewWGEBPFProxy(1, 1280) wgProxy := NewWGEBPFProxy(1, 1280)
_, _ = wgProxy.storeTurnConn(nil) _, _ = wgProxy.storeRelayedConn(nil)
wgProxy.lastUsedPort = 65535 wgProxy.lastUsedPort = 65535
p, _ := wgProxy.storeTurnConn(nil) p, _ := wgProxy.storeRelayedConn(nil)
if len(wgProxy.turnConnStore) != 2 { if len(wgProxy.relayedConnStore) != 2 {
t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.turnConnStore), 2) t.Errorf("invalid store size: %d, expected: %d", len(wgProxy.relayedConnStore), 2)
} }
if p != 2 { if p != 2 {
@@ -46,11 +46,11 @@ func TestWGEBPFProxy_portCalculation_maxConn(t *testing.T) {
wgProxy := NewWGEBPFProxy(1, 1280) wgProxy := NewWGEBPFProxy(1, 1280)
for i := 0; i < 65535; i++ { for i := 0; i < 65535; i++ {
_, _ = wgProxy.storeTurnConn(nil) _, _ = wgProxy.storeRelayedConn(nil)
} }
_, err := wgProxy.storeTurnConn(nil) _, err := wgProxy.storeRelayedConn(nil)
if err == nil { if err == nil {
t.Errorf("invalid turn conn store calculation") t.Errorf("invalid relayed conn store calculation")
} }
} }
+6 -6
View File
@@ -121,10 +121,10 @@ func NewProxyWrapper(proxy *WGEBPFProxy) *ProxyWrapper {
} }
} }
func (p *ProxyWrapper) AddTurnConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error { func (p *ProxyWrapper) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error {
addr, err := p.wgeBPFProxy.AddTurnConn(remoteConn) addr, err := p.wgeBPFProxy.AddRelayedConn(remoteConn)
if err != nil { if err != nil {
return fmt.Errorf("add turn conn: %w", err) return fmt.Errorf("add relayed conn: %w", err)
} }
headers, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, addr) headers, err := NewPacketHeaders(p.wgeBPFProxy.localWGListenPort, addr)
@@ -252,7 +252,7 @@ func (p *ProxyWrapper) CloseConn() error {
} }
func (p *ProxyWrapper) proxyToLocal(ctx context.Context) { func (p *ProxyWrapper) proxyToLocal(ctx context.Context) {
defer p.wgeBPFProxy.removeTurnConn(uint16(p.wgRelayedEndpointAddr.Port)) defer p.wgeBPFProxy.removeRelayedConn(uint16(p.wgRelayedEndpointAddr.Port))
buf := make([]byte, p.wgeBPFProxy.mtu+bufsize.WGBufferOverhead) buf := make([]byte, p.wgeBPFProxy.mtu+bufsize.WGBufferOverhead)
for { for {
@@ -273,7 +273,7 @@ func (p *ProxyWrapper) proxyToLocal(ctx context.Context) {
if ctx.Err() != nil { if ctx.Err() != nil {
return return
} }
log.Errorf("failed to write out turn pkg to local conn: %v", err) log.Errorf("failed to write out relayed pkg to local conn: %v", err)
} }
} }
} }
@@ -286,7 +286,7 @@ func (p *ProxyWrapper) readFromRemote(ctx context.Context, buf []byte) (int, err
} }
p.closeListener.Notify() p.closeListener.Notify()
if !errors.Is(err, io.EOF) { if !errors.Is(err, io.EOF) {
log.Errorf("failed to read from turn conn (endpoint: :%d): %s", p.wgRelayedEndpointAddr.Port, err) log.Errorf("failed to read from relayed conn (endpoint: :%d): %s", p.wgRelayedEndpointAddr.Port, err)
} }
return 0, err return 0, err
} }
+1 -1
View File
@@ -7,7 +7,7 @@ import (
// Proxy is a transfer layer between the relayed connection and the WireGuard // Proxy is a transfer layer between the relayed connection and the WireGuard
type Proxy interface { type Proxy interface {
AddTurnConn(ctx context.Context, endpoint *net.UDPAddr, remoteConn net.Conn) error AddRelayedConn(ctx context.Context, endpoint *net.UDPAddr, remoteConn net.Conn) error
EndpointAddr() *net.UDPAddr // EndpointAddr returns the address of the WireGuard peer endpoint EndpointAddr() *net.UDPAddr // EndpointAddr returns the address of the WireGuard peer endpoint
Work() // Work start or resume the proxy Work() // Work start or resume the proxy
Pause() // Pause to forward the packages from remote connection to WireGuard. The opposite way still works. Pause() // Pause to forward the packages from remote connection to WireGuard. The opposite way still works.
+2 -2
View File
@@ -95,7 +95,7 @@ func TestProxyCloseByRemoteConn(t *testing.T) {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
addr, _ := net.ResolveUDPAddr("udp", "100.108.135.221:51892") addr, _ := net.ResolveUDPAddr("udp", "100.108.135.221:51892")
relayedConn := newMockConn() relayedConn := newMockConn()
err := tt.proxy.AddTurnConn(ctx, addr, relayedConn) err := tt.proxy.AddRelayedConn(ctx, addr, relayedConn)
if err != nil { if err != nil {
t.Errorf("error: %v", err) t.Errorf("error: %v", err)
} }
@@ -157,7 +157,7 @@ func redirectTraffic(t *testing.T, proxy Proxy, wgPort int, endPointAddr *net.UD
_ = relayedServer.Close() _ = relayedServer.Close()
}() }()
if err := proxy.AddTurnConn(context.Background(), endPointAddr, relayedConn); err != nil { if err := proxy.AddRelayedConn(context.Background(), endPointAddr, relayedConn); err != nil {
t.Errorf("error: %v", err) t.Errorf("error: %v", err)
} }
defer func() { defer func() {
+6 -10
View File
@@ -10,8 +10,6 @@ import (
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
nbnet "github.com/netbirdio/netbird/client/net"
) )
// PrepareSenderRawSocketIPv4 creates and configures a raw socket for sending IPv4 packets // PrepareSenderRawSocketIPv4 creates and configures a raw socket for sending IPv4 packets
@@ -60,14 +58,12 @@ func prepareSenderRawSocket(family int, isIPv4 bool) (net.PacketConn, error) {
return nil, fmt.Errorf("binding to lo interface failed: %w", err) return nil, fmt.Errorf("binding to lo interface failed: %w", err)
} }
// Set the fwmark on the socket. // The socket is bound to lo and only ever sends to the local WireGuard
err = nbnet.SetSocketOpt(fd) // instance, a destination the local routing table resolves without help, so
if err != nil { // it carries no fwmark. Staying unmarked also keeps these packets out of
if closeErr := syscall.Close(fd); closeErr != nil { // third-party NAT rules that match on marks: such a rule rewriting the
log.Warnf("failed to close raw socket fd: %v", closeErr) // source would make WireGuard adopt the rewritten address as the peer
} // endpoint.
return nil, fmt.Errorf("setting fwmark failed: %w", err)
}
// Convert the file descriptor to a PacketConn. // Convert the file descriptor to a PacketConn.
file := os.NewFile(uintptr(fd), fmt.Sprintf("fd %d", fd)) file := os.NewFile(uintptr(fd), fmt.Sprintf("fd %d", fd))
@@ -0,0 +1,77 @@
//go:build linux && !android && privileged
package rawsocket
import (
"net"
"syscall"
"testing"
"golang.org/x/sys/unix"
nbnet "github.com/netbirdio/netbird/client/net"
)
// The sender sockets must stay unmarked: a NAT rule matching on fwmark that
// rewrites the source of an injected packet makes WireGuard adopt the rewritten
// address as the peer endpoint.
func TestSenderRawSocketsCarryNoFwmark(t *testing.T) {
// the mark is only ever applied when advanced routing is available, so
// without it the assertion below would hold for the wrong reason
nbnet.Init()
if !nbnet.AdvancedRouting() {
t.Skip("advanced routing unsupported, the sockets carry no mark either way")
}
tests := []struct {
name string
prepare func() (net.PacketConn, error)
// the proxy treats the IPv6 socket as optional, so a host without IPv6
// is a reason to skip rather than to fail
optional bool
}{
{name: "IPv4", prepare: PrepareSenderRawSocketIPv4},
{name: "IPv6", prepare: PrepareSenderRawSocketIPv6, optional: true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
conn, err := tc.prepare()
if err != nil {
if tc.optional {
t.Skipf("prepare raw socket: %v", err)
}
t.Fatalf("prepare raw socket: %v", err)
}
defer func() {
if err := conn.Close(); err != nil {
t.Logf("close raw socket: %v", err)
}
}()
syscallConn, ok := conn.(syscall.Conn)
if !ok {
t.Fatalf("raw socket %T does not expose a syscall conn", conn)
}
raw, err := syscallConn.SyscallConn()
if err != nil {
t.Fatalf("syscall conn: %v", err)
}
var mark int
var markErr error
if err := raw.Control(func(fd uintptr) {
mark, markErr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_MARK)
}); err != nil {
t.Fatalf("control: %v", err)
}
if markErr != nil {
t.Fatalf("get SO_MARK: %v", markErr)
}
if mark != 0 {
t.Errorf("SO_MARK = %#x, want 0", mark)
}
})
}
}
+5 -5
View File
@@ -119,9 +119,9 @@ func testRedirectAs(t *testing.T, proxy Proxy, wgPort int, nbAddr, p2pEndpoint *
} }
defer relayConn.Close() defer relayConn.Close()
// Add TURN connection to proxy // Add relayed connection to proxy
if err := proxy.AddTurnConn(ctx, nbAddr, relayConn); err != nil { if err := proxy.AddRelayedConn(ctx, nbAddr, relayConn); err != nil {
t.Fatalf("failed to add TURN connection: %v", err) t.Fatalf("failed to add relayed connection: %v", err)
} }
defer func() { defer func() {
if err := proxy.CloseConn(); err != nil { if err := proxy.CloseConn(); err != nil {
@@ -304,8 +304,8 @@ func TestRedirectAs_Multiple_Switches(t *testing.T) {
Port: 38746, Port: 38746,
} }
if err := proxy.AddTurnConn(ctx, nbAddr, relayConn); err != nil { if err := proxy.AddRelayedConn(ctx, nbAddr, relayConn); err != nil {
t.Fatalf("failed to add TURN connection: %v", err) t.Fatalf("failed to add relayed connection: %v", err)
} }
defer func() { defer func() {
if err := proxy.CloseConn(); err != nil { if err := proxy.CloseConn(); err != nil {
+2 -2
View File
@@ -51,12 +51,12 @@ func NewWGUDPProxy(wgPort int, mtu uint16) *WGUDPProxy {
return p return p
} }
// AddTurnConn // AddRelayedConn dials the local WireGuard port and stores the relayed connection.
// The provided Context must be non-nil. If the context expires before // The provided Context must be non-nil. If the context expires before
// the connection is complete, an error is returned. Once successfully // the connection is complete, an error is returned. Once successfully
// connected, any expiration of the context will not affect the // connected, any expiration of the context will not affect the
// connection. // connection.
func (p *WGUDPProxy) AddTurnConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error { func (p *WGUDPProxy) AddRelayedConn(ctx context.Context, _ *net.UDPAddr, remoteConn net.Conn) error {
dialer := net.Dialer{} dialer := net.Dialer{}
localConn, err := dialer.DialContext(ctx, "udp", fmt.Sprintf(":%d", p.localWGListenPort)) localConn, err := dialer.DialContext(ctx, "udp", fmt.Sprintf(":%d", p.localWGListenPort))
if err != nil { if err != nil {
+7 -8
View File
@@ -116,11 +116,11 @@ func (d *DefaultManager) ApplyFiltering(networkMap *mgmProto.NetworkMap, dnsRout
// firewall state, so an identical hash means an identical resulting ruleset. // firewall state, so an identical hash means an identical resulting ruleset.
func (d *DefaultManager) firewallConfigHash(networkMap *mgmProto.NetworkMap, dnsRouteFeatureFlag bool) (uint64, error) { func (d *DefaultManager) firewallConfigHash(networkMap *mgmProto.NetworkMap, dnsRouteFeatureFlag bool) (uint64, error) {
return hashstructure.Hash(struct { return hashstructure.Hash(struct {
PeerRules []*mgmProto.FirewallRule PeerRules []*mgmProto.FirewallRule
PeerRulesIsEmpty bool PeerRulesIsEmpty bool
RouteRules []*mgmProto.RouteFirewallRule RouteRules []*mgmProto.RouteFirewallRule
RouteRulesIsEmpty bool RouteRulesIsEmpty bool
DNSRouteFeatureFlag bool DNSRouteFeatureFlag bool
}{ }{
PeerRules: networkMap.GetFirewallRules(), PeerRules: networkMap.GetFirewallRules(),
PeerRulesIsEmpty: networkMap.GetFirewallRulesIsEmpty(), PeerRulesIsEmpty: networkMap.GetFirewallRulesIsEmpty(),
@@ -144,13 +144,13 @@ func (d *DefaultManager) applyPeerACLs(networkMap *mgmProto.NetworkMap) {
log.Warn("this peer is connected to a NetBird Management service with an older version. Allowing all traffic from connected peers") log.Warn("this peer is connected to a NetBird Management service with an older version. Allowing all traffic from connected peers")
rules = append(rules, rules = append(rules,
&mgmProto.FirewallRule{ &mgmProto.FirewallRule{
PeerIP: "0.0.0.0", PeerIP: "0.0.0.0", //nolint:staticcheck
Direction: mgmProto.RuleDirection_IN, Direction: mgmProto.RuleDirection_IN,
Action: mgmProto.RuleAction_ACCEPT, Action: mgmProto.RuleAction_ACCEPT,
Protocol: mgmProto.RuleProtocol_ALL, Protocol: mgmProto.RuleProtocol_ALL,
}, },
&mgmProto.FirewallRule{ &mgmProto.FirewallRule{
PeerIP: "0.0.0.0", PeerIP: "0.0.0.0", //nolint:staticcheck
Direction: mgmProto.RuleDirection_OUT, Direction: mgmProto.RuleDirection_OUT,
Action: mgmProto.RuleAction_ACCEPT, Action: mgmProto.RuleAction_ACCEPT,
Protocol: mgmProto.RuleProtocol_ALL, Protocol: mgmProto.RuleProtocol_ALL,
@@ -407,7 +407,6 @@ func (d *DefaultManager) getRuleGroupingSelector(rule *mgmProto.FirewallRule) st
return fmt.Sprintf("%v:%v:%v:%s:%v", strconv.Itoa(int(rule.Direction)), rule.Action, rule.Protocol, rule.Port, rule.PortInfo) return fmt.Sprintf("%v:%v:%v:%s:%v", strconv.Itoa(int(rule.Direction)), rule.Action, rule.Protocol, rule.Port, rule.PortInfo)
} }
// extractRuleIP extracts the peer IP from a firewall rule. // extractRuleIP extracts the peer IP from a firewall rule.
// If sourcePrefixes is populated (new management), decode the first entry and use its address. // If sourcePrefixes is populated (new management), decode the first entry and use its address.
// Otherwise fall back to the deprecated PeerIP string field (old management). // Otherwise fall back to the deprecated PeerIP string field (old management).
+4 -4
View File
@@ -5,9 +5,9 @@ import (
"net/netip" "net/netip"
"testing" "testing"
"github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/client/firewall" "github.com/netbirdio/netbird/client/firewall"
"github.com/netbirdio/netbird/client/iface" "github.com/netbirdio/netbird/client/iface"
@@ -87,7 +87,7 @@ func TestDefaultManager(t *testing.T) {
networkMap.FirewallRules = append( networkMap.FirewallRules = append(
networkMap.FirewallRules, networkMap.FirewallRules,
&mgmProto.FirewallRule{ &mgmProto.FirewallRule{
PeerIP: "10.93.0.3", PeerIP: "10.93.0.3", //nolint:staticcheck
Direction: mgmProto.RuleDirection_IN, Direction: mgmProto.RuleDirection_IN,
Action: mgmProto.RuleAction_DROP, Action: mgmProto.RuleAction_DROP,
Protocol: mgmProto.RuleProtocol_ICMP, Protocol: mgmProto.RuleProtocol_ICMP,
@@ -556,12 +556,12 @@ func TestApplyFilteringSkipsUnchangedConfig(t *testing.T) {
func buildNetworkMap(peerRules, routeRules int) *mgmProto.NetworkMap { func buildNetworkMap(peerRules, routeRules int) *mgmProto.NetworkMap {
nm := &mgmProto.NetworkMap{ nm := &mgmProto.NetworkMap{
FirewallRulesIsEmpty: peerRules == 0, FirewallRulesIsEmpty: peerRules == 0,
RoutesFirewallRulesIsEmpty: routeRules == 0, RoutesFirewallRulesIsEmpty: routeRules == 0,
} }
for i := range peerRules { for i := range peerRules {
nm.FirewallRules = append(nm.FirewallRules, &mgmProto.FirewallRule{ nm.FirewallRules = append(nm.FirewallRules, &mgmProto.FirewallRule{
PeerIP: fmt.Sprintf("10.%d.%d.%d", i>>16&0xff, i>>8&0xff, i&0xff), PeerIP: fmt.Sprintf("10.%d.%d.%d", i>>16&0xff, i>>8&0xff, i&0xff), //nolint:staticcheck
Direction: mgmProto.RuleDirection_IN, Direction: mgmProto.RuleDirection_IN,
Action: mgmProto.RuleAction_ACCEPT, Action: mgmProto.RuleAction_ACCEPT,
Protocol: mgmProto.RuleProtocol_TCP, Protocol: mgmProto.RuleProtocol_TCP,
+1 -1
View File
@@ -7,7 +7,7 @@ package mocks
import ( import (
reflect "reflect" reflect "reflect"
gomock "github.com/golang/mock/gomock" gomock "go.uber.org/mock/gomock"
wgdevice "golang.zx2c4.com/wireguard/device" wgdevice "golang.zx2c4.com/wireguard/device"
"github.com/netbirdio/netbird/client/iface/device" "github.com/netbirdio/netbird/client/iface/device"
+68 -104
View File
@@ -2,6 +2,7 @@ package internal
import ( import (
"context" "context"
"maps"
"os" "os"
"strconv" "strconv"
"sync" "sync"
@@ -14,6 +15,7 @@ import (
"github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/peerstore" "github.com/netbirdio/netbird/client/internal/peerstore"
"github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/route"
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
) )
// lazyForce is the resolved local decision for lazy connections, layered above the // lazyForce is the resolved local decision for lazy connections, layered above the
@@ -37,11 +39,13 @@ const (
// The only exception is ActivatePeer, which is safe for concurrent use so the // The only exception is ActivatePeer, which is safe for concurrent use so the
// DNS warm-up path can call it without contending on the engine mutex. // DNS warm-up path can call it without contending on the engine mutex.
type ConnMgr struct { type ConnMgr struct {
peerStore *peerstore.Store peerStore *peerstore.Store
statusRecorder *peer.Status statusRecorder *peer.Status
iface lazyconn.WGIface iface lazyconn.WGIface
force lazyForce force lazyForce
rosenpassEnabled bool // remoteLazyEnabled caches the account-wide lazy feature flag from management.
// It is the default for peers that do not carry a per-peer lazy hint.
remoteLazyEnabled bool
lazyConnMgr *manager.Manager lazyConnMgr *manager.Manager
// lazyConnMgrMu guards the lazyConnMgr pointer for readers outside the // lazyConnMgrMu guards the lazyConnMgr pointer for readers outside the
@@ -53,6 +57,10 @@ type ConnMgr struct {
// (re)armed (Mode A at arm time). Injected by the engine; nil disables the reconcile. // (re)armed (Mode A at arm time). Injected by the engine; nil disables the reconcile.
reconcileRoutedIPs func(peerKey string) error reconcileRoutedIPs func(peerKey string) error
// appliedExcludeList is the exclude set last handed to the lazy manager, kept so an
// unchanged set on the next sync skips the O(n) reconciliation.
appliedExcludeList map[string]bool
wg sync.WaitGroup wg sync.WaitGroup
lazyCtx context.Context lazyCtx context.Context
lazyCtxCancel context.CancelFunc lazyCtxCancel context.CancelFunc
@@ -66,78 +74,59 @@ func (e *ConnMgr) SetRoutedIPsReconciler(fn func(peerKey string) error) {
func NewConnMgr(engineConfig *EngineConfig, statusRecorder *peer.Status, peerStore *peerstore.Store, iface lazyconn.WGIface) *ConnMgr { func NewConnMgr(engineConfig *EngineConfig, statusRecorder *peer.Status, peerStore *peerstore.Store, iface lazyconn.WGIface) *ConnMgr {
e := &ConnMgr{ e := &ConnMgr{
peerStore: peerStore, peerStore: peerStore,
statusRecorder: statusRecorder, statusRecorder: statusRecorder,
iface: iface, iface: iface,
force: resolveLazyForce(engineConfig.LazyConnection), force: resolveLazyForce(engineConfig.LazyConnection),
rosenpassEnabled: engineConfig.RosenpassEnabled,
} }
return e return e
} }
// Start initializes the connection manager. It starts the lazy connection manager when a // Start initializes the connection manager. The lazy connection manager always runs so that
// local override forces it on; with no local override it waits for the management feature flag. // per-peer lazy defaults (e.g. proxy peers) work even when the account flag is off; the
// account flag and the local override decide the default lazy state per peer (see
// PeerLazyDefault). Rosenpass peers stay lazy-capable too: their connections just never idle
// on their own, since rosenpass rekey traffic keeps them active.
func (e *ConnMgr) Start(ctx context.Context) { func (e *ConnMgr) Start(ctx context.Context) {
if e.lazyConnMgr != nil { if e.lazyConnMgr != nil {
log.Errorf("lazy connection manager is already started") log.Errorf("lazy connection manager is already started")
return return
} }
switch e.force {
case lazyForceOff:
log.Infof("lazy connection manager is disabled by local override (%s or MDM policy)", lazyconn.EnvLazyConn)
e.statusRecorder.UpdateLazyConnection(false)
return
case lazyForceNone:
log.Infof("lazy connection manager is managed by the management feature flag")
e.statusRecorder.UpdateLazyConnection(false)
return
}
if e.rosenpassEnabled {
log.Warnf("rosenpass connection manager is enabled, lazy connection manager will not be started")
e.statusRecorder.UpdateLazyConnection(false)
return
}
e.initLazyManager(ctx) e.initLazyManager(ctx)
e.statusRecorder.UpdateLazyConnection(true) e.statusRecorder.UpdateLazyConnection(e.PeerLazyDefault(mgmProto.LazyState_LazyStateDefault))
} }
// UpdatedRemoteFeatureFlag is called when the remote feature flag is updated. // UpdatedRemoteFeatureFlag caches the account-wide lazy feature flag. The manager itself is
// If enabled, it initializes the lazy connection manager and start it. Do not need to call Start() again. // not started or stopped here; the per-sync exclude-list reconciliation moves normal peers
// If disabled, then it closes the lazy connection manager and open the connections to all peers. // between the lazy and always-active sets when the flag flips.
func (e *ConnMgr) UpdatedRemoteFeatureFlag(ctx context.Context, enabled bool) error { func (e *ConnMgr) UpdatedRemoteFeatureFlag(_ context.Context, enabled bool) error {
// a local override (NB_LAZY_CONN or local config) takes precedence over management e.remoteLazyEnabled = enabled
if e.force != lazyForceNone { if e.isStartedWithLazyMgr() {
return nil e.statusRecorder.UpdateLazyConnection(e.PeerLazyDefault(mgmProto.LazyState_LazyStateDefault))
}
return nil
}
// PeerLazyDefault reports whether a peer should be lazy. The local override
// (NB_LAZY_CONN/MDM) wins over everything; without a local override the
// management per-peer state applies (LazyStateLazy/Eager force the decision),
// and LazyStateDefault follows the account-wide flag.
func (e *ConnMgr) PeerLazyDefault(state mgmProto.LazyState) bool {
switch e.force {
case lazyForceOn:
return true
case lazyForceOff:
return false
} }
if enabled { switch state {
// if the lazy connection manager is already started, do not start it again case mgmProto.LazyState_LazyStateLazy:
if e.lazyConnMgr != nil { return true
return nil case mgmProto.LazyState_LazyStateEager:
} return false
default:
if e.rosenpassEnabled { return e.remoteLazyEnabled
log.Infof("rosenpass connection manager is enabled, lazy connection manager will not be started")
e.statusRecorder.UpdateLazyConnection(false)
return nil
}
log.Infof("lazy connection manager is enabled by the management feature flag")
e.initLazyManager(ctx)
e.statusRecorder.UpdateLazyConnection(true)
return e.addPeersToLazyConnManager()
} else {
if e.lazyConnMgr == nil {
e.statusRecorder.UpdateLazyConnection(false)
return nil
}
log.Infof("lazy connection manager is disabled by management feature flag")
e.closeManager(ctx)
e.statusRecorder.UpdateLazyConnection(false)
return nil
} }
} }
@@ -157,6 +146,13 @@ func (e *ConnMgr) SetExcludeList(ctx context.Context, peerIDs map[string]bool) {
return return
} }
// The exclude set is recomputed every sync but rarely changes; skip the O(n)
// store lookups and reconciliation when it matches what was already applied.
if maps.Equal(peerIDs, e.appliedExcludeList) {
return
}
e.appliedExcludeList = maps.Clone(peerIDs)
excludedPeers := make([]lazyconn.PeerConfig, 0, len(peerIDs)) excludedPeers := make([]lazyconn.PeerConfig, 0, len(peerIDs))
for peerID := range peerIDs { for peerID := range peerIDs {
@@ -192,12 +188,16 @@ func (e *ConnMgr) SetExcludeList(ctx context.Context, peerIDs map[string]bool) {
} }
} }
func (e *ConnMgr) AddPeerConn(ctx context.Context, peerKey string, conn *peer.Conn) (exists bool) { // AddPeerConn registers a peer connection. permanent requests an always-active connection
// (the peer belongs to the exclude set: a forwarder, or a peer that is not lazy by policy).
// Non-permanent peers are handed to the lazy manager. The subsequent SetExcludeList call
// reconciles membership for existing peers across flag flips.
func (e *ConnMgr) AddPeerConn(ctx context.Context, peerKey string, conn *peer.Conn, permanent bool) (exists bool) {
if success := e.peerStore.AddPeerConn(peerKey, conn); !success { if success := e.peerStore.AddPeerConn(peerKey, conn); !success {
return true return true
} }
if !e.isStartedWithLazyMgr() { if !e.isStartedWithLazyMgr() || permanent {
if err := conn.Open(ctx); err != nil { if err := conn.Open(ctx); err != nil {
conn.Log.Errorf("failed to open connection: %v", err) conn.Log.Errorf("failed to open connection: %v", err)
} }
@@ -296,6 +296,8 @@ func (e *ConnMgr) Close() {
e.lazyConnMgrMu.Lock() e.lazyConnMgrMu.Lock()
e.lazyConnMgr = nil e.lazyConnMgr = nil
e.lazyConnMgrMu.Unlock() e.lazyConnMgrMu.Unlock()
e.appliedExcludeList = nil
} }
func (e *ConnMgr) initLazyManager(engineCtx context.Context) { func (e *ConnMgr) initLazyManager(engineCtx context.Context) {
@@ -309,6 +311,8 @@ func (e *ConnMgr) initLazyManager(engineCtx context.Context) {
e.lazyCtx, e.lazyCtxCancel = context.WithCancel(engineCtx) e.lazyCtx, e.lazyCtxCancel = context.WithCancel(engineCtx)
e.lazyConnMgrMu.Unlock() e.lazyConnMgrMu.Unlock()
e.appliedExcludeList = nil
e.wg.Add(1) e.wg.Add(1)
go func() { go func() {
defer e.wg.Done() defer e.wg.Done()
@@ -316,46 +320,6 @@ func (e *ConnMgr) initLazyManager(engineCtx context.Context) {
}() }()
} }
func (e *ConnMgr) addPeersToLazyConnManager() error {
peers := e.peerStore.PeersPubKey()
lazyPeerCfgs := make([]lazyconn.PeerConfig, 0, len(peers))
for _, peerID := range peers {
var peerConn *peer.Conn
var exists bool
if peerConn, exists = e.peerStore.PeerConn(peerID); !exists {
log.Warnf("failed to find peer conn for peerID: %s", peerID)
continue
}
lazyPeerCfg := lazyconn.PeerConfig{
PublicKey: peerID,
AllowedIPs: peerConn.WgConfig().AllowedIps,
PeerConnID: peerConn.ConnID(),
Log: peerConn.Log,
}
lazyPeerCfgs = append(lazyPeerCfgs, lazyPeerCfg)
}
return e.lazyConnMgr.AddActivePeers(lazyPeerCfgs)
}
func (e *ConnMgr) closeManager(ctx context.Context) {
if e.lazyConnMgr == nil {
return
}
e.lazyCtxCancel()
e.wg.Wait()
e.lazyConnMgrMu.Lock()
e.lazyConnMgr = nil
e.lazyConnMgrMu.Unlock()
for _, peerID := range e.peerStore.PeersPubKey() {
e.peerStore.PeerConnOpen(ctx, peerID)
}
}
func (e *ConnMgr) isStartedWithLazyMgr() bool { func (e *ConnMgr) isStartedWithLazyMgr() bool {
return e.lazyConnMgr != nil && e.lazyCtxCancel != nil return e.lazyConnMgr != nil && e.lazyCtxCancel != nil
} }
+88
View File
@@ -16,6 +16,7 @@ import (
"github.com/netbirdio/netbird/client/internal/peer" "github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/peerstore" "github.com/netbirdio/netbird/client/internal/peerstore"
"github.com/netbirdio/netbird/monotime" "github.com/netbirdio/netbird/monotime"
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
) )
func TestResolveLazyForce(t *testing.T) { func TestResolveLazyForce(t *testing.T) {
@@ -138,4 +139,91 @@ func TestInactivityThresholdEnv(t *testing.T) {
} }
} }
func TestPeerLazyDefault(t *testing.T) {
tests := []struct {
name string
force lazyForce
remoteEnabled bool
state mgmProto.LazyState
want bool
}{
{name: "force on wins over eager state", force: lazyForceOn, state: mgmProto.LazyState_LazyStateEager, want: true},
{name: "force off wins over lazy state", force: lazyForceOff, remoteEnabled: true, state: mgmProto.LazyState_LazyStateLazy, want: false},
{name: "none, default, account off -> active", force: lazyForceNone, state: mgmProto.LazyState_LazyStateDefault, want: false},
{name: "none, default, account on -> lazy", force: lazyForceNone, remoteEnabled: true, state: mgmProto.LazyState_LazyStateDefault, want: true},
{name: "none, lazy state, account off -> lazy", force: lazyForceNone, state: mgmProto.LazyState_LazyStateLazy, want: true},
{name: "none, eager state, account on -> active", force: lazyForceNone, remoteEnabled: true, state: mgmProto.LazyState_LazyStateEager, want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
e := &ConnMgr{force: tt.force, remoteLazyEnabled: tt.remoteEnabled}
if got := e.PeerLazyDefault(tt.state); got != tt.want {
t.Fatalf("PeerLazyDefault(%v) = %v, want %v", tt.state, got, tt.want)
}
})
}
}
func durPtr(d time.Duration) *time.Duration { return &d } func durPtr(d time.Duration) *time.Duration { return &d }
// TestToExcludedLazyPeers covers the per-peer lazy classification (proxy vs
// normal, across the force/account-flag matrix). Forwarder-target exclusion is
// covered by TestToExcludedLazyPeers_ForwardTarget.
func TestToExcludedLazyPeers(t *testing.T) {
const (
normalKey = "normal"
lazyKey = "lazy-state"
eagerKey = "eager-state"
)
peers := []*mgmProto.RemotePeerConfig{
{WgPubKey: normalKey, AllowedIps: []string{"100.64.0.1/32"}},
{WgPubKey: lazyKey, AllowedIps: []string{"100.64.0.2/32"}, LazyState: mgmProto.LazyState_LazyStateLazy},
{WgPubKey: eagerKey, AllowedIps: []string{"100.64.0.3/32"}, LazyState: mgmProto.LazyState_LazyStateEager},
}
tests := []struct {
name string
force lazyForce
remoteEnabled bool
want map[string]bool
}{
{
name: "account off: lazy-state peer lazy, normal + eager active",
force: lazyForceNone, remoteEnabled: false,
want: map[string]bool{normalKey: true, eagerKey: true},
},
{
name: "account on: only eager-state peer active",
force: lazyForceNone, remoteEnabled: true,
want: map[string]bool{eagerKey: true},
},
{
name: "force off: everything active",
force: lazyForceOff, remoteEnabled: true,
want: map[string]bool{normalKey: true, lazyKey: true, eagerKey: true},
},
{
name: "force on: nothing active",
force: lazyForceOn, remoteEnabled: false,
want: map[string]bool{},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
e := &Engine{connMgr: &ConnMgr{force: tt.force, remoteLazyEnabled: tt.remoteEnabled}}
got := e.toExcludedLazyPeers(peers)
if len(got) != len(tt.want) {
t.Fatalf("toExcludedLazyPeers() = %v, want %v", got, tt.want)
}
for k := range tt.want {
if !got[k] {
t.Fatalf("expected peer %s excluded, got %v", k, got)
}
}
})
}
}
+45 -6
View File
@@ -39,6 +39,7 @@ import (
"github.com/netbirdio/netbird/client/internal/updater" "github.com/netbirdio/netbird/client/internal/updater"
"github.com/netbirdio/netbird/client/internal/updater/installer" "github.com/netbirdio/netbird/client/internal/updater/installer"
nbnet "github.com/netbirdio/netbird/client/net" nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/netevents"
cProto "github.com/netbirdio/netbird/client/proto" cProto "github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/ssh" "github.com/netbirdio/netbird/client/ssh"
sshconfig "github.com/netbirdio/netbird/client/ssh/config" sshconfig "github.com/netbirdio/netbird/client/ssh/config"
@@ -72,18 +73,31 @@ type ConnectClient struct {
fileDropManager *filedrop.Manager fileDropManager *filedrop.Manager
persistSyncResponse bool persistSyncResponse bool
// netMgr gates every reconnection loop on OS-reported network
// availability and sweeps connections on network change.
netMgr *netevents.Manager
}
// ConnectClientOption configures optional ConnectClient behavior.
type ConnectClientOption func(*ConnectClient)
// WithNetEvents injects the OS network event handling.
func WithNetEvents(events *netevents.Manager) ConnectClientOption {
return func(c *ConnectClient) { c.netMgr = events }
} }
func NewConnectClient( func NewConnectClient(
ctx context.Context, ctx context.Context,
config *profilemanager.Config, config *profilemanager.Config,
statusRecorder *peer.Status, statusRecorder *peer.Status,
opts ...ConnectClientOption,
) *ConnectClient { ) *ConnectClient {
// Derive the run context here so Stop owns the cancel that unblocks the run // Derive the run context here so Stop owns the cancel that unblocks the run
// loop. runCancel is set once at construction, so Stop can call it without // loop. runCancel is set once at construction, so Stop can call it without
// racing the run loop's startup. Callers therefore need not cancel before Stop. // racing the run loop's startup. Callers therefore need not cancel before Stop.
runCtx, runCancel := context.WithCancel(ctx) runCtx, runCancel := context.WithCancel(ctx)
return &ConnectClient{ c := &ConnectClient{
ctx: runCtx, ctx: runCtx,
runCancel: runCancel, runCancel: runCancel,
runExited: make(chan struct{}), runExited: make(chan struct{}),
@@ -91,6 +105,10 @@ func NewConnectClient(
statusRecorder: statusRecorder, statusRecorder: statusRecorder,
engineMutex: sync.Mutex{}, engineMutex: sync.Mutex{},
} }
for _, opt := range opts {
opt(c)
}
return c
} }
func (c *ConnectClient) SetUpdateManager(um *updater.Manager) { func (c *ConnectClient) SetUpdateManager(um *updater.Manager) {
@@ -282,6 +300,13 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
return nil return nil
} }
// suspend connection attempts while the OS reports no usable network
if waited, err := c.netMgr.Wait(c.ctx); err != nil {
return nil
} else if waited {
backOff.Reset()
}
state.Set(StatusConnecting) state.Set(StatusConnecting)
engineCtx, cancel := context.WithCancel(c.ctx) engineCtx, cancel := context.WithCancel(c.ctx)
@@ -293,7 +318,8 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
}() }()
log.Debugf("connecting to the Management service %s", c.config.ManagementURL.Host) log.Debugf("connecting to the Management service %s", c.config.ManagementURL.Host)
mgmClient, err := mgm.NewClient(engineCtx, c.config.ManagementURL.Host, myPrivateKey, mgmTlsEnabled) mgmClient, err := mgm.NewClient(engineCtx, c.config.ManagementURL.Host, myPrivateKey, mgmTlsEnabled,
mgm.WithNetEvents(c.netMgr))
if err != nil { if err != nil {
// On daemon shutdown / Down() the parent context is cancelled // On daemon shutdown / Down() the parent context is cancelled
// and the dial fails with "context canceled". Wrapping that // and the dial fails with "context canceled". Wrapping that
@@ -368,7 +394,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
}() }()
// with the global Netbird config in hand connect (just a connection, no stream yet) Signal // with the global Netbird config in hand connect (just a connection, no stream yet) Signal
signalClient, err := connectToSignal(engineCtx, loginResp.GetNetbirdConfig(), myPrivateKey) signalClient, err := connectToSignal(engineCtx, loginResp.GetNetbirdConfig(), myPrivateKey, c.netMgr)
if err != nil { if err != nil {
log.Error(err) log.Error(err)
return wrapErr(err) return wrapErr(err)
@@ -404,7 +430,8 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
engineConfig.StateDir = filepath.Dir(path) engineConfig.StateDir = filepath.Dir(path)
} }
relayManager := relayClient.NewManager(engineCtx, relayURLs, myPrivateKey.PublicKey().String(), engineConfig.MTU) relayManager := relayClient.NewManager(engineCtx, relayURLs, myPrivateKey.PublicKey().String(), engineConfig.MTU,
relayClient.WithNetEvents(c.netMgr))
c.statusRecorder.SetRelayMgr(relayManager) c.statusRecorder.SetRelayMgr(relayManager)
if len(relayURLs) > 0 { if len(relayURLs) > 0 {
if token != nil { if token != nil {
@@ -433,6 +460,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
ClientMetrics: c.clientMetrics, ClientMetrics: c.clientMetrics,
MetricsCtx: c.ctx, MetricsCtx: c.ctx,
FileDrop: c.fileDropManager, FileDrop: c.fileDropManager,
NetMgr: c.netMgr,
}, mobileDependency) }, mobileDependency)
engine.SetSyncResponsePersistence(c.persistSyncResponse) engine.SetSyncResponsePersistence(c.persistSyncResponse)
c.engine = engine c.engine = engine
@@ -489,6 +517,16 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
// status stream stuck at Connecting. // status stream stuck at Connecting.
err = backoff.Retry(operation, backoff.WithContext(backOff, c.ctx)) err = backoff.Retry(operation, backoff.WithContext(backOff, c.ctx))
if err != nil { if err != nil {
// Once the client context is cancelled backoff.WithContext surfaces the
// bare context error, and any attempt torn down mid-flight reports the
// same. That cancellation is the caller asking us to stop (Stop, Down or
// an engine restart), so exit cleanly instead of handing back a failure
// the caller would have to distinguish from a real one.
if c.ctx.Err() != nil && errors.Is(err, context.Canceled) {
log.Info("exiting client retry loop, context cancelled")
return nil
}
log.Debugf("exiting client retry loop due to unrecoverable error: %s", err) log.Debugf("exiting client retry loop due to unrecoverable error: %s", err)
if s, ok := gstatus.FromError(err); ok && (s.Code() == codes.PermissionDenied) { if s, ok := gstatus.FromError(err); ok && (s.Code() == codes.PermissionDenied) {
state.Set(StatusNeedsLogin) state.Set(StatusNeedsLogin)
@@ -682,7 +720,7 @@ func selectMTU(localMTU uint16, peerMTU int32) uint16 {
} }
// connectToSignal creates Signal Service client and established a connection // connectToSignal creates Signal Service client and established a connection
func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourPrivateKey wgtypes.Key) (*signal.GrpcClient, error) { func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourPrivateKey wgtypes.Key, netMgr *netevents.Manager) (*signal.GrpcClient, error) {
var sigTLSEnabled bool var sigTLSEnabled bool
if wtConfig.Signal.Protocol == mgmProto.HostConfig_HTTPS { if wtConfig.Signal.Protocol == mgmProto.HostConfig_HTTPS {
sigTLSEnabled = true sigTLSEnabled = true
@@ -690,7 +728,8 @@ func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourP
sigTLSEnabled = false sigTLSEnabled = false
} }
signalClient, err := signal.NewClient(ctx, wtConfig.Signal.Uri, ourPrivateKey, sigTLSEnabled) signalClient, err := signal.NewClient(ctx, wtConfig.Signal.Uri, ourPrivateKey, sigTLSEnabled,
signal.WithNetEvents(netMgr))
if err != nil { if err != nil {
log.Errorf("error while connecting to the Signal Exchange Service %s: %s", wtConfig.Signal.Uri, err) log.Errorf("error while connecting to the Signal Exchange Service %s: %s", wtConfig.Signal.Uri, err)
return nil, gstatus.Errorf(codes.FailedPrecondition, "failed connecting to Signal Service : %s", err) return nil, gstatus.Errorf(codes.FailedPrecondition, "failed connecting to Signal Service : %s", err)
+17
View File
@@ -0,0 +1,17 @@
package daemonaddr
import "strings"
// CarriesIdentity reports whether the control channel at addr conveys the
// connecting process's identity to the daemon. A Unix socket carries peer
// credentials and a named pipe carries the client's token. Nothing else does, TCP
// included, and there the daemon can authorize a privileged operation for nobody
// at all: see ResolveDaemonAddr, which says as much to anyone still reaching the
// Windows daemon on the address it served before it had a pipe.
//
// A client uses this to tell whether becoming privileged would get it anywhere.
// It answers from the scheme and nothing else, so an address it does not
// recognise counts as carrying no identity.
func CarriesIdentity(addr string) bool {
return strings.HasPrefix(addr, "unix://") || strings.HasPrefix(addr, pipeScheme)
}
@@ -0,0 +1,29 @@
package daemonaddr
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestCarriesIdentity(t *testing.T) {
tests := []struct {
addr string
want bool
}{
{"unix:///var/run/netbird.sock", true},
{"unix:///var/run/netbird/default.sock", true},
{"npipe://netbird", true},
{`npipe://\\.\pipe\ProtectedPrefix\Administrators\netbird`, true},
{"tcp://127.0.0.1:41731", false},
{"tcp://localhost:41731", false},
{"", false},
{"/var/run/netbird.sock", false},
}
for _, tt := range tests {
t.Run(tt.addr, func(t *testing.T) {
assert.Equal(t, tt.want, CarriesIdentity(tt.addr), "address %q", tt.addr)
})
}
}
+63 -37
View File
@@ -35,6 +35,8 @@ var (
// exported so a diagnostic reader reports the same locations that are written. // exported so a diagnostic reader reports the same locations that are written.
const ( const (
// NRPTKeyPrefix starts the name of every NRPT rule key this client creates. // NRPTKeyPrefix starts the name of every NRPT rule key this client creates.
// Older versions used different layouts under the same prefix: a single
// unsuffixed key, then one key per domain, now one key per batch of domains.
NRPTKeyPrefix = "NetBird-Match" NRPTKeyPrefix = "NetBird-Match"
// DNSPolicyConfigRoot holds the NRPT rules of the local policy store. // DNSPolicyConfigRoot holds the NRPT rules of the local policy store.
@@ -89,7 +91,6 @@ type registryConfigurator struct {
guid string guid string
routingAll bool routingAll bool
gpo bool gpo bool
nrptEntryCount int
origNameservers []netip.Addr origNameservers []netip.Addr
} }
@@ -322,14 +323,9 @@ func (r *registryConfigurator) applyDNSConfig(config HostDNSConfig, stateManager
} }
if len(matchDomains) != 0 { if len(matchDomains) != 0 {
count, err := r.addDNSMatchPolicy(matchDomains, config.ServerIP) if err := r.addDNSMatchPolicy(matchDomains, config.ServerIP); err != nil {
// Update count even on error to ensure cleanup covers partially created rules
r.nrptEntryCount = count
if err != nil {
return fmt.Errorf("add dns match policy: %w", err) return fmt.Errorf("add dns match policy: %w", err)
} }
} else {
r.nrptEntryCount = 0
} }
r.updateState(stateManager) r.updateState(stateManager)
@@ -345,9 +341,8 @@ func (r *registryConfigurator) applyDNSConfig(config HostDNSConfig, stateManager
func (r *registryConfigurator) updateState(stateManager *statemanager.Manager) { func (r *registryConfigurator) updateState(stateManager *statemanager.Manager) {
if err := stateManager.UpdateState(&ShutdownState{ if err := stateManager.UpdateState(&ShutdownState{
Guid: r.guid, Guid: r.guid,
GPO: r.gpo, GPO: r.gpo,
NRPTEntryCount: r.nrptEntryCount,
}); err != nil { }); err != nil {
log.Errorf("failed to update shutdown state: %s", err) log.Errorf("failed to update shutdown state: %s", err)
} }
@@ -362,7 +357,7 @@ func (r *registryConfigurator) addDNSSetupForAll(ip netip.Addr) error {
return nil return nil
} }
func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr) (int, error) { func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr) error {
// if the gpo key is present, we need to put our DNS settings there, otherwise our config might be ignored // if the gpo key is present, we need to put our DNS settings there, otherwise our config might be ignored
// see https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-gpnrpt/8cc31cb9-20cb-4140-9e85-3e08703b4745 // see https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-gpnrpt/8cc31cb9-20cb-4140-9e85-3e08703b4745
@@ -379,19 +374,17 @@ func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr
gpoPath := fmt.Sprintf("%s-%d", gpoDnsPolicyConfigMatchPath, ruleIndex) gpoPath := fmt.Sprintf("%s-%d", gpoDnsPolicyConfigMatchPath, ruleIndex)
if err := r.configureDNSPolicy(localPath, batchDomains, ip); err != nil { if err := r.configureDNSPolicy(localPath, batchDomains, ip); err != nil {
return ruleIndex, fmt.Errorf("configure DNS Local policy for rule %d: %w", ruleIndex, err) return fmt.Errorf("configure DNS Local policy for rule %d: %w", ruleIndex, err)
} }
// Increment immediately so the caller's cleanup path knows about this rule
ruleIndex++
if r.gpo { if r.gpo {
if err := r.configureDNSPolicy(gpoPath, batchDomains, ip); err != nil { if err := r.configureDNSPolicy(gpoPath, batchDomains, ip); err != nil {
return ruleIndex, fmt.Errorf("configure gpo DNS policy for rule %d: %w", ruleIndex-1, err) return fmt.Errorf("configure gpo DNS policy for rule %d: %w", ruleIndex, err)
} }
} }
log.Debugf("added NRPT rule %d with %d domains", ruleIndex-1, len(batchDomains)) log.Debugf("added NRPT rule %d with %d domains", ruleIndex, len(batchDomains))
ruleIndex++
} }
if r.gpo { if r.gpo {
@@ -401,7 +394,7 @@ func (r *registryConfigurator) addDNSMatchPolicy(domains []string, ip netip.Addr
} }
log.Infof("added %d NRPT rules for %d domains", ruleIndex, len(domains)) log.Infof("added %d NRPT rules for %d domains", ruleIndex, len(domains))
return ruleIndex, nil return nil
} }
func (r *registryConfigurator) configureDNSPolicy(policyPath string, domains []string, ip netip.Addr) error { func (r *registryConfigurator) configureDNSPolicy(policyPath string, domains []string, ip netip.Addr) error {
@@ -466,7 +459,7 @@ func (r *registryConfigurator) flushDNSCache() {
ret, _, err := dnsFlushResolverCacheFn.Call() ret, _, err := dnsFlushResolverCacheFn.Call()
if ret == 0 { if ret == 0 {
if err != nil && !errors.Is(err, syscall.Errno(0)) { if !errors.Is(err, syscall.Errno(0)) {
log.Errorf("DnsFlushResolverCache failed: %v", err) log.Errorf("DnsFlushResolverCache failed: %v", err)
return return
} }
@@ -534,28 +527,28 @@ func (r *registryConfigurator) restoreHostDNS() error {
return nil return nil
} }
// removeDNSMatchPolicies deletes every NRPT rule this client may have created,
// from the local and the GPO policy store. The rules are found by enumerating
// the registry, the only authoritative record of what was written. Cleanup must
// not depend on a rule count: the in-memory one is scoped to a single
// registryConfigurator and the persisted one is deleted on every clean
// disconnect, and a rule left behind keeps resolving names over an interface
// that is gone, until reboot discards the volatile key.
func (r *registryConfigurator) removeDNSMatchPolicies() error { func (r *registryConfigurator) removeDNSMatchPolicies() error {
var merr *multierror.Error var merr *multierror.Error
// Try to remove the base entries (for backward compatibility) for _, root := range []string{DNSPolicyConfigRoot, GPODNSPolicyConfigRoot} {
if err := removeRegistryKeyFromDNSPolicyConfig(dnsPolicyConfigMatchPath); err != nil { names, err := listNRPTRuleKeys(root)
merr = multierror.Append(merr, fmt.Errorf("remove local base entry: %w", err)) if err != nil {
} merr = multierror.Append(merr, fmt.Errorf("list rule keys under %s: %w", root, err))
continue
if err := removeRegistryKeyFromDNSPolicyConfig(gpoDnsPolicyConfigMatchPath); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove GPO base entry: %w", err))
}
for i := 0; i < r.nrptEntryCount; i++ {
localPath := fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i)
gpoPath := fmt.Sprintf("%s-%d", gpoDnsPolicyConfigMatchPath, i)
if err := removeRegistryKeyFromDNSPolicyConfig(localPath); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove local entry %d: %w", i, err))
} }
if err := removeRegistryKeyFromDNSPolicyConfig(gpoPath); err != nil { for _, name := range names {
merr = multierror.Append(merr, fmt.Errorf("remove GPO entry %d: %w", i, err)) path := root + `\` + name
if err := removeRegistryKeyFromDNSPolicyConfig(path); err != nil {
merr = multierror.Append(merr, fmt.Errorf("remove entry %s: %w", path, err))
}
} }
} }
@@ -570,6 +563,39 @@ func (r *registryConfigurator) restoreUncleanShutdownDNS() error {
return r.restoreHostDNS() return r.restoreHostDNS()
} }
// listNRPTRuleKeys returns the names of our NRPT rule keys under a policy store
// root. An absent root holds nothing to clean up, which is the normal state of
// the GPO store on a machine without DNS Client policy.
func listNRPTRuleKeys(root string) ([]string, error) {
k, err := registry.OpenKey(registry.LOCAL_MACHINE, root, registry.ENUMERATE_SUB_KEYS)
switch {
case errors.Is(err, registry.ErrNotExist), errors.Is(err, syscall.ERROR_PATH_NOT_FOUND):
// the GPO store is absent on a machine without DNS client policy
log.Debugf("HKEY_LOCAL_MACHINE\\%s does not exist", root)
return nil, nil
case err != nil:
// any other failure has to reach the caller: reporting no rules would
// report a successful cleanup while leaving the rules in place
return nil, fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", root, err)
}
defer closer(k)
names, err := k.ReadSubKeyNames(-1)
if err != nil {
return nil, fmt.Errorf("read subkey names: %w", err)
}
var ruleKeys []string
for _, name := range names {
// registry key names are case insensitive
if strings.HasPrefix(strings.ToLower(name), strings.ToLower(NRPTKeyPrefix)) {
ruleKeys = append(ruleKeys, name)
}
}
return ruleKeys, nil
}
func removeRegistryKeyFromDNSPolicyConfig(regKeyPath string) error { func removeRegistryKeyFromDNSPolicyConfig(regKeyPath string) error {
k, err := registry.OpenKey(registry.LOCAL_MACHINE, regKeyPath, registry.QUERY_VALUE) k, err := registry.OpenKey(registry.LOCAL_MACHINE, regKeyPath, registry.QUERY_VALUE)
if err != nil { if err != nil {
@@ -601,7 +627,7 @@ func refreshGroupPolicy() error {
) )
if ret == 0 { if ret == 0 {
if err != nil && !errors.Is(err, syscall.Errno(0)) { if !errors.Is(err, syscall.Errno(0)) {
return fmt.Errorf("RefreshPolicyEx failed: %w", err) return fmt.Errorf("RefreshPolicyEx failed: %w", err)
} }
return fmt.Errorf("RefreshPolicyEx failed") return fmt.Errorf("RefreshPolicyEx failed")
+63 -7
View File
@@ -25,7 +25,7 @@ func TestNRPTEntriesCleanupOnConfigChange(t *testing.T) {
// Create a test interface registry key so updateSearchDomains doesn't fail // Create a test interface registry key so updateSearchDomains doesn't fail
testGUID := "{12345678-1234-1234-1234-123456789ABC}" testGUID := "{12345678-1234-1234-1234-123456789ABC}"
interfacePath := `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\` + testGUID interfacePath := InterfaceConfigPath + `\` + testGUID
testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE) testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE)
require.NoError(t, err, "Should create test interface registry key") require.NoError(t, err, "Should create test interface registry key")
testKey.Close() testKey.Close()
@@ -56,7 +56,7 @@ func TestNRPTEntriesCleanupOnConfigChange(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
// Verify 3 NRPT rules exist // Verify 3 NRPT rules exist
assert.Equal(t, 3, cfg.nrptEntryCount, "Should create 3 NRPT rules for 125 domains") assert.Equal(t, 3, countNRPTRuleKeys(t), "Should create 3 NRPT rules for 125 domains")
for i := 0; i < 3; i++ { for i := 0; i < 3; i++ {
exists, err := registryKeyExists(fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i)) exists, err := registryKeyExists(fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i))
require.NoError(t, err) require.NoError(t, err)
@@ -81,7 +81,7 @@ func TestNRPTEntriesCleanupOnConfigChange(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
// Verify first 2 NRPT rules exist // Verify first 2 NRPT rules exist
assert.Equal(t, 2, cfg.nrptEntryCount, "Should create 2 NRPT rules for 75 domains") assert.Equal(t, 2, countNRPTRuleKeys(t), "Should create 2 NRPT rules for 75 domains")
for i := 0; i < 2; i++ { for i := 0; i < 2; i++ {
exists, err := registryKeyExists(fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i)) exists, err := registryKeyExists(fmt.Sprintf("%s-%d", dnsPolicyConfigMatchPath, i))
require.NoError(t, err) require.NoError(t, err)
@@ -106,9 +106,65 @@ func registryKeyExists(path string) (bool, error) {
return true, nil return true, nil
} }
// TestNRPTCleanupWithoutRuleCount verifies that rules written by a previous run
// are removed by a configurator that has no record of how many there are: an
// unclean exit loses the in-memory count and a clean disconnect deletes the
// persisted one, so cleanup cannot depend on either.
func TestNRPTCleanupWithoutRuleCount(t *testing.T) {
if testing.Short() {
t.Skip("skipping registry integration test in short mode")
}
defer cleanupRegistryKeys(t)
cleanupRegistryKeys(t)
testIP := netip.MustParseAddr("100.64.0.1")
// 75 domains produce two indexed rules, as the current layout does
domains := make([]string, 75)
for i := range domains {
domains[i] = fmt.Sprintf(".domain%d.com", i+1)
}
previousRun := &registryConfigurator{}
require.NoError(t, previousRun.addDNSMatchPolicy(domains, testIP))
// the unsuffixed key an older version would have written
require.NoError(t, previousRun.configureDNSPolicy(dnsPolicyConfigMatchPath, []string{".legacy.example.com"}, testIP))
// a policy owned by someone else, which cleanup must not touch
foreignPath := DNSPolicyConfigRoot + `\DnsPolicyConfigTestForeign`
foreignKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, foreignPath, registry.SET_VALUE)
require.NoError(t, err, "Should create foreign policy key")
foreignKey.Close()
defer func() {
_ = registry.DeleteKey(registry.LOCAL_MACHINE, foreignPath)
}()
require.Equal(t, 3, countNRPTRuleKeys(t), "Should have two indexed rules and the legacy one")
// a configurator that never applied a DNS config, as one built after a
// restart or from a shutdown state without a count is
freshRun := &registryConfigurator{}
require.NoError(t, freshRun.removeDNSMatchPolicies())
assert.Equal(t, 0, countNRPTRuleKeys(t), "Should remove every rule left by the previous run")
exists, err := registryKeyExists(foreignPath)
require.NoError(t, err)
assert.True(t, exists, "Should not remove a policy that is not ours")
}
func countNRPTRuleKeys(t *testing.T) int {
t.Helper()
names, err := listNRPTRuleKeys(DNSPolicyConfigRoot)
require.NoError(t, err, "Should list NRPT rule keys")
return len(names)
}
func cleanupRegistryKeys(*testing.T) { func cleanupRegistryKeys(*testing.T) {
// Clean up more entries to account for batching tests with many domains cfg := &registryConfigurator{}
cfg := &registryConfigurator{nrptEntryCount: 20}
_ = cfg.removeDNSMatchPolicies() _ = cfg.removeDNSMatchPolicies()
} }
@@ -125,7 +181,7 @@ func TestNRPTDomainBatching(t *testing.T) {
// Create a test interface registry key so updateSearchDomains doesn't fail // Create a test interface registry key so updateSearchDomains doesn't fail
testGUID := "{12345678-1234-1234-1234-123456789ABC}" testGUID := "{12345678-1234-1234-1234-123456789ABC}"
interfacePath := `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces\` + testGUID interfacePath := InterfaceConfigPath + `\` + testGUID
testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE) testKey, _, err := registry.CreateKey(registry.LOCAL_MACHINE, interfacePath, registry.SET_VALUE)
require.NoError(t, err, "Should create test interface registry key") require.NoError(t, err, "Should create test interface registry key")
testKey.Close() testKey.Close()
@@ -193,7 +249,7 @@ func TestNRPTDomainBatching(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
// Verify that exactly expectedRuleCount rules were created // Verify that exactly expectedRuleCount rules were created
assert.Equal(t, tc.expectedRuleCount, cfg.nrptEntryCount, assert.Equal(t, tc.expectedRuleCount, countNRPTRuleKeys(t),
"Should create %d NRPT rules for %d domains", tc.expectedRuleCount, tc.domainCount) "Should create %d NRPT rules for %d domains", tc.expectedRuleCount, tc.domainCount)
// Verify all expected rules exist // Verify all expected rules exist
@@ -224,6 +224,7 @@ func TestResolver_StaleTriggersAsyncRefresh(t *testing.T) {
} }
func TestResolver_ConcurrentStaleHitsCollapseRefresh(t *testing.T) { func TestResolver_ConcurrentStaleHitsCollapseRefresh(t *testing.T) {
semaphore := make(chan struct{})
r := NewResolver() r := NewResolver()
chain := newFakeChain() chain := newFakeChain()
chain.setAnswer("mgmt.example.com.", dns.TypeA, "10.0.0.2") chain.setAnswer("mgmt.example.com.", dns.TypeA, "10.0.0.2")
@@ -239,7 +240,7 @@ func TestResolver_ConcurrentStaleHitsCollapseRefresh(t *testing.T) {
break break
} }
} }
time.Sleep(50 * time.Millisecond) // hold inflight long enough to collide <-semaphore // block the call to force request collision
} }
r.SetChainResolver(chain, 50) r.SetChainResolver(chain, 50)
@@ -255,17 +256,17 @@ func TestResolver_ConcurrentStaleHitsCollapseRefresh(t *testing.T) {
var wg sync.WaitGroup var wg sync.WaitGroup
for i := 0; i < 50; i++ { for i := 0; i < 50; i++ {
wg.Add(1) wg.Go(func() {
go func() {
defer wg.Done()
queryA(t, r, "mgmt.example.com.") queryA(t, r, "mgmt.example.com.")
}() })
} }
assert.Eventually(t, func() bool { return inflight.Load() >= 1 }, 2*time.Second, 100*time.Millisecond)
close(semaphore)
wg.Wait() wg.Wait()
waitFor(t, 2*time.Second, func() bool { assert.Eventually(t, func() bool { return inflight.Load() == 0 }, 2*time.Second, 100*time.Millisecond)
return inflight.Load() == 0
})
calls := chain.callCount("mgmt.example.com.", dns.TypeA) calls := chain.callCount("mgmt.example.com.", dns.TypeA)
assert.LessOrEqual(t, calls, 2, "singleflight must collapse concurrent refreshes (got %d)", calls) assert.LessOrEqual(t, calls, 2, "singleflight must collapse concurrent refreshes (got %d)", calls)
+1 -1
View File
@@ -4,7 +4,7 @@ import (
"net" "net"
"testing" "testing"
"github.com/golang/mock/gomock" "go.uber.org/mock/gomock"
"github.com/google/gopacket" "github.com/google/gopacket"
"github.com/google/gopacket/layers" "github.com/google/gopacket/layers"
"github.com/miekg/dns" "github.com/miekg/dns"
@@ -9,7 +9,7 @@ import (
"os" "os"
"testing" "testing"
"github.com/golang/mock/gomock" "go.uber.org/mock/gomock"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes" "golang.zx2c4.com/wireguard/wgctrl/wgtypes"
@@ -5,9 +5,8 @@ import (
) )
type ShutdownState struct { type ShutdownState struct {
Guid string Guid string
GPO bool GPO bool
NRPTEntryCount int
} }
func (s *ShutdownState) Name() string { func (s *ShutdownState) Name() string {
@@ -16,9 +15,8 @@ func (s *ShutdownState) Name() string {
func (s *ShutdownState) Cleanup() error { func (s *ShutdownState) Cleanup() error {
manager := &registryConfigurator{ manager := &registryConfigurator{
guid: s.Guid, guid: s.Guid,
gpo: s.GPO, gpo: s.GPO,
nrptEntryCount: s.NRPTEntryCount,
} }
if err := manager.restoreUncleanShutdownDNS(); err != nil { if err := manager.restoreUncleanShutdownDNS(); err != nil {
+1 -1
View File
@@ -101,7 +101,7 @@ func (m *Manager) Start(fwdEntries []*ForwarderEntry) error {
m.dnsForwarder = NewDNSForwarder(listenAddress, dnsTTL, m.firewall, m.statusRecorder, m.wgIface) m.dnsForwarder = NewDNSForwarder(listenAddress, dnsTTL, m.firewall, m.statusRecorder, m.wgIface)
go func() { go func() {
if err := m.dnsForwarder.Listen(fwdEntries); err != nil { if err := m.dnsForwarder.Listen(fwdEntries); err != nil { //nolint:staticcheck
// todo handle close error if it is exists // todo handle close error if it is exists
log.Errorf("failed to start DNS forwarder, err: %v", err) log.Errorf("failed to start DNS forwarder, err: %v", err)
} }
+74
View File
@@ -0,0 +1,74 @@
// Package elevate re-runs this very executable under the operating system's own
// privilege-elevation mechanism and waits for it to finish.
//
// It exists so that a change the daemon restricts to root/administrator can be
// authorized from the GUI, by the user, at the moment they ask for it: Windows
// shows the UAC consent dialog, macOS the system authentication dialog, and
// Linux/FreeBSD the session's polkit agent. The credentials, where any are
// asked for, are collected by the operating system and never pass through
// NetBird.
//
// What the elevated process then does is the caller's business: it is the same
// binary, in a one-shot mode, and it is authorized by the daemon exactly like
// any other privileged caller, from the identity the kernel reports on the
// control channel. Nothing here grants privilege, and the daemon gains no new
// way to be talked into something: elevation only changes who is calling it.
package elevate
import (
"context"
"errors"
log "github.com/sirupsen/logrus"
)
// AppliedMarker is what the elevated process prints on standard output once it has
// done what it was run for.
//
// macOS's AuthorizationExecuteWithPrivileges reports no exit status and does not
// say which process it started, so there this line is the only evidence that the
// change was applied. The other platforms have an exit code and ignore it.
const AppliedMarker = "netbird-elevated: applied"
var (
// ErrDeclined reports that the user dismissed the prompt or did not
// authenticate. Nothing happened and nothing is wrong: a caller undoes its
// optimistic update and stays quiet.
ErrDeclined = errors.New("authorization declined")
// ErrUnavailable reports that this host has no elevation mechanism we can
// drive: no polkit on a Unix desktop, or an executable we decline to run as
// root. A caller falls back to telling the user which command to run.
ErrUnavailable = errors.New("no privilege elevation mechanism available")
)
// Run runs this executable with args under the platform's elevation mechanism
// and waits for it to exit. A non-zero exit is returned as an error, so the
// caller can treat a completed Run as the operation having succeeded.
//
// The args are the caller's own command line, so they cross no privilege
// boundary: only a user who has just authenticated as an administrator can get
// them run at all.
func Run(ctx context.Context, args ...string) error {
self, err := trustedSelf()
if err != nil {
return err
}
return run(ctx, self, args)
}
// Available reports whether Run has a mechanism to use on this host, so a caller
// can offer the prompt only when there is one and otherwise fall back to
// guidance the user can act on. It answers from what is installed, not from what
// the user is allowed to do: an administrator's password may still be required
// and may still not be given, which is ErrDeclined from Run.
func Available() bool {
if _, err := trustedSelf(); err != nil {
// Worth a line: this is also what a build run from a group-writable
// directory hits, and there is nothing in the UI to say why the offer is
// missing.
log.Debugf("not offering privilege elevation: %v", err)
return false
}
return mechanismAvailable()
}
+18
View File
@@ -0,0 +1,18 @@
package elevate
import "strings"
// noOutput stands in for a process that said nothing, so that a report of what it
// said still reads as a sentence.
const noOutput = "no output"
func firstLine(s string) string {
s = strings.TrimSpace(s)
if s == "" {
return noOutput
}
if i := strings.IndexByte(s, '\n'); i >= 0 {
return s[:i]
}
return s
}
+21
View File
@@ -0,0 +1,21 @@
package elevate
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestFirstLine(t *testing.T) {
tests := []struct{ in, want string }{
{in: "", want: noOutput},
{in: " \n ", want: noOutput},
{in: "one line", want: "one line"},
{in: "first\nsecond", want: "first"},
{in: "\nsecond\n", want: "second"},
}
for _, tt := range tests {
assert.Equal(t, tt.want, firstLine(tt.in), "input %q", tt.in)
}
}
+359
View File
@@ -0,0 +1,359 @@
package elevate
import (
"context"
"errors"
"fmt"
"os"
"runtime"
"strings"
"sync"
"syscall"
"unsafe"
"github.com/ebitengine/purego"
log "github.com/sirupsen/logrus"
)
// Authorization Services, reached through purego rather than cgo so the released
// binaries keep building with CGO_ENABLED=0.
//
// The prompt belongs to this process, which is what makes it carry the
// application's name and our own explanation. Going through osascript instead puts
// the very same trampoline behind a dialog attributed to osascript, and means
// handing a shell a command line to re-parse.
//
// # On AuthorizationExecuteWithPrivileges
//
// It is deprecated, and Apple's guidance (Quinn, "BSD Privilege Escalation on
// macOS", developer.apple.com/forums/thread/708765) is "while it still works, it's
// been deprecated for many years. Do not use it in a widely distributed product."
// It is used here anyway, knowingly, because the alternatives Apple offers are for
// *obtaining* ongoing privileges — an installer package, SMAppService, SMJobBless —
// and NetBird already has what they would install: a launchd daemon running as
// root. What is missing is only a way for an unprivileged client to ask it to act.
//
// The way to that without a deprecated call is to authorize the client instead of
// elevating one: the app takes the right with AuthorizationCreate, passes the
// AuthorizationExternalForm to the daemon, and the daemon checks it with
// AuthorizationCopyRights before acting — none of which is deprecated. It is the
// better design and it is where this should end up. It also means the daemon
// accepting an authorization over its control socket, which is a new way to be
// asked for privileged work and wants reviewing as such, so it is deliberately not
// bundled in with the rest of this.
//
// Until then, three things keep the deprecation from being a trap. Every symbol is
// resolved with an error rather than a panic, so a macOS that has dropped this
// function leaves the app offering the user a command instead of crashing on the
// way to a prompt. A failure to run the tool is reported as ErrUnavailable, so the
// fallback is the same one an agent-less Linux session gets. And the whole path
// runs under guard, which turns a panic out of the FFI layer into that same
// fallback.
//
// The trampoline passes on the environment it was given, so what it starts as root
// must be an executable this user's peers cannot influence: that is what
// trustedSelf refuses, and what signing the binary settles for the loader.
const (
securityFramework = "/System/Library/Frameworks/Security.framework/Security"
libSystem = "/usr/lib/libSystem.B.dylib"
// trampoline is what the framework hands the tool to. Present on every macOS,
// and worth confirming before offering a prompt rather than mid-prompt.
trampoline = "/usr/libexec/security_authtrampoline"
)
// rightExecute is the right an administrator holds, and what
// AuthorizationExecuteWithPrivileges requires of us.
const rightExecute = "system.privilege.admin"
// promptKey is kAuthorizationEnvironmentPrompt, which puts a sentence of ours above
// the system's in the dialog. It is about the change rather than the mechanism.
const (
promptKey = "prompt"
promptText = "NetBird needs to change a setting that grants SSH access to this computer."
)
// OSStatus values from SecBase.h that mean something to us; anything else is
// reported as it comes.
const (
errAuthorizationSuccess = 0
errAuthorizationDenied = -60005
errAuthorizationCanceled = -60006
errAuthorizationInteractionNotAllowed = -60007
errAuthorizationToolExecuteFailure = -60031
errAuthorizationToolEnvironmentError = -60032
)
// AuthorizationFlags from Authorization.h.
const (
flagDefaults = 0
flagInteractionAllowed = 1 << 0
flagExtendRights = 1 << 1
flagDestroyRights = 1 << 3
flagPreAuthorize = 1 << 4
)
// authorizationItem mirrors AuthorizationItem: a name, and a value the name gives
// meaning to. 32 bytes on both amd64 and arm64.
type authorizationItem struct {
name *byte
valueLength uintptr
value unsafe.Pointer
// flags is reserved by the API and always zero. Declared because the layout
// is the contract: without it the struct is 24 bytes where C reads 32.
flags uint32 //nolint:unused // part of the C layout
}
// authorizationItemSet mirrors AuthorizationItemSet, which serves as both an
// AuthorizationRights and an AuthorizationEnvironment.
type authorizationItemSet struct {
count uint32
items *authorizationItem
}
var (
authorizationCreate func(rights, environment *authorizationItemSet, flags uint32, authorization *uintptr) int32
authorizationExecuteWithPrivileges func(authorization uintptr, pathToTool string, options uint32, arguments *uintptr, communicationsPipe *uintptr) int32
authorizationFree func(authorization uintptr, flags uint32) int32
fileno func(stream uintptr) int32
fclose func(stream uintptr) int32
loadOnce sync.Once
loadErr error
)
// load resolves the functions once. A framework that cannot be opened, or a symbol
// that is no longer there, leaves the host without a mechanism rather than taking
// the process down with it: see the note on deprecation above.
func load() error {
loadOnce.Do(func() { loadErr = guard("loading Security.framework", resolve) })
return loadErr
}
// guard turns a panic out of the FFI layer into an error, so an API that has
// changed under us costs the user a prompt rather than the window they were
// clicking in. purego panics on a signature it cannot map, and this is the one
// place in the client that calls a deprecated system function.
//
// It catches Go panics, which is what purego raises. A fault inside the framework
// itself is not a panic and not recoverable; the layout the tests pin down is what
// stands between us and that.
func guard(what string, fn func() error) (err error) {
defer func() {
r := recover()
if r == nil {
return
}
log.Errorf("%s panicked: %v", what, r)
err = fmt.Errorf("%w: %s: %v", ErrUnavailable, what, r)
}()
return fn()
}
func resolve() error {
security, err := purego.Dlopen(securityFramework, purego.RTLD_LAZY|purego.RTLD_GLOBAL)
if err != nil {
return fmt.Errorf("open %s: %w", securityFramework, err)
}
system, err := purego.Dlopen(libSystem, purego.RTLD_LAZY|purego.RTLD_GLOBAL)
if err != nil {
return fmt.Errorf("open %s: %w", libSystem, err)
}
// purego.RegisterLibFunc panics on a symbol it cannot find, which is not how a
// deprecated function's disappearance should reach the user.
for _, fn := range []struct {
ptr any
handle uintptr
name string
}{
{&authorizationCreate, security, "AuthorizationCreate"},
{&authorizationExecuteWithPrivileges, security, "AuthorizationExecuteWithPrivileges"},
{&authorizationFree, security, "AuthorizationFree"},
{&fileno, system, "fileno"},
{&fclose, system, "fclose"},
} {
symbol, err := purego.Dlsym(fn.handle, fn.name)
if err != nil {
return fmt.Errorf("resolve %s: %w", fn.name, err)
}
if symbol == 0 {
return fmt.Errorf("resolve %s: not present on this system", fn.name)
}
purego.RegisterFunc(fn.ptr, symbol)
}
return nil
}
// run asks the system to run self as root: first for the right, which is what puts
// up the authentication dialog and collects the password or takes the Touch ID,
// then for the tool. The credentials go to the system's authorization trampoline
// and never to us.
//
// The context bounds only our own waiting; the dialog belongs to the system and
// closes when the user answers it.
func run(ctx context.Context, self string, args []string) error {
if err := load(); err != nil {
return fmt.Errorf("%w: %v", ErrUnavailable, err)
}
return guard("asking for privileges", func() error {
authorization, err := authorize()
if err != nil {
return err
}
defer authorizationFree(authorization, flagDestroyRights)
return execute(ctx, authorization, self, args)
})
}
func mechanismAvailable() bool {
if err := load(); err != nil {
return false
}
info, err := os.Stat(trampoline)
return err == nil && !info.IsDir()
}
// authorize obtains the right, prompting for it. A dismissed dialog comes back as
// errAuthorizationCanceled and a password given up on as errAuthorizationDenied;
// both are the user's answer rather than a failure.
func authorize() (uintptr, error) {
var pinner runtime.Pinner
defer pinner.Unpin()
rights := itemSet(&pinner, authorizationItem{name: cString(&pinner, rightExecute)})
environment := itemSet(&pinner, promptItem(&pinner))
var authorization uintptr
status := authorizationCreate(rights, environment,
flagDefaults|flagInteractionAllowed|flagPreAuthorize|flagExtendRights, &authorization)
switch status {
case errAuthorizationSuccess:
return authorization, nil
case errAuthorizationCanceled, errAuthorizationDenied:
return 0, ErrDeclined
case errAuthorizationInteractionNotAllowed:
// Nowhere to put a dialog, so there is nobody to ask: a launch daemon, or
// a session with no window server.
return 0, fmt.Errorf("%w: this session cannot show an authorization prompt", ErrUnavailable)
default:
return 0, fmt.Errorf("request %s: OSStatus %d", rightExecute, status)
}
}
// execute runs the tool with the right in hand and waits for it by reading the pipe
// it is given until the tool closes it.
//
// AuthorizationExecuteWithPrivileges reports no exit status and does not say what
// process it started, which is why the one-shot says so itself: what it prints is
// the only evidence that the change was applied.
func execute(ctx context.Context, authorization uintptr, self string, args []string) error {
var pinner runtime.Pinner
defer pinner.Unpin()
argv := make([]uintptr, 0, len(args)+1)
for _, arg := range args {
argv = append(argv, uintptr(unsafe.Pointer(cString(&pinner, arg))))
}
argv = append(argv, 0)
pinner.Pin(&argv[0])
var pipe uintptr
status := authorizationExecuteWithPrivileges(authorization, self, flagDefaults, &argv[0], &pipe)
switch status {
case errAuthorizationSuccess:
case errAuthorizationCanceled:
return ErrDeclined
case errAuthorizationToolExecuteFailure, errAuthorizationToolEnvironmentError:
// The right was granted and the tool still did not start. Nothing the user
// can do about it from here, so point them at the command instead.
return fmt.Errorf("%w: the system would not run %s elevated (OSStatus %d)", ErrUnavailable, self, status)
default:
return fmt.Errorf("run %s elevated: OSStatus %d", self, status)
}
out, err := readPipe(ctx, pipe)
if err != nil {
return err
}
return checkApplied(out)
}
// checkApplied reads the one-shot's report, which stands in for the exit status
// there is no way to ask for here. A run that said nothing did not apply the
// change, whatever else went on.
func checkApplied(out string) error {
if !strings.Contains(out, AppliedMarker) {
return fmt.Errorf("elevated netbird did not report the change as applied: %s", firstLine(out))
}
return nil
}
// readPipe drains the tool's output, which ends when the tool exits and is
// therefore also how we wait for it.
func readPipe(ctx context.Context, pipe uintptr) (string, error) {
if pipe == 0 {
return "", nil
}
defer fclose(pipe)
fd := int(fileno(pipe))
if fd < 0 {
return "", nil
}
var out strings.Builder
buf := make([]byte, 4096)
for {
if err := ctx.Err(); err != nil {
return out.String(), err
}
n, err := syscall.Read(fd, buf)
if n > 0 {
out.Write(buf[:n])
}
switch {
case errors.Is(err, syscall.EINTR):
// A signal landed mid-read, which says nothing about the tool.
continue
case err != nil:
log.Debugf("read the elevated process's output: %v", err)
return out.String(), nil
case n <= 0:
// End of file: the tool closed the pipe, which is how it exiting
// reaches us.
return out.String(), nil
}
}
}
// itemSet builds an AuthorizationItemSet over items, pinned for the call.
func itemSet(pinner *runtime.Pinner, items ...authorizationItem) *authorizationItemSet {
pinner.Pin(&items[0])
set := &authorizationItemSet{count: uint32(len(items)), items: &items[0]}
pinner.Pin(set)
return set
}
// promptItem is the environment entry carrying our sentence for the dialog.
func promptItem(pinner *runtime.Pinner) authorizationItem {
value := []byte(promptText)
pinner.Pin(&value[0])
return authorizationItem{
name: cString(pinner, promptKey),
valueLength: uintptr(len(value)),
value: unsafe.Pointer(&value[0]),
}
}
// cString returns a NUL-terminated copy of s, pinned so the C side may hold it for
// the duration of the call.
func cString(pinner *runtime.Pinner, s string) *byte {
b := append([]byte(s), 0)
pinner.Pin(&b[0])
return &b[0]
}
+111
View File
@@ -0,0 +1,111 @@
package elevate
import (
"errors"
"runtime"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// The framework has to load and the symbols have to resolve, or nothing else here
// means anything.
func TestSecurityFrameworkLoads(t *testing.T) {
require.NoError(t, load(), "Security.framework must open")
for name, fn := range map[string]any{
"AuthorizationCreate": authorizationCreate,
"AuthorizationExecuteWithPrivileges": authorizationExecuteWithPrivileges,
"AuthorizationFree": authorizationFree,
"fileno": fileno,
"fclose": fclose,
} {
assert.NotNil(t, fn, "%s must resolve", name)
}
}
// A request with no interaction allowed exercises the whole call — the rights and
// environment structs, and the OSStatus that comes back — without a dialog anybody
// has to answer. What the system decides is its business; that it decides at all is
// what this asserts.
func TestAuthorizationCreateWithoutInteraction(t *testing.T) {
if err := load(); err != nil {
t.Skipf("Security.framework did not open: %v", err)
}
var pinner runtime.Pinner
defer pinner.Unpin()
rights := itemSet(&pinner, authorizationItem{name: cString(&pinner, rightExecute)})
environment := itemSet(&pinner, promptItem(&pinner))
require.EqualValues(t, 1, rights.count, "the rights struct layout must match the C one")
var authorization uintptr
status := authorizationCreate(rights, environment, flagDefaults|flagExtendRights, &authorization)
switch status {
case errAuthorizationSuccess:
// Credentials were already cached for this session.
authorizationFree(authorization, flagDestroyRights)
case errAuthorizationDenied, errAuthorizationInteractionNotAllowed:
// The expected answers when nobody may be asked.
default:
require.Failf(t, "unknown OSStatus", "AuthorizationCreate returned %d, want a status we recognise", status)
}
}
// Asking with a right nobody has must not be mistaken for a declined prompt: the
// caller would report nothing at all.
func TestAuthorizeUnknownRightIsNotDeclined(t *testing.T) {
if err := load(); err != nil {
t.Skipf("Security.framework did not open: %v", err)
}
var pinner runtime.Pinner
defer pinner.Unpin()
rights := itemSet(&pinner, authorizationItem{name: cString(&pinner, "io.netbird.right.that.does.not.exist")})
var authorization uintptr
status := authorizationCreate(rights, nil, flagDefaults|flagExtendRights, &authorization)
if status == errAuthorizationSuccess {
authorizationFree(authorization, flagDestroyRights)
}
assert.NotEqual(t, int32(errAuthorizationSuccess), status, "a right that does not exist must not be granted")
}
func TestMechanismAvailable(t *testing.T) {
assert.True(t, mechanismAvailable(), "the trampoline exists on every macOS")
}
// The one-shot's report is what stands in for an exit status here, so a run that
// says nothing must not read as success.
func TestCheckApplied(t *testing.T) {
require.NoError(t, checkApplied(AppliedMarker+"\n"), "the report the one-shot prints")
require.NoError(t, checkApplied("some warning\n"+AppliedMarker+"\n"), "the report after other output")
assert.Error(t, checkApplied(""), "a run that printed nothing did not apply the change")
assert.Error(t, checkApplied("dyld: library not loaded\n"), "output that is not the report")
}
// A panic out of the FFI layer has to reach the caller as "no mechanism", which is
// the outcome that offers the user the command instead of taking the window down.
func TestGuardTurnsAPanicIntoUnavailable(t *testing.T) {
err := guard("pretending to call something", func() error {
panic("purego: signature it cannot map")
})
require.ErrorIs(t, err, ErrUnavailable, "a panic must read as a missing mechanism")
assert.Contains(t, err.Error(), "pretending to call something", "what panicked")
}
// guard wraps every darwin path, so what a caller switches on has to survive it.
func TestGuardPassesErrorsThrough(t *testing.T) {
sentinel := errors.New("the call itself failed")
assert.ErrorIs(t, guard("calling", func() error { return sentinel }), sentinel,
"the error it was given")
assert.ErrorIs(t, guard("calling", func() error { return ErrDeclined }), ErrDeclined,
"a declined prompt stays declined")
assert.NoError(t, guard("calling", func() error { return nil }), "a call that worked")
}
+117
View File
@@ -0,0 +1,117 @@
//go:build linux
package elevate
import (
"context"
"errors"
"fmt"
"io"
"os"
"os/exec"
"strings"
)
// pkexec exit codes that are about the authorization rather than about the program
// we asked it to run. The manual page reserves both.
const (
// exitDismissed is returned when the user dismissed the authentication
// dialog.
exitDismissed = 126
// exitNotAuthorized is returned when the authorization was not obtained. That
// covers the user saying no as well as pkexec having had nobody to ask: see
// noAgentMarkers.
exitNotAuthorized = 127
)
// exitNotAuthorized covers three different endings that only pkexec's own words
// tell apart, so they are matched here. Read with LC_ALL=C so the words are the
// ones written below.
//
// refusedMarker is a refusal: the user said no, gave up on the password, or holds
// an account that may not elevate at all.
const refusedMarker = "Not authorized"
// noAgentMarkers say pkexec had no way to ask: no agent registered for the
// session, and no controlling terminal for the textual agent it falls back to.
var noAgentMarkers = []string{"authentication agent", "controlling terminal"}
// run asks polkit to run self as root. pkexec hands the request to the session's
// polkit agent, which is what prompts and what collects any password; we see only
// its verdict.
//
// The environment is otherwise deliberately not passed through: pkexec clears it
// bar a small allowlist, and the one-shot needs nothing from it.
func run(ctx context.Context, self string, args []string) error {
pkexec, err := exec.LookPath("pkexec")
if err != nil {
return fmt.Errorf("%w: pkexec is not installed", ErrUnavailable)
}
cmd := exec.CommandContext(ctx, pkexec, append([]string{self}, args...)...)
// C locale so pkexec's own diagnostics are the ones noAgentMarkers knows.
cmd.Env = append(os.Environ(), "LC_ALL=C")
var stderr strings.Builder
cmd.Stderr = &stderr
// The one-shot reports itself on stdout for macOS's sake, where there is no
// exit status to read. Here there is one, so that line is noise.
cmd.Stdout = io.Discard
err = cmd.Run()
if err == nil {
return nil
}
var exitErr *exec.ExitError
if !errors.As(err, &exitErr) {
return fmt.Errorf("run pkexec: %w", err)
}
// Matched against everything pkexec said, reported as one line: a complaint
// that is not the first thing printed still has to be recognised, and reading
// it as a refusal would swallow it.
full := stderr.String()
out := firstLine(full)
switch exitErr.ExitCode() {
case exitDismissed:
return ErrDeclined
case exitNotAuthorized:
return notAuthorized(full, out)
default:
return fmt.Errorf("elevated netbird exited with %d: %s", exitErr.ExitCode(), out)
}
}
// notAuthorized sorts out the three endings pkexec reports as exitNotAuthorized.
//
// It also returns that code when the authorization succeeded and it then could
// not run the program, so a refusal has to be recognised rather than assumed:
// reading every one of these as "the user said no" would revert the control in
// silence on a host where elevation is broken.
func notAuthorized(full, out string) error {
switch {
case hasAny(full, noAgentMarkers):
return fmt.Errorf("%w: polkit had no way to ask: %s", ErrUnavailable, out)
case out == noOutput, strings.Contains(full, refusedMarker):
// The user said no, which needs no message; that an account barred from
// elevating altogether lands here too is why the reason is kept.
return fmt.Errorf("%w: %s", ErrDeclined, out)
default:
return fmt.Errorf("pkexec could not run elevated netbird: %s", out)
}
}
func hasAny(s string, markers []string) bool {
for _, marker := range markers {
if strings.Contains(s, marker) {
return true
}
}
return false
}
func mechanismAvailable() bool {
_, err := exec.LookPath("pkexec")
return err == nil
}
+110
View File
@@ -0,0 +1,110 @@
//go:build linux
package elevate
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// fakePkexec puts a pkexec on PATH that exits with the given code, so the
// mapping from polkit's exit codes onto our errors can be exercised without a
// polkit agent.
func fakePkexec(t *testing.T, exitCode int, stderr string) {
t.Helper()
dir := t.TempDir()
script := fmt.Sprintf("#!/bin/sh\necho %s >&2\nexit %d\n", shellQuote(stderr), exitCode)
require.NoError(t, os.WriteFile(filepath.Join(dir, "pkexec"), []byte(script), 0o700), "write the fake pkexec")
t.Setenv("PATH", dir)
}
func shellQuote(s string) string {
return "'" + strings.ReplaceAll(s, "'", `'\''`) + "'"
}
func TestRunMapsPkexecExitCodes(t *testing.T) {
tests := []struct {
name string
exitCode int
stderr string
wantErr error
}{
{name: "applied", exitCode: 0},
{
name: "dialog dismissed",
exitCode: exitDismissed,
stderr: "Error executing command as another user: Request dismissed",
wantErr: ErrDeclined,
},
{
// What a graphical agent reports for a cancelled prompt. Not a
// failure: the user was asked and answered.
name: "prompt cancelled",
exitCode: exitNotAuthorized,
stderr: "Error executing command as another user: Not authorized",
wantErr: ErrDeclined,
},
{
// The same status, but pkexec never got to ask anybody.
name: "no agent and no terminal to fall back on",
exitCode: exitNotAuthorized,
stderr: "Error creating textual authentication agent: Error opening current controlling terminal for the process (`/dev/tty'): No such device or address",
wantErr: ErrUnavailable,
},
{
// And the same status again once the authorization succeeded and
// pkexec could not run what it had been authorized to run. Reading
// that as a refusal would revert the control in silence on a host
// where elevation is broken.
name: "authorized but not runnable",
exitCode: exitNotAuthorized,
stderr: "Error executing command as another user: No such file or directory",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
fakePkexec(t, tt.exitCode, tt.stderr)
err := run(context.Background(), "/nonexistent/netbird-ui", []string{"--flag"})
switch {
case tt.wantErr != nil:
require.ErrorIs(t, err, tt.wantErr, "exit %d said %q", tt.exitCode, tt.stderr)
case tt.exitCode == 0:
require.NoError(t, err, "a pkexec that exited cleanly applied the change")
default:
require.Error(t, err, "exit %d said %q", tt.exitCode, tt.stderr)
assert.NotErrorIs(t, err, ErrDeclined, "not the user's answer")
assert.NotErrorIs(t, err, ErrUnavailable, "not a missing mechanism")
}
})
}
}
// An exit code that is not polkit's is the one-shot's own failure, and has to
// stay distinguishable from a declined prompt: the caller reports it.
func TestRunReportsOneShotFailure(t *testing.T) {
fakePkexec(t, 3, "the one-shot said no")
err := run(context.Background(), "/nonexistent/netbird-ui", nil)
require.Error(t, err, "a one-shot that failed is not a prompt that was answered")
assert.NotErrorIs(t, err, ErrDeclined, "not the user's answer")
assert.NotErrorIs(t, err, ErrUnavailable, "not a missing mechanism")
}
func TestRunWithoutPkexecIsUnavailable(t *testing.T) {
t.Setenv("PATH", t.TempDir())
err := run(context.Background(), "/nonexistent/netbird-ui", nil)
require.ErrorIs(t, err, ErrUnavailable, "no pkexec means no mechanism")
assert.False(t, mechanismAvailable(), "mechanismAvailable without pkexec on PATH")
}
@@ -0,0 +1,19 @@
//go:build !windows && !darwin && !linux
package elevate
import "context"
// run reports that this platform has no elevation prompt to drive.
//
// The desktop app is the only caller and is not built for any of these: mobile
// and WASM have no local user to ask, and the FreeBSD client ships without a UI.
// pkexec would be the mechanism there, and run_unix.go is what to widen if that
// changes.
func run(context.Context, string, []string) error {
return ErrUnavailable
}
func mechanismAvailable() bool {
return false
}
+187
View File
@@ -0,0 +1,187 @@
package elevate
import (
"context"
"errors"
"fmt"
"runtime"
"unsafe"
log "github.com/sirupsen/logrus"
"golang.org/x/sys/windows"
)
const (
// seeMaskNoCloseProcess keeps the started process's handle open in
// hProcess so we can wait for it.
seeMaskNoCloseProcess = 0x00000040
// seeMaskNoAsync makes ShellExecuteExW finish its work before returning,
// which it must when the calling thread does not pump messages.
seeMaskNoAsync = 0x00000100
// seeMaskFlagNoUI suppresses the shell's own error dialogs; the UAC consent
// dialog is not one of them and still appears.
seeMaskFlagNoUI = 0x00000400
// swHide: the one-shot has no window to show.
swHide = 0
)
// shellExecuteInfoW mirrors SHELLEXECUTEINFOW. The field order and Go's own
// padding match the C layout on both 386 and amd64.
type shellExecuteInfoW struct {
cbSize uint32
fMask uint32
hwnd windows.HWND
lpVerb *uint16
lpFile *uint16
lpParameters *uint16
lpDirectory *uint16
nShow int32
hInstApp windows.Handle
lpIDList uintptr
lpClass *uint16
hkeyClass windows.Handle
dwHotKey uint32
hIconOrMonitor windows.Handle
hProcess windows.Handle
}
var (
shell32 = windows.NewLazySystemDLL("shell32.dll")
procShellExecuteEx = shell32.NewProc("ShellExecuteExW")
)
// run starts self elevated with the "runas" verb, which is what raises the UAC
// consent dialog, and waits for it to finish. Windows decides whether consent is
// enough or an administrator's credentials are needed, and collects them itself.
func run(ctx context.Context, self string, args []string) error {
verb, err := windows.UTF16PtrFromString("runas")
if err != nil {
return fmt.Errorf("encode verb: %w", err)
}
file, err := windows.UTF16PtrFromString(self)
if err != nil {
return fmt.Errorf("encode %s: %w", self, err)
}
params, err := windows.UTF16PtrFromString(windows.ComposeCommandLine(args))
if err != nil {
return fmt.Errorf("encode arguments: %w", err)
}
info := shellExecuteInfoW{
fMask: seeMaskNoCloseProcess | seeMaskNoAsync | seeMaskFlagNoUI,
hwnd: ownerWindow(),
lpVerb: verb,
lpFile: file,
lpParameters: params,
nShow: swHide,
}
info.cbSize = uint32(unsafe.Sizeof(info))
process, err := shellExecute(&info)
if err != nil {
return err
}
defer func() {
if err := windows.CloseHandle(process); err != nil {
log.Debugf("close elevated process handle: %v", err)
}
}()
return waitForProcess(ctx, process)
}
// shellExecute performs the call itself. ShellExecuteExW wants COM initialised on
// the calling thread, so the goroutine is pinned to one for the duration and COM
// is set up on it; an "already initialised, different mode" answer is fine,
// because then somebody else has done it for us.
func shellExecute(info *shellExecuteInfoW) (windows.Handle, error) {
runtime.LockOSThread()
defer runtime.UnlockOSThread()
switch err := windows.CoInitializeEx(0, windows.COINIT_APARTMENTTHREADED); {
case err == nil, isHResult(err, windows.S_FALSE):
// Ours, or already initialised in the same mode: either way this call
// counts and has to be balanced.
defer windows.CoUninitialize()
case isHResult(err, windows.RPC_E_CHANGED_MODE):
// The thread is already in the other apartment model. ShellExecuteExW
// works there too, and there is nothing of ours to balance.
default:
return 0, fmt.Errorf("initialise COM: %w", err)
}
ret, _, lastErr := procShellExecuteEx.Call(uintptr(unsafe.Pointer(info)))
if ret != 0 {
return info.hProcess, nil
}
if errors.Is(lastErr, windows.ERROR_CANCELLED) {
return 0, ErrDeclined
}
return 0, fmt.Errorf("run elevated: %w", lastErr)
}
// ownerWindow returns this process's foreground window, and 0 when the window in
// front belongs to somebody else or cannot be attributed. ShellExecuteExW takes it
// as the parent for the UI it raises, which is what keeps the consent dialog in
// front of the window the user was just clicking in instead of behind it. It is
// also what a remote-desktop session needs to place the dialog at all when the
// secure desktop is switched off.
func ownerWindow() windows.HWND {
hwnd := windows.GetForegroundWindow()
if hwnd == 0 {
return 0
}
var pid uint32
if _, err := windows.GetWindowThreadProcessId(hwnd, &pid); err != nil {
log.Debugf("cannot attribute the foreground window, raising the prompt without an owner: %v", err)
return 0
}
if pid != windows.GetCurrentProcessId() {
return 0
}
return hwnd
}
// isHResult reports whether err carries the given HRESULT. CoInitializeEx
// returns its HRESULT as an Errno, so the comparison is on the raw value.
func isHResult(err error, hresult windows.Handle) bool {
var errno windows.Errno
return errors.As(err, &errno) && uintptr(errno) == uintptr(hresult)
}
func waitForProcess(ctx context.Context, process windows.Handle) error {
// The wait is interruptible so a cancelled context stops us waiting on a
// consent dialog nobody is going to answer. The elevated process is not
// ours to kill, and it either applies the change or does not.
for {
event, err := windows.WaitForSingleObject(process, 250)
if err != nil {
return fmt.Errorf("wait for the elevated process: %w", err)
}
if event == uint32(windows.WAIT_OBJECT_0) {
break
}
if err := ctx.Err(); err != nil {
return err
}
}
var code uint32
if err := windows.GetExitCodeProcess(process, &code); err != nil {
return fmt.Errorf("read the elevated process's exit code: %w", err)
}
if code != 0 {
return fmt.Errorf("elevated netbird exited with %d", code)
}
return nil
}
// mechanismAvailable is true on Windows: UAC prompts for consent when the user
// is an administrator and for an administrator's credentials when they are not,
// so there is always something to ask.
func mechanismAvailable() bool {
return true
}
+40
View File
@@ -0,0 +1,40 @@
package elevate
import (
"fmt"
"os"
"path/filepath"
)
// trustedSelf returns the path of this executable, provided it is one we are
// willing to have run as root.
//
// The check is what keeps elevation from becoming a way to launder someone
// else's code into a root process: the user consents to NetBird being elevated,
// having been shown NetBird's name, so what runs must be the file NetBird was
// installed as and not something a third party could have swapped for it. An
// executable only its owner can write is that; anything wider is refused, and
// the caller falls back to showing the command instead.
//
// The owner writing to their own executable is not part of that threat: code
// running as the user can already prompt them for anything, and could just as
// well ask them to run the command by hand. What matters is that no *other*
// unprivileged account can reach it.
func trustedSelf() (string, error) {
exe, err := os.Executable()
if err != nil {
return "", fmt.Errorf("locate this executable: %w", err)
}
// Resolve symlinks so the checks below apply to the file that would actually
// be executed, not to a link somebody else may control.
resolved, err := filepath.EvalSymlinks(exe)
if err != nil {
return "", fmt.Errorf("resolve %s: %w", exe, err)
}
if err := checkOnlyOwnerWritable(resolved); err != nil {
return "", fmt.Errorf("%w: %s cannot be trusted to run as root: %w", ErrUnavailable, resolved, err)
}
return resolved, nil
}
@@ -0,0 +1,10 @@
package elevate
// adminWriteGIDs are the groups whose write access to an executable does not
// widen who could authorize elevating it.
//
// macOS installs applications as root:admin, mode 0775, /Applications included,
// so requiring owner-only write would reject every normal install. Group admin
// (gid 80) is exactly the set of accounts that can answer the authentication
// dialog, so its write access grants nothing the prompt would not.
var adminWriteGIDs = []uint32{0, 80}
@@ -0,0 +1,9 @@
//go:build !windows && !darwin
package elevate
// adminWriteGIDs are the groups whose write access to an executable does not
// widen who could authorize elevating it. Only root's own group qualifies here:
// a distribution installs into root-owned directories, and there is no
// system-wide administrators group that both writes them and answers polkit.
var adminWriteGIDs = []uint32{0}
+119
View File
@@ -0,0 +1,119 @@
//go:build !windows
package elevate
import (
"errors"
"fmt"
"os"
"path/filepath"
"slices"
"strconv"
"syscall"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/getent"
)
// checkOnlyOwnerWritable reports an error unless path, and every directory leading
// to it, is owned by either root or this user and writable by nobody who could not
// already act as its owner. A writable directory is as good as a writable file,
// since anything in it can be replaced, so the whole chain is checked.
func checkOnlyOwnerWritable(path string) error {
self := uint32(os.Getuid())
for dir := path; ; dir = filepath.Dir(dir) {
info, err := os.Lstat(dir)
if err != nil {
return fmt.Errorf("stat %s: %w", dir, err)
}
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return errors.New("file ownership is unavailable on this platform")
}
if stat.Uid != 0 && stat.Uid != self {
return fmt.Errorf("%s is owned by uid %d, neither root nor this user", dir, stat.Uid)
}
if err := checkWriteBits(dir, info, stat.Uid, stat.Gid); err != nil {
return err
}
if parent := filepath.Dir(dir); parent == dir {
return nil
}
}
}
func checkWriteBits(path string, info os.FileInfo, uid, gid uint32) error {
// On a directory the sticky bit stands in for the write bits: whoever may
// write there still cannot replace an entry they do not own, which is the
// only thing that would matter to us. /tmp is the usual example.
sticky := info.IsDir() && info.Mode()&os.ModeSticky != 0
return writeBitsAllow(path, info.Mode().Perm(), sticky, groupWriteAllowed(uid, gid))
}
// writeBitsAllow decides on the permission bits alone, given whether the group's
// write access has been vouched for.
func writeBitsAllow(path string, perm os.FileMode, sticky, groupAllowed bool) error {
if sticky {
return nil
}
if perm&0o020 != 0 && !groupAllowed {
return fmt.Errorf("%s is writable by a group with members other than its owner (%v)", path, perm)
}
if perm&0o002 != 0 {
return fmt.Errorf("%s is world-writable (%v)", path, perm)
}
return nil
}
// groupWriteAllowed reports whether a group's write access to a file owned by uid
// puts it in reach of anyone who could not already act as that owner.
//
// Two ways it does not. A group in adminWriteGIDs holds the accounts that can
// answer the elevation prompt anyway. And a user private group is how Debian,
// Ubuntu and Fedora ship: their umask of 002 makes a home directory and
// everything built in it group-writable, so refusing that would refuse every
// build not installed from a package.
func groupWriteAllowed(uid, gid uint32) bool {
if slices.Contains(adminWriteGIDs, gid) {
return true
}
group, err := getent.LookupGroupID(strconv.FormatUint(uint64(gid), 10))
if err != nil {
log.Debugf("cannot look up group %d, treating it as shared: %v", gid, err)
return false
}
owner, err := getent.LookupUserID(strconv.FormatUint(uint64(uid), 10))
if err != nil {
log.Debugf("cannot look up uid %d, treating its group as shared: %v", uid, err)
return false
}
if group.Name != owner.Username {
return false
}
return !groupHasOtherMembers(group.Name, owner.Username)
}
// groupHasOtherMembers reports whether the group lists a member besides owner.
//
// Sharing the owner's name is what a user private group is recognised by, and it
// says nothing about who is in it: a group that has since gained a member is
// still named that way, and that member can write whatever the group can. So the
// membership is read rather than assumed. A group whose members cannot be
// listed, because no source on this host describes it, is treated as shared:
// the name alone cannot vouch for who writes through it.
func groupHasOtherMembers(name, owner string) bool {
members, err := getent.GroupMembers(name)
if err != nil {
log.Debugf("cannot list the members of group %q, treating it as shared: %v", name, err)
return true
}
return slices.ContainsFunc(members, func(member string) bool { return member != owner })
}
@@ -0,0 +1,148 @@
//go:build !windows
package elevate
import (
"os"
"os/user"
"path/filepath"
"strconv"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// ownerOnlyDir is t.TempDir() with the write bits tightened. testing creates its
// numbered directory with 0777 minus the umask, so under the common 002 umask it
// is group-writable and would fail the check under test on its own.
func ownerOnlyDir(t *testing.T) string {
t.Helper()
dir := t.TempDir()
require.NoError(t, os.Chmod(dir, 0o755), "tighten the temporary directory")
return dir
}
// writeExecutable creates a plain executable file, the shape trustedSelf checks.
func writeExecutable(t *testing.T, dir string) string {
t.Helper()
path := filepath.Join(dir, "netbird-ui")
require.NoError(t, os.WriteFile(path, []byte("#!/bin/sh\n"), 0o755), "write the executable")
require.NoError(t, os.Chmod(path, 0o755), "set the executable's mode")
return path
}
func TestCheckOnlyOwnerWritableAcceptsOwnerOnly(t *testing.T) {
err := checkOnlyOwnerWritable(writeExecutable(t, ownerOnlyDir(t)))
assert.NoError(t, err, "an owner-only writable executable is trustworthy")
}
func TestCheckOnlyOwnerWritableRejectsWorldWritableFile(t *testing.T) {
path := writeExecutable(t, ownerOnlyDir(t))
require.NoError(t, os.Chmod(path, 0o777), "make the executable world-writable")
assert.Error(t, checkOnlyOwnerWritable(path), "a world-writable executable must be refused")
}
// The permission policy on its own, without a filesystem to arrange: whether the
// group has been vouched for is the only thing that makes group write acceptable.
func TestWriteBitsAllow(t *testing.T) {
tests := []struct {
name string
perm os.FileMode
sticky bool
groupAllowed bool
wantErr bool
}{
{name: "owner only", perm: 0o755},
{name: "group write in a private group", perm: 0o775, groupAllowed: true},
{name: "group write in a shared group", perm: 0o775, wantErr: true},
{name: "world write", perm: 0o777, groupAllowed: true, wantErr: true},
{name: "world write on a sticky directory", perm: 0o777, sticky: true},
{name: "group write on a sticky directory", perm: 0o775, sticky: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := writeBitsAllow("/path", tt.perm, tt.sticky, tt.groupAllowed)
if tt.wantErr {
assert.Error(t, err, "perm %v, sticky %v, group allowed %v", tt.perm, tt.sticky, tt.groupAllowed)
return
}
assert.NoError(t, err, "perm %v, sticky %v, group allowed %v", tt.perm, tt.sticky, tt.groupAllowed)
})
}
}
// A build under a home directory on a distribution with a 002 umask, which is what
// a locally built or tarball-installed binary looks like. Its group has no members
// but its owner, so it is as good as owner-only.
//
// Whether this host is such a distribution is read from the environment rather than
// from groupWriteAllowed: asking the function under test whether to run would let
// it skip its own coverage away if it regressed to refusing everything.
func TestCheckOnlyOwnerWritableAcceptsOwnPrivateGroup(t *testing.T) {
requirePrivatePrimaryGroup(t)
dir := ownerOnlyDir(t)
path := writeExecutable(t, dir)
require.NoError(t, os.Chmod(dir, 0o775), "make the directory group-writable")
require.NoError(t, os.Chmod(path, 0o775), "make the executable group-writable")
err := checkOnlyOwnerWritable(path)
assert.NoError(t, err, "group write in the owner's own private group reaches nobody else")
}
// A group whose membership no source can answer for is treated as shared: the
// private-group allowance must not stand on a name nobody can vouch for. The
// membership listing itself lives in the getent package and is tested there.
func TestGroupHasOtherMembersRejectsAnUnknownGroup(t *testing.T) {
assert.True(t, groupHasOtherMembers("nonexistent_group_xyzzy_12345", "vma"),
"a group no source describes")
}
// A writable directory is as good as a writable file: whoever can write the
// directory can put a different binary at the same path.
func TestCheckOnlyOwnerWritableRejectsWritableDirectory(t *testing.T) {
dir := filepath.Join(ownerOnlyDir(t), "bin")
require.NoError(t, os.Mkdir(dir, 0o755), "create the directory")
path := writeExecutable(t, dir)
require.NoError(t, os.Chmod(dir, 0o777), "make the directory world-writable")
assert.Error(t, checkOnlyOwnerWritable(path), "an executable in a world-writable directory must be refused")
}
// A sticky world-writable directory is exempt: the sticky bit is what stops one
// user replacing another's entries. /tmp is why this matters.
func TestCheckOnlyOwnerWritableAcceptsStickyDirectory(t *testing.T) {
dir := filepath.Join(ownerOnlyDir(t), "sticky")
require.NoError(t, os.Mkdir(dir, 0o755), "create the directory")
path := writeExecutable(t, dir)
require.NoError(t, os.Chmod(dir, 0o777|os.ModeSticky), "make the directory sticky and world-writable")
err := checkOnlyOwnerWritable(path)
assert.NoError(t, err, "the sticky bit stops another user replacing the executable")
}
func TestCheckOnlyOwnerWritableRejectsMissingFile(t *testing.T) {
err := checkOnlyOwnerWritable(filepath.Join(ownerOnlyDir(t), "absent"))
assert.Error(t, err, "an executable that is not there must be refused")
}
// requirePrivatePrimaryGroup skips unless this user's primary group is their own,
// which is what the user-private-group allowance is about.
func requirePrivatePrimaryGroup(t *testing.T) {
t.Helper()
self, err := user.Current()
require.NoError(t, err, "look up the test user")
group, err := user.LookupGroupId(strconv.Itoa(os.Getgid()))
require.NoError(t, err, "look up the test user's primary group")
if group.Name != self.Username {
t.Skipf("the test user's primary group is %q, not their own, so there is nothing to assert here", group.Name)
}
if groupHasOtherMembers(group.Name, self.Username) {
t.Skipf("group %q has other members, so it is not a private group", group.Name)
}
}
+215
View File
@@ -0,0 +1,215 @@
package elevate
import (
"errors"
"fmt"
"path/filepath"
"slices"
"unsafe"
"golang.org/x/sys/windows"
)
const (
// fileDeleteChild is FILE_DELETE_CHILD, which x/sys does not define: the
// right to delete an entry of a directory without holding DELETE on it.
fileDeleteChild = 0x00000040
// accessAllowedCallbackACEType is an allow ACE with a condition appended to
// the ACCESS_ALLOWED_ACE layout, so its trustee is still at SidStart.
accessAllowedCallbackACEType = 0x9
// The allow ACE types that carry object GUIDs ahead of the trustee, so the
// SID is not at SidStart. They occur on directory-service objects rather
// than files, and are refused rather than skipped: see aceTrustee.
accessAllowedObjectACEType = 0x5
accessAllowedCallbackObjectACEType = 0xB
)
// fileWriteAccess are the rights that let a trustee rewrite or replace a file,
// or take it over and then do so.
const fileWriteAccess = windows.FILE_WRITE_DATA | windows.FILE_APPEND_DATA |
windows.DELETE | windows.WRITE_DAC | windows.WRITE_OWNER |
windows.GENERIC_WRITE | windows.GENERIC_ALL
// dirWriteAccess are the rights over a directory that let a trustee replace an
// entry somebody else owns. Creating a new entry is not one of them, which is
// what the Unix sticky bit says in one bit: the root of every volume grants
// BUILTIN\Users the right to add directories under it, and that reaches nothing
// already there.
const dirWriteAccess = fileDeleteChild | windows.DELETE |
windows.WRITE_DAC | windows.WRITE_OWNER | windows.GENERIC_ALL
// trustedInstallerSID owns much of what Windows itself installs. x/sys has no
// well-known constant for it.
const trustedInstallerSID = "S-1-5-80-956008885-3418522649-1831038044-1853292631-2271478464"
// checkOnlyOwnerWritable reports an error unless path, and every directory
// leading to it, is owned by an account that can elevate (or by this user) and
// grants write access to nobody else. A writable directory is as good as a
// writable file, since an entry in it can be replaced, so the whole chain is
// checked.
func checkOnlyOwnerWritable(path string) error {
owners, err := trustedOwners()
if err != nil {
return err
}
writers, err := trustedWriters(owners)
if err != nil {
return err
}
writeAccess := windows.ACCESS_MASK(fileWriteAccess)
for target := path; ; target = filepath.Dir(target) {
if err := checkSecurity(target, writeAccess, owners, writers); err != nil {
return err
}
if parent := filepath.Dir(target); parent == target {
return nil
}
writeAccess = dirWriteAccess
}
}
// trustedOwners are the accounts we accept as the owner of the executable and of
// the directories above it: the ones that can already answer the UAC prompt,
// plus this user, whose own executable is theirs to write. Code running as the
// user could prompt them for anything anyway; what matters is that no *other*
// unprivileged account can reach it.
func trustedOwners() ([]*windows.SID, error) {
self, err := currentUserSID()
if err != nil {
return nil, err
}
owners := []*windows.SID{self}
for _, wellKnown := range []windows.WELL_KNOWN_SID_TYPE{
windows.WinLocalSystemSid,
windows.WinBuiltinAdministratorsSid,
} {
sid, err := windows.CreateWellKnownSid(wellKnown)
if err != nil {
return nil, fmt.Errorf("build well-known SID %d: %w", wellKnown, err)
}
owners = append(owners, sid)
}
installer, err := windows.StringToSid(trustedInstallerSID)
if err != nil {
return nil, fmt.Errorf("parse TrustedInstaller SID: %w", err)
}
return append(owners, installer), nil
}
// trustedWriters are the trustees whose write access does not widen who could
// decide what runs behind the prompt. The owners, and CREATOR OWNER, which
// resolves to the object's owner and is therefore already vetted.
func trustedWriters(owners []*windows.SID) ([]*windows.SID, error) {
creatorOwner, err := windows.CreateWellKnownSid(windows.WinCreatorOwnerSid)
if err != nil {
return nil, fmt.Errorf("build the CREATOR OWNER SID: %w", err)
}
return append(slices.Clone(owners), creatorOwner), nil
}
func checkSecurity(path string, writeAccess windows.ACCESS_MASK, owners, writers []*windows.SID) error {
sd, err := windows.GetNamedSecurityInfo(path, windows.SE_FILE_OBJECT,
windows.OWNER_SECURITY_INFORMATION|windows.DACL_SECURITY_INFORMATION)
if err != nil {
return fmt.Errorf("read security descriptor of %s: %w", path, err)
}
owner, _, err := sd.Owner()
if err != nil {
return fmt.Errorf("read owner of %s: %w", path, err)
}
if !containsSID(owners, owner) {
return fmt.Errorf("%s is owned by %s, which is neither this user nor an account that can elevate", path, owner)
}
dacl, _, err := sd.DACL()
if err != nil {
return fmt.Errorf("read DACL of %s: %w", path, err)
}
// A NULL DACL grants everyone everything; only an absent security
// descriptor would have got us here without one, and neither is trustworthy.
if dacl == nil {
return fmt.Errorf("%s has no DACL, so it grants write access to everyone", path)
}
return checkDACL(path, dacl, writeAccess, writers)
}
// checkDACL refuses an ACL that grants write access to a trustee outside
// writers.
//
// An allowlist, because the trustees that must not have it cannot be listed: an
// ACE naming an ordinary user account hands that account the same power as one
// naming Everyone, and only the accounts that may hold it are knowable.
func checkDACL(path string, dacl *windows.ACL, writeAccess windows.ACCESS_MASK, writers []*windows.SID) error {
for i := uint32(0); i < uint32(dacl.AceCount); i++ {
var ace *windows.ACCESS_ALLOWED_ACE
if err := windows.GetAce(dacl, i, &ace); err != nil {
return fmt.Errorf("read ACE %d of %s: %w", i, path, err)
}
// An inherit-only ACE says what children of this object get, not what
// this object grants.
if ace.Header.AceFlags&windows.INHERIT_ONLY_ACE != 0 {
continue
}
if ace.Mask&writeAccess == 0 {
continue
}
// Only an allow ACE grants anything; a deny ACE narrows what one gave.
if !isAllowACE(ace.Header.AceType) {
continue
}
trustee, err := aceTrustee(ace)
if err != nil {
return fmt.Errorf("read the trustee of ACE %d of %s: %w", i, path, err)
}
if !containsSID(writers, trustee) {
return fmt.Errorf("%s grants write access to %s", path, trustee)
}
}
return nil
}
// isAllowACE reports whether an ACE type grants rights, rather than denying,
// auditing or labelling them.
func isAllowACE(aceType uint8) bool {
switch aceType {
case windows.ACCESS_ALLOWED_ACE_TYPE, accessAllowedCallbackACEType,
accessAllowedObjectACEType, accessAllowedCallbackObjectACEType:
return true
default:
return false
}
}
// aceTrustee returns who an allow ACE grants its rights to. An ACE whose trustee
// cannot be located is an error rather than something to skip past: being unable
// to read who is being given write access is a refusal.
func aceTrustee(ace *windows.ACCESS_ALLOWED_ACE) (*windows.SID, error) {
switch ace.Header.AceType {
case windows.ACCESS_ALLOWED_ACE_TYPE, accessAllowedCallbackACEType:
//nolint:gosec // SidStart is the first uint32 of the variable-length SID that follows the ACE header.
return (*windows.SID)(unsafe.Pointer(&ace.SidStart)), nil
default:
return nil, errors.New("an object-type allow ACE does not carry its trustee where we can read it")
}
}
func containsSID(sids []*windows.SID, sid *windows.SID) bool {
return slices.ContainsFunc(sids, sid.Equals)
}
func currentUserSID() (*windows.SID, error) {
token := windows.GetCurrentProcessToken()
user, err := token.GetTokenUser()
if err != nil {
return nil, fmt.Errorf("read this process's user: %w", err)
}
return user.User.Sid, nil
}
@@ -0,0 +1,126 @@
package elevate
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/sys/windows"
)
// A file the test user created under their own profile, which is what a per-user
// install looks like. The whole chain up to the volume root is walked, so this is
// also what says the walk does not refuse an ordinary Windows installation: the
// root of every volume grants BUILTIN\Users rights that are not ours to worry
// about.
func TestCheckOnlyOwnerWritableAcceptsOwnFile(t *testing.T) {
err := checkOnlyOwnerWritable(writeExecutable(t))
assert.NoError(t, err, "a file the test user owns, under directories only administrators can write")
}
// Write access held by an account that cannot answer the UAC prompt means that
// account decides what runs behind it, whoever the ACE names. The trustees that
// must not have it cannot be listed, so the check names the ones that may.
func TestCheckOnlyOwnerWritableRejectsUntrustedWriters(t *testing.T) {
tests := []struct {
name string
wellKnown windows.WELL_KNOWN_SID_TYPE
}{
{name: "everyone", wellKnown: windows.WinWorldSid},
{name: "authenticated users", wellKnown: windows.WinAuthenticatedUserSid},
{name: "builtin users", wellKnown: windows.WinBuiltinUsersSid},
// A service account, which no denylist of the obvious groups would name
// and which cannot elevate any more than Everyone can.
{name: "local service", wellKnown: windows.WinLocalServiceSid},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
path := writeExecutable(t)
grantWrite(t, path, tt.wellKnown)
assert.Error(t, checkOnlyOwnerWritable(path),
"write access for %s must be refused", tt.name)
})
}
}
// The masks are the policy: on a file any write reaches its contents, while on a
// directory only deleting or taking over an entry reaches something already
// there. Adding an entry does not, which is why the walk survives a volume root.
func TestWriteAccessMasks(t *testing.T) {
assert.NotZero(t, fileWriteAccess&windows.FILE_WRITE_DATA, "writing a file's data reaches its contents")
assert.NotZero(t, fileWriteAccess&windows.FILE_APPEND_DATA, "appending to a file reaches its contents")
assert.Zero(t, dirWriteAccess&windows.FILE_WRITE_DATA, "adding a file to a directory replaces nothing")
assert.Zero(t, dirWriteAccess&windows.FILE_APPEND_DATA, "adding a subdirectory replaces nothing")
assert.NotZero(t, dirWriteAccess&fileDeleteChild, "deleting an entry replaces it")
assert.NotZero(t, dirWriteAccess&windows.DELETE, "deleting the directory takes its entries with it")
}
func TestIsAllowACE(t *testing.T) {
tests := []struct {
name string
aceType uint8
want bool
}{
{name: "allowed", aceType: windows.ACCESS_ALLOWED_ACE_TYPE, want: true},
{name: "allowed callback", aceType: accessAllowedCallbackACEType, want: true},
{name: "allowed object", aceType: accessAllowedObjectACEType, want: true},
{name: "allowed callback object", aceType: accessAllowedCallbackObjectACEType, want: true},
{name: "denied", aceType: windows.ACCESS_DENIED_ACE_TYPE},
// SYSTEM_AUDIT_ACE_TYPE, which x/sys does not define: an ACE that records
// access rather than granting it.
{name: "audit", aceType: 0x2},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assert.Equal(t, tt.want, isAllowACE(tt.aceType), "ACE type %#x", tt.aceType)
})
}
}
// writeExecutable creates a plain file under the test's own directory, the shape
// trustedSelf checks.
func writeExecutable(t *testing.T) string {
t.Helper()
path := filepath.Join(t.TempDir(), "netbird-ui.exe")
require.NoError(t, os.WriteFile(path, []byte("MZ"), 0o755), "write the executable")
return path
}
// grantWrite replaces the file's DACL with one that grants a well-known trustee
// everything, keeping the test user's own access so the file stays deletable.
func grantWrite(t *testing.T, path string, wellKnown windows.WELL_KNOWN_SID_TYPE) {
t.Helper()
trustee, err := windows.CreateWellKnownSid(wellKnown)
require.NoError(t, err, "build the trustee SID")
self, err := currentUserSID()
require.NoError(t, err, "read the test user's SID")
acl, err := windows.ACLFromEntries([]windows.EXPLICIT_ACCESS{
fullControl(self, windows.TRUSTEE_IS_USER),
fullControl(trustee, windows.TRUSTEE_IS_WELL_KNOWN_GROUP),
}, nil)
require.NoError(t, err, "build the ACL")
require.NoError(t, windows.SetNamedSecurityInfo(path, windows.SE_FILE_OBJECT,
windows.DACL_SECURITY_INFORMATION|windows.PROTECTED_DACL_SECURITY_INFORMATION,
nil, nil, acl, nil), "set the DACL")
}
func fullControl(sid *windows.SID, trusteeType uint32) windows.EXPLICIT_ACCESS {
return windows.EXPLICIT_ACCESS{
AccessPermissions: windows.GENERIC_ALL,
AccessMode: windows.GRANT_ACCESS,
Trustee: windows.TRUSTEE{
TrusteeForm: windows.TRUSTEE_IS_SID,
TrusteeType: windows.TRUSTEE_TYPE(trusteeType),
TrusteeValue: windows.TrusteeValueFromSID(sid),
},
}
}
+32 -47
View File
@@ -60,6 +60,7 @@ import (
"github.com/netbirdio/netbird/client/internal/syncstore" "github.com/netbirdio/netbird/client/internal/syncstore"
"github.com/netbirdio/netbird/client/internal/updater" "github.com/netbirdio/netbird/client/internal/updater"
"github.com/netbirdio/netbird/client/jobexec" "github.com/netbirdio/netbird/client/jobexec"
"github.com/netbirdio/netbird/client/netevents"
cProto "github.com/netbirdio/netbird/client/proto" cProto "github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/system" "github.com/netbirdio/netbird/client/system"
nbdns "github.com/netbirdio/netbird/dns" nbdns "github.com/netbirdio/netbird/dns"
@@ -183,6 +184,9 @@ type EngineServices struct {
ClientMetrics *metrics.ClientMetrics ClientMetrics *metrics.ClientMetrics
MetricsCtx context.Context MetricsCtx context.Context
FileDrop *filedrop.Manager FileDrop *filedrop.Manager
// NetMgr gates the reconnection loops on OS-reported network
// availability; nil disables gating.
NetMgr *netevents.Manager
} }
// Engine is a mechanism responsible for reacting on Signal and Management stream events and managing connections to the remote peers. // Engine is a mechanism responsible for reacting on Signal and Management stream events and managing connections to the remote peers.
@@ -206,6 +210,10 @@ type Engine struct {
config *EngineConfig config *EngineConfig
mobileDep MobileDependency mobileDep MobileDependency
// netMgr gates the peer reconnection guards on OS-reported network
// availability; nil disables gating.
netMgr *netevents.Manager
// STUNs is a list of STUN servers used by ICE // STUNs is a list of STUN servers used by ICE
STUNs []*stun.URI STUNs []*stun.URI
// TURNs is a list of STUN servers used by ICE // TURNs is a list of STUN servers used by ICE
@@ -343,6 +351,7 @@ func NewEngine(
syncMsgMux: &sync.Mutex{}, syncMsgMux: &sync.Mutex{},
config: config, config: config,
mobileDep: mobileDep, mobileDep: mobileDep,
netMgr: services.NetMgr,
STUNs: []*stun.URI{}, STUNs: []*stun.URI{},
TURNs: []*stun.URI{}, TURNs: []*stun.URI{},
networkSerial: 0, networkSerial: 0,
@@ -872,8 +881,7 @@ func (e *Engine) modifyPeers(peersUpdate []*mgmProto.RemotePeerConfig) error {
} }
// third, add the peer connections again // third, add the peer connections again
for _, p := range modified { for _, p := range modified {
err := e.addNewPeer(p) if err := e.addNewPeer(p); err != nil {
if err != nil {
return err return err
} }
} }
@@ -1497,8 +1505,12 @@ func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap) error {
return nil return nil
} }
if err := e.connMgr.UpdatedRemoteFeatureFlag(e.ctx, networkMap.GetPeerConfig().GetLazyConnectionEnabled()); err != nil { // Only update the flag when the sync carries a peer config; a nil peer config
log.Errorf("failed to update lazy connection feature flag: %v", err) // (e.g. a partial update) must not reset the cached flag to false.
if peerConfig := networkMap.GetPeerConfig(); peerConfig != nil {
if err := e.connMgr.UpdatedRemoteFeatureFlag(e.ctx, peerConfig.GetLazyConnectionEnabled()); err != nil {
log.Errorf("failed to update lazy connection feature flag: %v", err)
}
} }
if e.firewall != nil { if e.firewall != nil {
@@ -1564,8 +1576,7 @@ func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap) error {
// Ingress forward rules // Ingress forward rules
done = e.phase("forward_rules") done = e.phase("forward_rules")
forwardingRules, err := e.updateForwardRules(networkMap.GetForwardingRules()) if _, err := e.updateForwardRules(networkMap.GetForwardingRules()); err != nil {
if err != nil {
log.Errorf("failed to update forward rules, err: %v", err) log.Errorf("failed to update forward rules, err: %v", err)
} }
done() done()
@@ -1583,8 +1594,7 @@ func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap) error {
// must set the exclude list after the peers are added. Without it the manager can not figure out the peers parameters from the store // must set the exclude list after the peers are added. Without it the manager can not figure out the peers parameters from the store
done = e.phase("lazy_exclude") done = e.phase("lazy_exclude")
excludedLazyPeers := e.toExcludedLazyPeers(forwardingRules, remotePeers) e.connMgr.SetExcludeList(e.ctx, e.toExcludedLazyPeers(remotePeers))
e.connMgr.SetExcludeList(e.ctx, excludedLazyPeers)
done() done()
e.networkSerial = serial e.networkSerial = serial
@@ -1828,15 +1838,15 @@ func addrToString(addr netip.Addr) string {
// addNewPeers adds peers that were not know before but arrived from the Management service with the update // addNewPeers adds peers that were not know before but arrived from the Management service with the update
func (e *Engine) addNewPeers(peersUpdate []*mgmProto.RemotePeerConfig) error { func (e *Engine) addNewPeers(peersUpdate []*mgmProto.RemotePeerConfig) error {
for _, p := range peersUpdate { for _, p := range peersUpdate {
err := e.addNewPeer(p) if err := e.addNewPeer(p); err != nil {
if err != nil {
return err return err
} }
} }
return nil return nil
} }
// addNewPeer add peer if connection doesn't exist // addNewPeer add peer if connection doesn't exist. A peer that is not lazy by
// policy gets an always-active connection instead.
func (e *Engine) addNewPeer(peerConfig *mgmProto.RemotePeerConfig) error { func (e *Engine) addNewPeer(peerConfig *mgmProto.RemotePeerConfig) error {
peerKey := peerConfig.GetWgPubKey() peerKey := peerConfig.GetWgPubKey()
peerIPs := make([]netip.Prefix, 0, len(peerConfig.GetAllowedIps())) peerIPs := make([]netip.Prefix, 0, len(peerConfig.GetAllowedIps()))
@@ -1871,7 +1881,8 @@ func (e *Engine) addNewPeer(peerConfig *mgmProto.RemotePeerConfig) error {
log.Warnf("error adding peer %s to status recorder, got error: %v", peerKey, err) log.Warnf("error adding peer %s to status recorder, got error: %v", peerKey, err)
} }
if exists := e.connMgr.AddPeerConn(e.ctx, peerKey, conn); exists { permanent := !e.connMgr.PeerLazyDefault(peerConfig.GetLazyState())
if exists := e.connMgr.AddPeerConn(e.ctx, peerKey, conn, permanent); exists {
conn.Close(false) conn.Close(false)
return fmt.Errorf("peer already exists: %s", peerKey) return fmt.Errorf("peer already exists: %s", peerKey)
} }
@@ -1905,6 +1916,7 @@ func (e *Engine) createPeerConn(pubKey string, allowedIPs []netip.Prefix, agentV
PermissiveMode: e.config.RosenpassPermissive, PermissiveMode: e.config.RosenpassPermissive,
}, },
ICEConfig: e.createICEConfig(), ICEConfig: e.createICEConfig(),
NetMgr: e.netMgr,
} }
serviceDependencies := peer.ServiceDependencies{ serviceDependencies := peer.ServiceDependencies{
@@ -2580,7 +2592,7 @@ func (e *Engine) SetCapture(pc device.PacketCapture) error {
} }
afc := capture.NewAFPacketCapture(intf.Name(), sess) afc := capture.NewAFPacketCapture(intf.Name(), sess)
if err := afc.Start(); err != nil { if err := afc.Start(); err != nil { //nolint:staticcheck // always errors on non-Linux builds
return fmt.Errorf("start AF_PACKET capture on %s: %w", intf.Name(), err) return fmt.Errorf("start AF_PACKET capture on %s: %w", intf.Name(), err)
} }
e.afpacketCapture = afc e.afpacketCapture = afc
@@ -2669,46 +2681,19 @@ func (e *Engine) updateForwardRules(rules []*mgmProto.ForwardingRule) ([]firewal
return forwardingRules, nberrors.FormatErrorOrNil(merr) return forwardingRules, nberrors.FormatErrorOrNil(merr)
} }
func (e *Engine) toExcludedLazyPeers(rules []firewallManager.ForwardRule, peers []*mgmProto.RemotePeerConfig) map[string]bool { // toExcludedLazyPeers returns the peers that must have an always-active
// connection: those that are not lazy by policy (the per-peer lazy state or the
// account flag, subject to the local override).
func (e *Engine) toExcludedLazyPeers(peers []*mgmProto.RemotePeerConfig) map[string]bool {
excludedPeers := make(map[string]bool) excludedPeers := make(map[string]bool)
for _, p := range peers {
// Ingress forward targets: inbound forwarded traffic is initiated remotely and if !e.connMgr.PeerLazyDefault(p.GetLazyState()) {
// cannot wake a lazy connection, so the peer routing the target must stay excludedPeers[p.GetWgPubKey()] = true
// permanently connected. AllowedIPs are already parsed on the peer conn, so
// reuse those typed prefixes instead of re-parsing the network map strings.
for _, r := range rules {
for _, p := range peers {
if e.peerRoutesAddr(p, r.TranslatedAddress) {
log.Infof("exclude forwarder peer from lazy connection: %s", p.GetWgPubKey())
excludedPeers[p.GetWgPubKey()] = true
}
} }
} }
return excludedPeers return excludedPeers
} }
// peerRoutesAddr reports whether the peer is a router for addr, matched against
// the peer's already-parsed AllowedIPs from the store (the same typed value the
// lazy manager consumes) rather than re-parsing the network map strings.
func (e *Engine) peerRoutesAddr(p *mgmProto.RemotePeerConfig, addr netip.Addr) bool {
prefixes, ok := e.peerStore.AllowedIPs(p.GetWgPubKey())
if !ok {
return false
}
return prefixesContain(prefixes, addr)
}
// prefixesContain reports whether addr falls within any of the prefixes.
func prefixesContain(prefixes []netip.Prefix, addr netip.Addr) bool {
for _, prefix := range prefixes {
if prefix.Contains(addr) {
return true
}
}
return false
}
// isChecksEqual checks if two slices of checks are equal. // isChecksEqual checks if two slices of checks are equal.
func isChecksEqual(checks1, checks2 []*mgmProto.Checks) bool { func isChecksEqual(checks1, checks2 []*mgmProto.Checks) bool {
normalize := func(checks []*mgmProto.Checks) []string { normalize := func(checks []*mgmProto.Checks) []string {
@@ -1,87 +0,0 @@
package internal
import (
"net/netip"
"testing"
"github.com/stretchr/testify/require"
firewallManager "github.com/netbirdio/netbird/client/firewall/manager"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/peerstore"
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
)
func TestPrefixesContain(t *testing.T) {
tests := []struct {
name string
prefixes []string
addr string
want bool
}{
{name: "own overlay /32 matches", prefixes: []string{"100.110.8.145/32"}, addr: "100.110.8.145", want: true},
{name: "addr inside routed subnet", prefixes: []string{"10.121.0.0/16"}, addr: "10.121.208.4", want: true},
{name: "addr outside subnet", prefixes: []string{"10.121.0.0/16"}, addr: "10.122.0.1", want: false},
{name: "different /32", prefixes: []string{"100.110.8.145/32"}, addr: "100.110.8.146", want: false},
{name: "ipv6 /128 matches", prefixes: []string{"fd00::1/128"}, addr: "fd00::1", want: true},
{name: "no prefixes", prefixes: nil, addr: "10.121.208.4", want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
prefixes := make([]netip.Prefix, 0, len(tt.prefixes))
for _, p := range tt.prefixes {
prefixes = append(prefixes, netip.MustParsePrefix(p))
}
require.Equal(t, tt.want, prefixesContain(prefixes, netip.MustParseAddr(tt.addr)))
})
}
}
// TestToExcludedLazyPeers_ForwardTarget guards a regression: the forward-target
// peer (the peer routing a ForwardRule.TranslatedAddress) must be excluded from
// lazy connections, matched via the peer's already-parsed AllowedIPs.
func TestToExcludedLazyPeers_ForwardTarget(t *testing.T) {
const targetPeerKey = "cccccccccccccccccccccccccccccccccccccccccc0="
const otherPeerKey = "dddddddddddddddddddddddddddddddddddddddddd0="
store := peerstore.NewConnStore()
store.AddPeerConn(targetPeerKey, newTestConn(t, targetPeerKey, "100.110.8.145/32"))
store.AddPeerConn(otherPeerKey, newTestConn(t, otherPeerKey, "100.110.9.10/32"))
e := &Engine{peerStore: store}
peers := []*mgmProto.RemotePeerConfig{
{WgPubKey: targetPeerKey, AllowedIps: []string{"100.110.8.145/32"}},
{WgPubKey: otherPeerKey, AllowedIps: []string{"100.110.9.10/32"}},
}
rules := []firewallManager.ForwardRule{
{TranslatedAddress: netip.MustParseAddr("100.110.8.145")},
}
excluded := e.toExcludedLazyPeers(rules, peers)
require.True(t, excluded[targetPeerKey], "forward-target peer must be excluded from lazy connections")
require.False(t, excluded[otherPeerKey], "non-target peer must not be excluded")
require.Len(t, excluded, 1)
}
func TestToExcludedLazyPeers_NoRules(t *testing.T) {
e := &Engine{peerStore: peerstore.NewConnStore()}
peers := []*mgmProto.RemotePeerConfig{
{WgPubKey: "peer-a", AllowedIps: []string{"100.110.8.145/32"}},
}
require.Empty(t, e.toExcludedLazyPeers(nil, peers))
}
func newTestConn(t *testing.T, key, allowedIP string) *peer.Conn {
t.Helper()
conn, err := peer.NewConn(peer.ConnConfig{
Key: key,
WgConfig: peer.WgConfig{AllowedIps: []netip.Prefix{netip.MustParsePrefix(allowedIP)}},
}, peer.ServiceDependencies{})
require.NoError(t, err)
return conn
}
+1 -1
View File
@@ -12,7 +12,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/golang/mock/gomock" "go.uber.org/mock/gomock"
"github.com/google/uuid" "github.com/google/uuid"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
+2 -1
View File
@@ -279,7 +279,8 @@ func TestEngine_UpdateNetworkMap(t *testing.T) {
}, MobileDependency{}) }, MobileDependency{})
wgIface := &MockWGIface{ wgIface := &MockWGIface{
NameFunc: func() string { return "utun102" }, NameFunc: func() string { return "utun102" },
IsUserspaceBindFunc: func() bool { return true },
RemovePeerFunc: func(peerKey string) error { RemovePeerFunc: func(peerKey string) error {
return nil return nil
}, },
+36
View File
@@ -0,0 +1,36 @@
//go:build cgo && !osusergo && !windows
package getent
import "os/user"
// Built with cgo, os/user resolves through libc (getpwnam_r and friends),
// which goes through the host's NSS stack natively. Whatever it fails to
// find, the getent command would not find either, so there is nothing to
// fall back to.
// LookupUser looks up a user by name.
func LookupUser(username string) (*user.User, error) {
return user.Lookup(username)
}
// LookupUserID looks up a user by UID.
func LookupUserID(uid string) (*user.User, error) {
return user.LookupId(uid)
}
// CurrentUser returns the user this process runs as.
func CurrentUser() (*user.User, error) {
return user.Current()
}
// LookupGroupID looks up a group by GID.
func LookupGroupID(gid string) (*user.Group, error) {
return user.LookupGroupId(gid)
}
// GroupIDs returns the IDs of the groups the user is a member of; libc's
// getgrouplist handles NSS groups natively.
func GroupIDs(u *user.User) ([]string, error) {
return u.GroupIds()
}
+6
View File
@@ -0,0 +1,6 @@
// Package getent resolves users and groups through the host's NSS stack.
// Built without cgo, os/user reads /etc/passwd and /etc/group alone and misses
// anything LDAP, SSSD or winbind provide; the getent and id commands resolve
// through NSS whatever the build. The lookups here try the standard library
// first, which needs no subprocess, and fall back to those commands.
package getent
@@ -1,4 +1,4 @@
package server package getent
import ( import (
"os/user" "os/user"
@@ -10,38 +10,48 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
func TestLookupWithGetent_CurrentUser(t *testing.T) { func TestLookupUser_CurrentUser(t *testing.T) {
// The current user should always be resolvable on any platform // The current user should always be resolvable on any platform
current, err := user.Current() current, err := user.Current()
require.NoError(t, err) require.NoError(t, err)
u, err := lookupWithGetent(current.Username) u, err := LookupUser(current.Username)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, current.Username, u.Username) assert.Equal(t, current.Username, u.Username)
assert.Equal(t, current.Uid, u.Uid) assert.Equal(t, current.Uid, u.Uid)
assert.Equal(t, current.Gid, u.Gid) assert.Equal(t, current.Gid, u.Gid)
} }
func TestLookupWithGetent_NonexistentUser(t *testing.T) { func TestLookupUser_NonexistentUser(t *testing.T) {
_, err := lookupWithGetent("nonexistent_user_xyzzy_12345") _, err := LookupUser("nonexistent_user_xyzzy_12345")
require.Error(t, err, "should fail for nonexistent user") require.Error(t, err, "should fail for nonexistent user")
} }
func TestCurrentUserWithGetent(t *testing.T) { func TestLookupUserID_CurrentUser(t *testing.T) {
current, err := user.Current()
require.NoError(t, err)
u, err := LookupUserID(current.Uid)
require.NoError(t, err)
assert.Equal(t, current.Username, u.Username)
assert.Equal(t, current.Uid, u.Uid)
}
func TestCurrentUser(t *testing.T) {
stdUser, err := user.Current() stdUser, err := user.Current()
require.NoError(t, err) require.NoError(t, err)
u, err := currentUserWithGetent() u, err := CurrentUser()
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, stdUser.Uid, u.Uid) assert.Equal(t, stdUser.Uid, u.Uid)
assert.Equal(t, stdUser.Username, u.Username) assert.Equal(t, stdUser.Username, u.Username)
} }
func TestGroupIdsWithFallback_CurrentUser(t *testing.T) { func TestGroupIDs_CurrentUser(t *testing.T) {
current, err := user.Current() current, err := user.Current()
require.NoError(t, err) require.NoError(t, err)
groups, err := groupIdsWithFallback(current) groups, err := GroupIDs(current)
require.NoError(t, err) require.NoError(t, err)
require.NotEmpty(t, groups, "current user should have at least one group") require.NotEmpty(t, groups, "current user should have at least one group")
@@ -53,32 +63,30 @@ func TestGroupIdsWithFallback_CurrentUser(t *testing.T) {
} }
} }
func TestGetShellFromGetent_CurrentUser(t *testing.T) { func TestUserShell_CurrentUser(t *testing.T) {
if runtime.GOOS == "windows" {
// Windows stub always returns empty, which is correct
shell := getShellFromGetent("1000")
assert.Empty(t, shell, "Windows stub should return empty")
return
}
current, err := user.Current() current, err := user.Current()
require.NoError(t, err) require.NoError(t, err)
// getent may not be available on all systems (e.g., macOS without Homebrew getent) // getent may not be available on all systems (e.g., macOS without
shell := getShellFromGetent(current.Uid) // Homebrew getent), and Windows has no login shells at all.
shell, err := UserShell(current.Uid)
if err != nil {
t.Logf("UserShell failed, getent may not be available: %v", err)
return
}
if shell == "" { if shell == "" {
t.Log("getShellFromGetent returned empty, getent may not be available") t.Log("UserShell returned empty, the user has no shell set")
return return
} }
assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell) assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell)
} }
func TestLookupWithGetent_RootUser(t *testing.T) { func TestLookupUser_RootUser(t *testing.T) {
if runtime.GOOS == "windows" { if runtime.GOOS == "windows" {
t.Skip("no root user on Windows") t.Skip("no root user on Windows")
} }
u, err := lookupWithGetent("root") u, err := LookupUser("root")
if err != nil { if err != nil {
t.Skip("root user not available on this system") t.Skip("root user not available on this system")
} }
@@ -86,25 +94,25 @@ func TestLookupWithGetent_RootUser(t *testing.T) {
} }
// TestIntegration_FullLookupChain exercises the complete user lookup chain // TestIntegration_FullLookupChain exercises the complete user lookup chain
// against the real system, testing that all wrappers (lookupWithGetent, // against the real system, testing that all wrappers (LookupUser,
// currentUserWithGetent, groupIdsWithFallback, getShellFromGetent) produce // CurrentUser, GroupIDs, UserShell) produce consistent and correct results
// consistent and correct results when composed together. // when composed together.
func TestIntegration_FullLookupChain(t *testing.T) { func TestIntegration_FullLookupChain(t *testing.T) {
// Step 1: currentUserWithGetent must resolve the running user. // Step 1: CurrentUser must resolve the running user.
current, err := currentUserWithGetent() current, err := CurrentUser()
require.NoError(t, err, "currentUserWithGetent must resolve the running user") require.NoError(t, err, "CurrentUser must resolve the running user")
require.NotEmpty(t, current.Uid) require.NotEmpty(t, current.Uid)
require.NotEmpty(t, current.Username) require.NotEmpty(t, current.Username)
// Step 2: lookupWithGetent by the same username must return matching identity. // Step 2: LookupUser by the same username must return matching identity.
byName, err := lookupWithGetent(current.Username) byName, err := LookupUser(current.Username)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, current.Uid, byName.Uid, "lookup by name should return same UID") assert.Equal(t, current.Uid, byName.Uid, "lookup by name should return same UID")
assert.Equal(t, current.Gid, byName.Gid, "lookup by name should return same GID") assert.Equal(t, current.Gid, byName.Gid, "lookup by name should return same GID")
assert.Equal(t, current.HomeDir, byName.HomeDir, "lookup by name should return same home") assert.Equal(t, current.HomeDir, byName.HomeDir, "lookup by name should return same home")
// Step 3: groupIdsWithFallback must return at least the primary GID. // Step 3: GroupIDs must return at least the primary GID.
groups, err := groupIdsWithFallback(current) groups, err := GroupIDs(current)
require.NoError(t, err) require.NoError(t, err)
require.NotEmpty(t, groups, "user must have at least one group") require.NotEmpty(t, groups, "user must have at least one group")
@@ -119,29 +127,20 @@ func TestIntegration_FullLookupChain(t *testing.T) {
} }
} }
assert.True(t, foundPrimary, "primary GID %s should appear in supplementary groups", current.Gid) assert.True(t, foundPrimary, "primary GID %s should appear in supplementary groups", current.Gid)
// Step 4: getShellFromGetent should either return a valid shell path or empty
// (empty is OK when getent is not available, e.g. macOS without Homebrew getent).
if runtime.GOOS != "windows" {
shell := getShellFromGetent(current.Uid)
if shell != "" {
assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell)
}
}
} }
// TestIntegration_LookupAndGroupsConsistency verifies that a user resolved via // TestIntegration_LookupAndGroupsConsistency verifies that a user resolved via
// lookupWithGetent can have their groups resolved via groupIdsWithFallback, // LookupUser can have their groups resolved via GroupIDs, testing the handoff
// testing the handoff between the two functions as used by the SSH server. // between the two functions as used by the SSH server.
func TestIntegration_LookupAndGroupsConsistency(t *testing.T) { func TestIntegration_LookupAndGroupsConsistency(t *testing.T) {
current, err := user.Current() current, err := user.Current()
require.NoError(t, err) require.NoError(t, err)
// Simulate the SSH server flow: lookup user, then get their groups. // Simulate the SSH server flow: lookup user, then get their groups.
resolved, err := lookupWithGetent(current.Username) resolved, err := LookupUser(current.Username)
require.NoError(t, err) require.NoError(t, err)
groups, err := groupIdsWithFallback(resolved) groups, err := GroupIDs(resolved)
require.NoError(t, err) require.NoError(t, err)
require.NotEmpty(t, groups, "resolved user must have groups") require.NotEmpty(t, groups, "resolved user must have groups")
@@ -154,19 +153,3 @@ func TestIntegration_LookupAndGroupsConsistency(t *testing.T) {
} }
} }
} }
// TestIntegration_ShellLookupChain tests the full shell resolution chain
// (getShellFromPasswd -> getShellFromGetent -> $SHELL -> default) on Unix.
func TestIntegration_ShellLookupChain(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("Unix shell lookup not applicable on Windows")
}
current, err := user.Current()
require.NoError(t, err)
// getUserShell is the top-level function used by the SSH server.
shell := getUserShell(current.Uid)
require.NotEmpty(t, shell, "getUserShell must always return a shell")
assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell)
}
+110
View File
@@ -0,0 +1,110 @@
//go:build (!cgo || osusergo) && !windows
package getent
import (
"os"
"os/user"
"strconv"
log "github.com/sirupsen/logrus"
)
// Without cgo, os/user only reads /etc/passwd and /etc/group and misses
// NSS-provided users and groups; the getent and id commands go through the
// host's NSS stack.
// LookupUser looks up a user by name, falling back to getent if os/user fails.
func LookupUser(username string) (*user.User, error) {
u, err := user.Lookup(username)
if err == nil {
return u, nil
}
stdErr := err
log.Debugf("os/user.Lookup(%q) failed, trying getent: %v", username, err)
u, _, getentErr := passwdLookup(username)
if getentErr != nil {
log.Debugf("getent fallback for %q also failed: %v", username, getentErr)
return nil, stdErr
}
return u, nil
}
// LookupUserID looks up a user by UID, falling back to getent if os/user fails.
func LookupUserID(uid string) (*user.User, error) {
u, err := user.LookupId(uid)
if err == nil {
return u, nil
}
stdErr := err
log.Debugf("os/user.LookupId(%q) failed, trying getent: %v", uid, err)
u, _, getentErr := passwdLookup(uid)
if getentErr != nil {
log.Debugf("getent fallback for uid %s also failed: %v", uid, getentErr)
return nil, stdErr
}
return u, nil
}
// CurrentUser returns the user this process runs as, falling back to getent
// if os/user fails.
func CurrentUser() (*user.User, error) {
u, err := user.Current()
if err == nil {
return u, nil
}
stdErr := err
uid := strconv.Itoa(os.Getuid())
log.Debugf("os/user.Current() failed, trying getent with UID %s: %v", uid, err)
u, _, getentErr := passwdLookup(uid)
if getentErr != nil {
return nil, stdErr
}
return u, nil
}
// LookupGroupID looks up a group by GID, falling back to getent if os/user
// fails.
func LookupGroupID(gid string) (*user.Group, error) {
g, err := user.LookupGroupId(gid)
if err == nil {
return g, nil
}
stdErr := err
log.Debugf("os/user.LookupGroupId(%q) failed, trying getent: %v", gid, err)
g, _, getentErr := groupLookup(gid)
if getentErr != nil {
log.Debugf("getent fallback for gid %s also failed: %v", gid, getentErr)
return nil, stdErr
}
return g, nil
}
// GroupIDs returns the IDs of the groups the user is a member of.
// NOTE: unlike the lookups above, which try the standard library first, this
// intentionally tries `id -G` first because without cgo, user.GroupIds only
// reads /etc/group and silently returns incomplete results for NSS users
// (no error, just missing groups). The id command goes through NSS and
// returns the full set.
func GroupIDs(u *user.User) ([]string, error) {
ids, err := idGroups(u.Username)
if err == nil {
return ids, nil
}
log.Debugf("id -G %q failed, falling back to user.GroupIds(): %v", u.Username, err)
ids, stdErr := u.GroupIds()
if stdErr != nil {
return nil, stdErr
}
return ids, nil
}
+224
View File
@@ -0,0 +1,224 @@
//go:build !windows
package getent
import (
"bufio"
"context"
"fmt"
"os"
"os/exec"
"os/user"
"runtime"
"strings"
"time"
log "github.com/sirupsen/logrus"
)
const commandTimeout = 5 * time.Second
// groupFile lists which accounts are in which group, for hosts where the
// getent command is not available (macOS ships without it).
const groupFile = "/etc/group"
// UserShell returns the login shell getent reports for the user with this UID.
// It reaches shells that /etc/passwd does not list, because getent resolves
// through the host's NSS stack.
func UserShell(uid string) (string, error) {
_, shell, err := passwdLookup(uid)
if err != nil {
return "", err
}
return shell, nil
}
// GroupMembers returns the names of the group's members: from getent, which
// resolves through NSS, or from /etc/group where getent is not available. A
// group neither source describes is an error; an empty member list is not,
// since accounts with the group as their primary one are not listed in it.
func GroupMembers(name string) ([]string, error) {
_, members, err := groupLookup(name)
if err == nil {
return members, nil
}
log.Debugf("getent cannot list group %q, reading %s: %v", name, groupFile, err)
return groupMembersFromFile(groupFile, name)
}
// passwdLookup executes `getent passwd <query>`, where query is a username or
// UID, and returns the user and login shell.
func passwdLookup(query string) (*user.User, string, error) {
out, err := run("passwd", query)
if err != nil {
return nil, "", err
}
return parsePasswd(string(out))
}
// groupLookup executes `getent group <query>`, where query is a group name or
// GID, and returns the group and its member names.
func groupLookup(query string) (*user.Group, []string, error) {
out, err := run("group", query)
if err != nil {
return nil, nil, err
}
return parseGroup(string(out))
}
// run executes `getent <database> <key>` with a timeout.
func run(database, key string) ([]byte, error) {
if !validateInput(key) {
return nil, fmt.Errorf("invalid getent input: %q", key)
}
ctx, cancel := context.WithTimeout(context.Background(), commandTimeout)
defer cancel()
out, err := exec.CommandContext(ctx, "getent", database, key).Output()
if err != nil {
return nil, fmt.Errorf("getent %s %s: %w", database, key, err)
}
return out, nil
}
// parsePasswd parses getent passwd output: "name:x:uid:gid:gecos:home:shell"
func parsePasswd(output string) (*user.User, string, error) {
fields := strings.SplitN(strings.TrimSpace(output), ":", 8)
if len(fields) < 6 {
return nil, "", fmt.Errorf("unexpected getent output (need 6+ fields): %q", output)
}
if fields[0] == "" || fields[2] == "" || fields[3] == "" {
return nil, "", fmt.Errorf("missing required fields in getent output: %q", output)
}
var shell string
if len(fields) >= 7 {
shell = fields[6]
}
return &user.User{
Username: fields[0],
Uid: fields[2],
Gid: fields[3],
Name: fields[4],
HomeDir: fields[5],
}, shell, nil
}
// parseGroup parses getent group output: "name:x:gid:member,member"
func parseGroup(output string) (*user.Group, []string, error) {
fields := strings.SplitN(strings.TrimSpace(output), ":", 4)
if len(fields) < 3 {
return nil, nil, fmt.Errorf("unexpected getent output (need 3+ fields): %q", output)
}
if fields[0] == "" || fields[2] == "" {
return nil, nil, fmt.Errorf("missing required fields in getent output: %q", output)
}
var members []string
if len(fields) >= 4 {
members = splitMembers(fields[3])
}
return &user.Group{Name: fields[0], Gid: fields[2]}, members, nil
}
func splitMembers(list string) []string {
var members []string
for member := range strings.SplitSeq(list, ",") {
if member != "" {
members = append(members, member)
}
}
return members
}
// groupMembersFromFile finds the group's member list in a file of /etc/group's
// format. A group the file does not describe, because it comes from LDAP or
// another NSS source, is an error rather than an empty list.
func groupMembersFromFile(path, name string) ([]string, error) {
file, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("open %s: %w", path, err)
}
defer func() {
if err := file.Close(); err != nil {
log.Debugf("close %s: %v", path, err)
}
}()
scanner := bufio.NewScanner(file)
for scanner.Scan() {
// name:password:gid:member,member
fields := strings.Split(scanner.Text(), ":")
if len(fields) < 4 || fields[0] != name {
continue
}
return splitMembers(fields[3]), nil
}
if err := scanner.Err(); err != nil {
return nil, fmt.Errorf("read %s: %w", path, err)
}
return nil, fmt.Errorf("%s does not describe group %q", path, name)
}
// validateInput checks that the input is safe to pass to getent or id.
// Allows POSIX usernames, numeric IDs, and common NSS extensions
// (@ for Kerberos, $ for Samba, + for NIS compat). A leading hyphen is
// rejected so the input can never be parsed as a command-line flag.
func validateInput(input string) bool {
maxLen := 32
if runtime.GOOS == "linux" {
maxLen = 256
}
if len(input) == 0 || len(input) > maxLen {
return false
}
if input[0] == '-' {
return false
}
for _, r := range input {
if isAllowedChar(r) {
continue
}
return false
}
return true
}
func isAllowedChar(r rune) bool {
if r >= 'a' && r <= 'z' || r >= 'A' && r <= 'Z' || r >= '0' && r <= '9' {
return true
}
switch r {
case '.', '_', '-', '@', '+', '$':
return true
}
return false
}
// idGroups runs `id -G <username>` and returns the space-separated group IDs.
func idGroups(username string) ([]string, error) {
if !validateInput(username) {
return nil, fmt.Errorf("invalid username for id command: %q", username)
}
ctx, cancel := context.WithTimeout(context.Background(), commandTimeout)
defer cancel()
out, err := exec.CommandContext(ctx, "id", "-G", username).Output()
if err != nil {
return nil, fmt.Errorf("id -G %s: %w", username, err)
}
trimmed := strings.TrimSpace(string(out))
if trimmed == "" {
return nil, fmt.Errorf("id -G %s: empty output", username)
}
return strings.Fields(trimmed), nil
}
@@ -1,10 +1,12 @@
//go:build !windows //go:build !windows
package server package getent
import ( import (
"os"
"os/exec" "os/exec"
"os/user" "os/user"
"path/filepath"
"runtime" "runtime"
"strconv" "strconv"
"testing" "testing"
@@ -13,7 +15,7 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
func TestParseGetentPasswd(t *testing.T) { func TestParsePasswd(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
input string input string
@@ -128,7 +130,7 @@ func TestParseGetentPasswd(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
u, shell, err := parseGetentPasswd(tt.input) u, shell, err := parsePasswd(tt.input)
if tt.wantErr { if tt.wantErr {
require.Error(t, err) require.Error(t, err)
if tt.errContains != "" { if tt.errContains != "" {
@@ -147,7 +149,120 @@ func TestParseGetentPasswd(t *testing.T) {
} }
} }
func TestValidateGetentInput(t *testing.T) { func TestParseGroup(t *testing.T) {
tests := []struct {
name string
input string
wantGroup *user.Group
wantMembers []string
wantErr bool
}{
{
name: "no members",
input: "vma:x:1000:\n",
wantGroup: &user.Group{Name: "vma", Gid: "1000"},
},
{
name: "one member",
input: "sudo:x:27:alice",
wantGroup: &user.Group{Name: "sudo", Gid: "27"},
wantMembers: []string{"alice"},
},
{
name: "several members",
input: "docker:x:998:alice,bob\n",
wantGroup: &user.Group{Name: "docker", Gid: "998"},
wantMembers: []string{"alice", "bob"},
},
{
name: "too few fields",
input: "bad:x",
wantErr: true,
},
{
name: "empty group name",
input: ":x:1000:alice",
wantErr: true,
},
{
name: "empty GID",
input: "vma:x::alice",
wantErr: true,
},
{
name: "empty input",
input: "",
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
g, members, err := parseGroup(tt.input)
if tt.wantErr {
require.Error(t, err)
return
}
require.NoError(t, err)
assert.Equal(t, tt.wantGroup.Name, g.Name, "group name")
assert.Equal(t, tt.wantGroup.Gid, g.Gid, "GID")
assert.Equal(t, tt.wantMembers, members, "members")
})
}
}
func TestGroupMembersFromFile(t *testing.T) {
tests := []struct {
name string
entry string
want []string
}{
{name: "no members", entry: "vma:x:1000:"},
{name: "only the owner", entry: "vma:x:1000:vma", want: []string{"vma"}},
{name: "two members", entry: "vma:x:1000:vma,bob", want: []string{"vma", "bob"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "group")
body := "root:x:0:\n" + tt.entry + "\nsudo:x:27:vma\n"
require.NoError(t, os.WriteFile(path, []byte(body), 0o644), "write the group file")
members, err := groupMembersFromFile(path, "vma")
require.NoError(t, err, "entry %q", tt.entry)
assert.Equal(t, tt.want, members, "entry %q", tt.entry)
})
}
}
// A group the file does not describe, because it comes from LDAP or another
// NSS source, is an error rather than an empty member list: the caller must
// be able to tell "no members" from "no answer".
func TestGroupMembersFromFileUnknownGroup(t *testing.T) {
path := filepath.Join(t.TempDir(), "group")
require.NoError(t, os.WriteFile(path, []byte("root:x:0:\n"), 0o644), "write the group file")
_, err := groupMembersFromFile(path, "vma")
assert.Error(t, err, "a group the file does not describe")
_, err = groupMembersFromFile(filepath.Join(t.TempDir(), "absent"), "vma")
assert.Error(t, err, "no group file at all")
}
// GroupMembers on the root group, which every Unix has, whichever source
// answers for it.
func TestGroupMembers_RootGroup(t *testing.T) {
rootGroup := "root"
switch runtime.GOOS {
case "darwin", "dragonfly", "freebsd", "netbsd", "openbsd":
rootGroup = "wheel"
}
_, err := GroupMembers(rootGroup)
assert.NoError(t, err, "the %s group must be describable", rootGroup)
}
func TestValidateInput(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
input string input string
@@ -180,7 +295,7 @@ func TestValidateGetentInput(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
assert.Equal(t, tt.want, validateGetentInput(tt.input)) assert.Equal(t, tt.want, validateInput(tt.input))
}) })
} }
} }
@@ -193,12 +308,12 @@ func makeLongString(n int) string {
return string(b) return string(b)
} }
func TestRunGetent_RootUser(t *testing.T) { func TestPasswdLookup_RootUser(t *testing.T) {
if _, err := exec.LookPath("getent"); err != nil { if _, err := exec.LookPath("getent"); err != nil {
t.Skip("getent not available on this system") t.Skip("getent not available on this system")
} }
u, shell, err := runGetent("root") u, shell, err := passwdLookup("root")
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "root", u.Username) assert.Equal(t, "root", u.Username)
assert.Equal(t, "0", u.Uid) assert.Equal(t, "0", u.Uid)
@@ -206,44 +321,55 @@ func TestRunGetent_RootUser(t *testing.T) {
assert.NotEmpty(t, shell, "root should have a shell") assert.NotEmpty(t, shell, "root should have a shell")
} }
func TestRunGetent_ByUID(t *testing.T) { func TestPasswdLookup_ByUID(t *testing.T) {
if _, err := exec.LookPath("getent"); err != nil { if _, err := exec.LookPath("getent"); err != nil {
t.Skip("getent not available on this system") t.Skip("getent not available on this system")
} }
u, _, err := runGetent("0") u, _, err := passwdLookup("0")
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "root", u.Username) assert.Equal(t, "root", u.Username)
assert.Equal(t, "0", u.Uid) assert.Equal(t, "0", u.Uid)
} }
func TestRunGetent_NonexistentUser(t *testing.T) { func TestPasswdLookup_NonexistentUser(t *testing.T) {
if _, err := exec.LookPath("getent"); err != nil { if _, err := exec.LookPath("getent"); err != nil {
t.Skip("getent not available on this system") t.Skip("getent not available on this system")
} }
_, _, err := runGetent("nonexistent_user_xyzzy_12345") _, _, err := passwdLookup("nonexistent_user_xyzzy_12345")
assert.Error(t, err) assert.Error(t, err)
} }
func TestRunGetent_InvalidInput(t *testing.T) { func TestPasswdLookup_InvalidInput(t *testing.T) {
_, _, err := runGetent("") _, _, err := passwdLookup("")
assert.Error(t, err) assert.Error(t, err)
_, _, err = runGetent("user\x00name") _, _, err = passwdLookup("user\x00name")
assert.Error(t, err) assert.Error(t, err)
} }
func TestRunGetent_NotAvailable(t *testing.T) { func TestPasswdLookup_NotAvailable(t *testing.T) {
if _, err := exec.LookPath("getent"); err == nil { if _, err := exec.LookPath("getent"); err == nil {
t.Skip("getent is available, can't test missing case") t.Skip("getent is available, can't test missing case")
} }
_, _, err := runGetent("root") _, _, err := passwdLookup("root")
assert.Error(t, err, "should fail when getent is not installed") assert.Error(t, err, "should fail when getent is not installed")
} }
func TestRunIdGroups_CurrentUser(t *testing.T) { func TestGroupLookup_RootGroup(t *testing.T) {
if _, err := exec.LookPath("getent"); err != nil {
t.Skip("getent not available on this system")
}
g, _, err := groupLookup("0")
require.NoError(t, err)
assert.Equal(t, "0", g.Gid, "GID 0 resolves to the root group")
assert.NotEmpty(t, g.Name, "the root group has a name")
}
func TestIdGroups_CurrentUser(t *testing.T) {
if _, err := exec.LookPath("id"); err != nil { if _, err := exec.LookPath("id"); err != nil {
t.Skip("id not available on this system") t.Skip("id not available on this system")
} }
@@ -251,7 +377,7 @@ func TestRunIdGroups_CurrentUser(t *testing.T) {
current, err := user.Current() current, err := user.Current()
require.NoError(t, err) require.NoError(t, err)
groups, err := runIdGroups(current.Username) groups, err := idGroups(current.Username)
require.NoError(t, err) require.NoError(t, err)
require.NotEmpty(t, groups, "current user should have at least one group") require.NotEmpty(t, groups, "current user should have at least one group")
@@ -261,20 +387,20 @@ func TestRunIdGroups_CurrentUser(t *testing.T) {
} }
} }
func TestRunIdGroups_NonexistentUser(t *testing.T) { func TestIdGroups_NonexistentUser(t *testing.T) {
if _, err := exec.LookPath("id"); err != nil { if _, err := exec.LookPath("id"); err != nil {
t.Skip("id not available on this system") t.Skip("id not available on this system")
} }
_, err := runIdGroups("nonexistent_user_xyzzy_12345") _, err := idGroups("nonexistent_user_xyzzy_12345")
assert.Error(t, err) assert.Error(t, err)
} }
func TestRunIdGroups_InvalidInput(t *testing.T) { func TestIdGroups_InvalidInput(t *testing.T) {
_, err := runIdGroups("") _, err := idGroups("")
assert.Error(t, err) assert.Error(t, err)
_, err = runIdGroups("user\x00name") _, err = idGroups("user\x00name")
assert.Error(t, err) assert.Error(t, err)
} }
@@ -286,7 +412,7 @@ func TestGetentResultsMatchStdlib(t *testing.T) {
current, err := user.Current() current, err := user.Current()
require.NoError(t, err) require.NoError(t, err)
getentUser, _, err := runGetent(current.Username) getentUser, _, err := passwdLookup(current.Username)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, current.Username, getentUser.Username, "username should match") assert.Equal(t, current.Username, getentUser.Username, "username should match")
@@ -303,7 +429,7 @@ func TestGetentResultsMatchStdlib_ByUID(t *testing.T) {
current, err := user.Current() current, err := user.Current()
require.NoError(t, err) require.NoError(t, err)
getentUser, _, err := runGetent(current.Uid) getentUser, _, err := passwdLookup(current.Uid)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, current.Username, getentUser.Username, "username should match when looked up by UID") assert.Equal(t, current.Username, getentUser.Username, "username should match when looked up by UID")
@@ -323,12 +449,12 @@ func TestIdGroupsMatchStdlib(t *testing.T) {
t.Skip("os/user.GroupIds() not working, likely CGO_ENABLED=0") t.Skip("os/user.GroupIds() not working, likely CGO_ENABLED=0")
} }
idGroups, err := runIdGroups(current.Username) idGroupIDs, err := idGroups(current.Username)
require.NoError(t, err) require.NoError(t, err)
// Deduplicate both lists: id -G can return duplicates (e.g., root in Docker) // Deduplicate both lists: id -G can return duplicates (e.g., root in Docker)
// and ElementsMatch treats duplicates as distinct. // and ElementsMatch treats duplicates as distinct.
assert.ElementsMatch(t, uniqueStrings(stdGroups), uniqueStrings(idGroups), "id -G should return same groups as os/user") assert.ElementsMatch(t, uniqueStrings(stdGroups), uniqueStrings(idGroupIDs), "id -G should return same groups as os/user")
} }
func uniqueStrings(ss []string) []string { func uniqueStrings(ss []string) []string {
@@ -343,71 +469,3 @@ func uniqueStrings(ss []string) []string {
} }
return out return out
} }
// TestGetShellFromPasswd_CurrentUser verifies that getShellFromPasswd correctly
// reads the current user's shell from /etc/passwd by comparing it against what
// getent reports (which goes through NSS).
func TestGetShellFromPasswd_CurrentUser(t *testing.T) {
current, err := user.Current()
require.NoError(t, err)
shell := getShellFromPasswd(current.Uid)
if shell == "" {
t.Skip("current user not found in /etc/passwd (may be an NSS-only user)")
}
assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell)
if _, err := exec.LookPath("getent"); err == nil {
_, getentShell, getentErr := runGetent(current.Uid)
if getentErr == nil && getentShell != "" {
assert.Equal(t, getentShell, shell, "shell from /etc/passwd should match getent")
}
}
}
// TestGetShellFromPasswd_RootUser verifies that getShellFromPasswd can read
// root's shell from /etc/passwd. Root is guaranteed to be in /etc/passwd on
// any standard Unix system.
func TestGetShellFromPasswd_RootUser(t *testing.T) {
shell := getShellFromPasswd("0")
require.NotEmpty(t, shell, "root (UID 0) must be in /etc/passwd")
assert.True(t, shell[0] == '/', "root shell should be an absolute path, got %q", shell)
}
// TestGetShellFromPasswd_NonexistentUID verifies that getShellFromPasswd
// returns empty for a UID that doesn't exist in /etc/passwd.
func TestGetShellFromPasswd_NonexistentUID(t *testing.T) {
shell := getShellFromPasswd("4294967294")
assert.Empty(t, shell, "nonexistent UID should return empty shell")
}
// TestGetShellFromPasswd_MatchesGetentForKnownUsers reads /etc/passwd directly
// and cross-validates every entry against getent to ensure parseGetentPasswd
// and getShellFromPasswd agree on shell values.
func TestGetShellFromPasswd_MatchesGetentForKnownUsers(t *testing.T) {
if _, err := exec.LookPath("getent"); err != nil {
t.Skip("getent not available")
}
// Pick a few well-known system UIDs that are virtually always in /etc/passwd.
uids := []string{"0"} // root
current, err := user.Current()
require.NoError(t, err)
uids = append(uids, current.Uid)
for _, uid := range uids {
passwdShell := getShellFromPasswd(uid)
if passwdShell == "" {
continue
}
_, getentShell, err := runGetent(uid)
if err != nil {
continue
}
assert.Equal(t, getentShell, passwdShell, "shell mismatch for UID %s", uid)
}
}
+36
View File
@@ -0,0 +1,36 @@
//go:build windows
package getent
import (
"errors"
"os/user"
)
// Windows does not use NSS or getent; os/user resolves accounts there
// without cgo, so everything delegates to it.
// LookupUser looks up a user by name.
func LookupUser(username string) (*user.User, error) {
return user.Lookup(username)
}
// LookupUserID looks up a user by UID.
func LookupUserID(uid string) (*user.User, error) {
return user.LookupId(uid)
}
// CurrentUser returns the user this process runs as.
func CurrentUser() (*user.User, error) {
return user.Current()
}
// GroupIDs returns the IDs of the groups the user is a member of.
func GroupIDs(u *user.User) ([]string, error) {
return u.GroupIds()
}
// UserShell is unanswerable on Windows, which has no login-shell database.
func UserShell(string) (string, error) {
return "", errors.ErrUnsupported
}
+16
View File
@@ -91,6 +91,12 @@ func SelfDelegatesTo() (Identity, bool) {
return selfIdentity, true return selfIdentity, true
} }
// The values PrivilegedActorKey returns.
const (
ActorKeyAdministrator = "administrator"
ActorKeyRoot = "root"
)
// PrivilegedActor names the principal a privileged operation requires, for use // PrivilegedActor names the principal a privileged operation requires, for use
// in messages shown to the user. // in messages shown to the user.
func PrivilegedActor() string { func PrivilegedActor() string {
@@ -100,6 +106,16 @@ func PrivilegedActor() string {
return "root" return "root"
} }
// PrivilegedActorKey identifies that principal without wording it, for a client
// that writes its own message in the user's language. The words PrivilegedActor
// returns are English, and a translated sentence cannot borrow them.
func PrivilegedActorKey() string {
if runtime.GOOS == "windows" {
return ActorKeyAdministrator
}
return ActorKeyRoot
}
// ElevatedCommand renders a command so that running it grants the privileges the // ElevatedCommand renders a command so that running it grants the privileges the
// operation needs. Windows has no in-line equivalent of sudo, so the command is // operation needs. Windows has no in-line equivalent of sudo, so the command is
// returned unchanged and the user is expected to run it from an elevated // returned unchanged and the user is expected to run it from an elevated
+9 -5
View File
@@ -26,6 +26,7 @@ import (
"github.com/netbirdio/netbird/client/internal/portforward" "github.com/netbirdio/netbird/client/internal/portforward"
"github.com/netbirdio/netbird/client/internal/rosenpass" "github.com/netbirdio/netbird/client/internal/rosenpass"
"github.com/netbirdio/netbird/client/internal/stdnet" "github.com/netbirdio/netbird/client/internal/stdnet"
"github.com/netbirdio/netbird/client/netevents"
"github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/route"
relayClient "github.com/netbirdio/netbird/shared/relay/client" relayClient "github.com/netbirdio/netbird/shared/relay/client"
) )
@@ -93,6 +94,10 @@ type ConnConfig struct {
// ICEConfig ICE protocol configuration // ICEConfig ICE protocol configuration
ICEConfig icemaker.Config ICEConfig icemaker.Config
// NetMgr gates the reconnection guard on OS-reported network
// availability; nil disables gating.
NetMgr *netevents.Manager
} }
type Conn struct { type Conn struct {
@@ -254,7 +259,7 @@ func (conn *Conn) open(engineCtx context.Context, firstPacket []byte) error {
conn.handshaker.AddICEListener(conn.workerICE.OnNewOffer) conn.handshaker.AddICEListener(conn.workerICE.OnNewOffer)
} }
conn.guard = guard.NewGuard(conn.Log, conn.isConnectedOnAllWay, conn.config.Timeout, conn.srWatcher) conn.guard = guard.NewGuard(conn.Log, conn.isConnectedOnAllWay, conn.config.Timeout, conn.srWatcher, conn.config.NetMgr)
conn.wg.Add(1) conn.wg.Add(1)
go func() { go func() {
@@ -440,7 +445,7 @@ func (conn *Conn) onICEConnectionIsReady(priority conntype.ConnPriority, iceConn
conn.dumpState.NewLocalProxy() conn.dumpState.NewLocalProxy()
wgProxy, err = conn.newProxy(iceConnInfo.RemoteConn) wgProxy, err = conn.newProxy(iceConnInfo.RemoteConn)
if err != nil { if err != nil {
conn.Log.Errorf("failed to add turn net.Conn to local proxy: %v", err) conn.Log.Errorf("failed to add relayed net.Conn to local proxy: %v", err)
return return
} }
ep = wgProxy.EndpointAddr() ep = wgProxy.EndpointAddr()
@@ -878,9 +883,8 @@ func (conn *Conn) newProxy(remoteConn net.Conn) (wgproxy.Proxy, error) {
} }
wgProxy := conn.config.WgConfig.WgInterface.GetProxy() wgProxy := conn.config.WgConfig.WgInterface.GetProxy()
if err := wgProxy.AddTurnConn(conn.ctx, udpAddr, remoteConn); err != nil { if err := wgProxy.AddRelayedConn(conn.ctx, udpAddr, remoteConn); err != nil {
conn.Log.Errorf("failed to add turn net.Conn to local proxy: %v", err) return nil, fmt.Errorf("add relayed conn to proxy: %w", err)
return nil, err
} }
return wgProxy, nil return wgProxy, nil
} }
+44 -5
View File
@@ -22,6 +22,12 @@ const (
type connStatusFunc func() ConnStatus type connStatusFunc func() ConnStatus
// NetworkWatcher is the availability view the guard gates reconnects on.
type NetworkWatcher interface {
IsOnline() bool
Changed() <-chan struct{}
}
// Guard is responsible for the reconnection logic. // Guard is responsible for the reconnection logic.
// It will trigger to send an offer to the peer then has connection issues. // It will trigger to send an offer to the peer then has connection issues.
// Watch these events: // Watch these events:
@@ -31,20 +37,26 @@ type connStatusFunc func() ConnStatus
// - Relayed connection disconnected // - Relayed connection disconnected
// - ICE candidate changes // - ICE candidate changes
type Guard struct { type Guard struct {
log *log.Entry log *log.Entry
isConnectedOnAllWay connStatusFunc isConnectedOnAllWay connStatusFunc
timeout time.Duration timeout time.Duration
srWatcher *SRWatcher srWatcher *SRWatcher
// netWatcher gates reconnect attempts on OS-reported network availability;
// nil disables gating.
netWatcher NetworkWatcher
relayedConnDisconnected chan struct{} relayedConnDisconnected chan struct{}
iCEConnDisconnected chan struct{} iCEConnDisconnected chan struct{}
} }
func NewGuard(log *log.Entry, isConnectedFn connStatusFunc, timeout time.Duration, srWatcher *SRWatcher) *Guard { // NewGuard creates a reconnection guard for a peer connection. A nil netWatcher
// disables network availability gating.
func NewGuard(log *log.Entry, isConnectedFn connStatusFunc, timeout time.Duration, srWatcher *SRWatcher, netWatcher NetworkWatcher) *Guard {
return &Guard{ return &Guard{
log: log, log: log,
isConnectedOnAllWay: isConnectedFn, isConnectedOnAllWay: isConnectedFn,
timeout: timeout, timeout: timeout,
srWatcher: srWatcher, srWatcher: srWatcher,
netWatcher: netWatcher,
relayedConnDisconnected: make(chan struct{}, 1), relayedConnDisconnected: make(chan struct{}, 1),
iCEConnDisconnected: make(chan struct{}, 1), iCEConnDisconnected: make(chan struct{}, 1),
} }
@@ -96,9 +108,19 @@ func (g *Guard) reconnectLoopWithRetry(ctx context.Context, callback func()) {
iceState := &iceRetryState{log: g.log} iceState := &iceRetryState{log: g.log}
defer iceState.reset() defer iceState.reset()
var netChanged <-chan struct{}
if g.netWatcher != nil {
netChanged = g.netWatcher.Changed()
}
for { for {
select { select {
case <-tickerChannel: case <-tickerChannel:
// skip attempts while the OS reports no usable network; the
// netChanged case below resumes the loop once it returns
if g.netWatcher != nil && !g.netWatcher.IsOnline() {
continue
}
switch g.isConnectedOnAllWay() { switch g.isConnectedOnAllWay() {
case ConnStatusConnected: case ConnStatusConnected:
// all good, nothing to do // all good, nothing to do
@@ -135,6 +157,23 @@ func (g *Guard) reconnectLoopWithRetry(ctx context.Context, callback func()) {
tickerChannel = ticker.C tickerChannel = ticker.C
iceState.reset() iceState.reset()
case <-netChanged:
// Re-arm for the next transition before acting on this one.
netChanged = g.netWatcher.Changed()
if !g.netWatcher.IsOnline() {
continue
}
// Ticks skipped while offline drove the backoff towards its
// maximum without ever attempting, and left the ICE budget
// frozen — possibly in hourly mode. Recover on our own so the
// peer does not depend on a signal or relay event that never
// comes when both stayed up across the outage.
g.log.Debugf("network is back, reset reconnection ticker")
ticker.Stop()
ticker = g.newReconnectTicker(ctx)
tickerChannel = ticker.C
iceState.reset()
case <-ctx.Done(): case <-ctx.Done():
g.log.Debugf("context is done, stop reconnect loop") g.log.Debugf("context is done, stop reconnect loop")
return return
@@ -15,7 +15,7 @@ import (
func newTestGuard(status connStatusFunc) *Guard { func newTestGuard(status connStatusFunc) *Guard {
srw := NewSRWatcher(nil, nil, nil, ice.Config{}) srw := NewSRWatcher(nil, nil, nil, ice.Config{})
return NewGuard(log.WithField("test", "guard"), status, 50*time.Millisecond, srw) return NewGuard(log.WithField("test", "guard"), status, 50*time.Millisecond, srw, nil)
} }
// countBackoffTickerGoroutines returns how many goroutines are currently sitting // countBackoffTickerGoroutines returns how many goroutines are currently sitting
@@ -0,0 +1,107 @@
package guard
import (
"context"
"sync/atomic"
"testing"
"time"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/peer/ice"
"github.com/netbirdio/netbird/client/netevents/netstate"
)
// newTestGuardWithNetState builds a guard with a realistic MaxInterval: the
// backoff must be able to grow well past the outage, as it does in production
// where the timeout is seconds to minutes.
func newTestGuardWithNetState(status connStatusFunc, netState *netstate.State) *Guard {
srw := NewSRWatcher(nil, nil, nil, ice.Config{})
return NewGuard(log.WithField("test", "guard"), status, 30*time.Second, srw, netState)
}
// TestGuard_RecoversAfterOfflineToOnline covers a peer that stays disconnected
// across a network outage while neither signal nor relay reports an event —
// both stayed up, as on a short airplane mode toggle over Wi-Fi.
//
// Every tick taken while offline is skipped, but it still advances the
// exponential backoff, so by the time the network returns the next tick can be
// tens of seconds away. Without an explicit reaction to the transition the
// peer waits out that interval for a recovery that could start immediately.
func TestGuard_RecoversAfterOfflineToOnline(t *testing.T) {
netState := netstate.New()
var attempts atomic.Int32
g := newTestGuardWithNetState(func() ConnStatus { return ConnStatusDisconnected }, netState)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
// Start from the reconnect ticker (800ms initial interval), the state a
// peer is in after it loses its connection.
go g.Start(ctx, func() { attempts.Add(1) })
g.SetRelayedConnDisconnected()
// Let the backoff climb: 0.8s, 1.6s, 3.2s, 6.4s ... every tick is skipped
// while offline, but each one doubles the wait for the next.
netState.Set(false)
time.Sleep(8 * time.Second)
offlineAttempts := attempts.Load()
if offlineAttempts != 0 {
t.Fatalf("callback ran %d times while offline, want 0", offlineAttempts)
}
netState.Set(true)
// The next organic tick is now several seconds out, so anything within
// this window can only come from reacting to the transition itself.
pollCtx, stopPolling := context.WithTimeout(ctx, 2*time.Second)
defer stopPolling()
select {
case <-pollCtx.Done():
t.Fatal("peer was not retried within 2s of the network coming back, " +
"with neither a signal nor a relay event to fall back on")
case <-pollUntil(pollCtx, func() bool { return attempts.Load() > 0 }):
}
}
// TestGuard_OfflineTransitionDoesNotRetry checks the other direction: going
// offline must not itself trigger an attempt.
func TestGuard_OfflineTransitionDoesNotRetry(t *testing.T) {
netState := netstate.New()
var attempts atomic.Int32
g := newTestGuardWithNetState(func() ConnStatus { return ConnStatusDisconnected }, netState)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go g.Start(ctx, func() { attempts.Add(1) })
netState.Set(false)
time.Sleep(5 * time.Second)
if got := attempts.Load(); got != 0 {
t.Fatalf("callback ran %d times after going offline, want 0", got)
}
}
// pollUntil closes the returned channel once cond holds. It gives up when ctx
// is done, so the polling goroutine never outlives the test that started it.
func pollUntil(ctx context.Context, cond func() bool) <-chan struct{} {
done := make(chan struct{})
go func() {
for {
if cond() {
close(done)
return
}
select {
case <-ctx.Done():
return
case <-time.After(10 * time.Millisecond):
}
}
}()
return done
}
+38 -24
View File
@@ -81,14 +81,19 @@ type Handshaker struct {
func NewHandshaker(log *log.Entry, config ConnConfig, signaler *Signaler, ice *WorkerICE, relay *WorkerRelay, metricsStages *MetricsStages) *Handshaker { func NewHandshaker(log *log.Entry, config ConnConfig, signaler *Signaler, ice *WorkerICE, relay *WorkerRelay, metricsStages *MetricsStages) *Handshaker {
h := &Handshaker{ h := &Handshaker{
log: log, log: log,
config: config, config: config,
signaler: signaler, signaler: signaler,
ice: ice, ice: ice,
relay: relay, relay: relay,
metricsStages: metricsStages, metricsStages: metricsStages,
remoteOffersCh: make(chan OfferAnswer), // Buffered by one so an offer or answer that arrives between Open launching
remoteAnswerCh: make(chan OfferAnswer), // the Listen goroutine and it reaching its receive is held rather than
// dropped. A peer activated by an incoming signal receives the remote's
// message in that window; an unbuffered channel skips it as "receiver not
// ready", and the connection cannot proceed until the remote re-sends.
remoteOffersCh: make(chan OfferAnswer, 1),
remoteAnswerCh: make(chan OfferAnswer, 1),
} }
// assume remote supports ICE until we learn otherwise from received offers // assume remote supports ICE until we learn otherwise from received offers
h.remoteICESupported.Store(ice != nil) h.remoteICESupported.Store(ice != nil)
@@ -162,29 +167,38 @@ func (h *Handshaker) SendOffer() error {
return h.sendOffer() return h.sendOffer()
} }
// OnRemoteOffer handles an offer from the remote peer and returns true if the message was accepted, false otherwise // OnRemoteOffer hands an offer to Listen without blocking, keeping only the most
// doesn't block, discards the message if connection wasn't ready // recent one if several arrive before Listen reads them.
func (h *Handshaker) OnRemoteOffer(offer OfferAnswer) { func (h *Handshaker) OnRemoteOffer(offer OfferAnswer) {
select { enqueueLatest(h.remoteOffersCh, offer)
case h.remoteOffersCh <- offer:
return
default:
h.log.Warnf("skipping remote offer message because receiver not ready")
// connection might not be ready yet to receive so we ignore the message
return
}
} }
// OnRemoteAnswer handles an offer from the remote peer and returns true if the message was accepted, false otherwise // OnRemoteAnswer hands an answer to Listen without blocking, keeping only the most
// doesn't block, discards the message if connection wasn't ready // recent one if several arrive before Listen reads them.
func (h *Handshaker) OnRemoteAnswer(answer OfferAnswer) { func (h *Handshaker) OnRemoteAnswer(answer OfferAnswer) {
enqueueLatest(h.remoteAnswerCh, answer)
}
// enqueueLatest delivers msg on a one-slot channel without blocking. When the slot
// already holds an unread message the older one is discarded in favor of msg, so a
// message arriving before Listen starts reading is held rather than dropped, and
// the newest wins if several arrive first. Safe because there is a single producer
// (the engine loop): after draining the stale value the send always has room.
func enqueueLatest(ch chan OfferAnswer, msg OfferAnswer) {
select { select {
case h.remoteAnswerCh <- answer: case ch <- msg:
return return
default: default:
// connection might not be ready yet to receive so we ignore the message }
h.log.Warnf("skipping remote answer message because receiver not ready")
return select {
case <-ch:
default:
}
select {
case ch <- msg:
default:
} }
} }
+63
View File
@@ -0,0 +1,63 @@
package peer
import (
"testing"
"time"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
)
func newTestHandshaker(t *testing.T) *Handshaker {
t.Helper()
// The tests exercise the answer path, whose Listen branch dispatches to the
// relay listener without sending an answer, so no signaler/ICE/relay is needed.
return NewHandshaker(log.WithField("test", t.Name()), ConnConfig{}, nil, nil, nil, nil)
}
// TestHandshakerHoldsSignalArrivingBeforeListen covers the case where a peer is
// activated by an incoming signal: the remote's offer/answer arrives in the same
// step that opens the connection, before the Listen loop starts reading. The
// message must be held rather than dropped, or the connection cannot proceed until
// the remote re-sends. This is the path taken when an eager peer connects to a
// lazily-managed one.
func TestHandshakerHoldsSignalArrivingBeforeListen(t *testing.T) {
h := newTestHandshaker(t)
processed := make(chan *OfferAnswer, 4)
h.AddRelayListener(func(o *OfferAnswer) { processed <- o })
// Delivered before Listen is reading, as when the peer is woken by the remote's
// signal and the message is delivered right after Open.
h.OnRemoteAnswer(OfferAnswer{WgListenPort: 51820})
go h.Listen(t.Context())
select {
case <-processed:
case <-time.After(2 * time.Second):
assert.Fail(t, "remote-answer dispatch: signal delivered before Listen was ready was dropped")
}
}
// TestHandshakerKeepsLatestSignalBeforeListen covers several signals arriving
// before Listen reads: the newest must win (matching the latest-offer contract),
// rather than the first being kept and later ones discarded.
func TestHandshakerKeepsLatestSignalBeforeListen(t *testing.T) {
h := newTestHandshaker(t)
processed := make(chan *OfferAnswer, 4)
h.AddRelayListener(func(o *OfferAnswer) { processed <- o })
h.OnRemoteAnswer(OfferAnswer{WgListenPort: 1111})
h.OnRemoteAnswer(OfferAnswer{WgListenPort: 2222})
go h.Listen(t.Context())
select {
case got := <-processed:
assert.Equal(t, 2222, got.WgListenPort, "remote-answer dispatch: the latest queued signal should be processed")
case <-time.After(2 * time.Second):
assert.Fail(t, "remote-answer dispatch: queued signal was dropped")
}
}
+29
View File
@@ -1,11 +1,40 @@
package peer package peer
// ClientState identifies the client connection state delivered via
// Listener.OnStateChanged.
type ClientState int
// Client states. The numeric values cross the gomobile boundary (the mobile
// bindings re-export them as integer constants), so they are a wire format:
// append new states at the end, never reorder or insert.
const (
ClientStateDisconnected ClientState = iota
ClientStateConnected
ClientStateConnecting
ClientStateDisconnecting
// ClientStateNoNetwork is an overlay state: it is never stored as the
// last notification, only derived from ClientStateConnecting while the
// OS reports no usable network (see notifier.effectiveState).
ClientStateNoNetwork
)
// Listener is a callback type about the NetBird network connection state // Listener is a callback type about the NetBird network connection state
type Listener interface { type Listener interface {
// OnStateChanged reports every client state transition. New states are
// delivered only through this callback; the per-state callbacks below
// are kept for compatibility and will be removed once all consumers
// have migrated.
OnStateChanged(state ClientState)
// Deprecated: consume OnStateChanged instead.
OnConnected() OnConnected()
// Deprecated: consume OnStateChanged instead.
OnDisconnected() OnDisconnected()
// Deprecated: consume OnStateChanged instead.
OnConnecting() OnConnecting()
// Deprecated: consume OnStateChanged instead.
OnDisconnecting() OnDisconnecting()
OnAddressChanged(string, string) OnAddressChanged(string, string)
OnPeersListChanged(int) OnPeersListChanged(int)
} }
+81 -30
View File
@@ -4,31 +4,64 @@ import (
"sync" "sync"
) )
const (
stateDisconnected = iota
stateConnected
stateConnecting
stateDisconnecting
)
type notifier struct { type notifier struct {
// publishLock orders state publication: it is held across computing the
// effective state and handing it to the listener, so a transition cannot
// overtake a newer one and leave the listener on a stale state.
publishLock sync.Mutex
serverStateLock sync.Mutex serverStateLock sync.Mutex
listenersLock sync.Mutex listenersLock sync.Mutex
listener Listener listener Listener
currentClientState bool currentClientState bool
lastNotification int lastNotification ClientState
lastNumberOfPeers int lastNumberOfPeers int
lastFqdnAddress string lastFqdnAddress string
lastIPAddress string lastIPAddress string
networkAvailable bool
} }
func newNotifier() *notifier { func newNotifier() *notifier {
return &notifier{} return &notifier{
networkAvailable: true,
}
}
// effectiveState maps the computed state to what listeners should see:
// while the OS reports no usable network, "Connecting" would be a lie —
// connection attempts are suspended — so it is reported as NoNetwork.
// Caller must hold serverStateLock.
func (n *notifier) effectiveState(state ClientState) ClientState {
if !n.networkAvailable && state == ClientStateConnecting {
return ClientStateNoNetwork
}
return state
}
// setNetworkAvailable records the OS network availability and re-notifies
// the listener when the flag flips the effective state (Connecting <->
// NoNetwork).
func (n *notifier) setNetworkAvailable(available bool) {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock()
if n.networkAvailable == available {
n.serverStateLock.Unlock()
return
}
previous := n.effectiveState(n.lastNotification)
n.networkAvailable = available
current := n.effectiveState(n.lastNotification)
n.serverStateLock.Unlock()
if previous != current {
n.notify(current)
}
} }
func (n *notifier) setListener(listener Listener) { func (n *notifier) setListener(listener Listener) {
n.serverStateLock.Lock() n.serverStateLock.Lock()
lastNotification := n.lastNotification lastNotification := n.effectiveState(n.lastNotification)
numOfPeers := n.lastNumberOfPeers numOfPeers := n.lastNumberOfPeers
fqdnAddress := n.lastFqdnAddress fqdnAddress := n.lastFqdnAddress
address := n.lastIPAddress address := n.lastIPAddress
@@ -52,6 +85,9 @@ func (n *notifier) removeListener() {
} }
func (n *notifier) updateServerStates(mgmState bool, signalState bool) { func (n *notifier) updateServerStates(mgmState bool, signalState bool) {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock() n.serverStateLock.Lock()
calculatedState := n.calculateState(mgmState, signalState) calculatedState := n.calculateState(mgmState, signalState)
@@ -61,43 +97,54 @@ func (n *notifier) updateServerStates(mgmState bool, signalState bool) {
} }
n.lastNotification = calculatedState n.lastNotification = calculatedState
effective := n.effectiveState(calculatedState)
n.serverStateLock.Unlock() n.serverStateLock.Unlock()
n.notify(calculatedState) n.notify(effective)
} }
func (n *notifier) clientStart() { func (n *notifier) clientStart() {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock() n.serverStateLock.Lock()
n.currentClientState = true n.currentClientState = true
n.lastNotification = stateConnecting n.lastNotification = ClientStateConnecting
effective := n.effectiveState(ClientStateConnecting)
n.serverStateLock.Unlock() n.serverStateLock.Unlock()
n.notify(stateConnecting) n.notify(effective)
} }
func (n *notifier) clientStop() { func (n *notifier) clientStop() {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock() n.serverStateLock.Lock()
n.currentClientState = false n.currentClientState = false
n.lastNotification = stateDisconnected n.lastNotification = ClientStateDisconnected
n.serverStateLock.Unlock() n.serverStateLock.Unlock()
n.notify(stateDisconnected) n.notify(ClientStateDisconnected)
} }
func (n *notifier) clientTearDown() { func (n *notifier) clientTearDown() {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock() n.serverStateLock.Lock()
n.currentClientState = false n.currentClientState = false
n.lastNotification = stateDisconnecting n.lastNotification = ClientStateDisconnecting
n.serverStateLock.Unlock() n.serverStateLock.Unlock()
n.notify(stateDisconnecting) n.notify(ClientStateDisconnecting)
} }
func (n *notifier) isServerStateChanged(newState int) bool { func (n *notifier) isServerStateChanged(newState ClientState) bool {
return n.lastNotification != newState return n.lastNotification != newState
} }
func (n *notifier) notify(state int) { func (n *notifier) notify(state ClientState) {
n.listenersLock.Lock() n.listenersLock.Lock()
listener := n.listener listener := n.listener
n.listenersLock.Unlock() n.listenersLock.Unlock()
@@ -109,20 +156,20 @@ func (n *notifier) notify(state int) {
notifyListener(listener, state) notifyListener(listener, state)
} }
func (n *notifier) calculateState(managementConn, signalConn bool) int { func (n *notifier) calculateState(managementConn, signalConn bool) ClientState {
if managementConn && signalConn { if managementConn && signalConn {
return stateConnected return ClientStateConnected
} }
if !managementConn && !signalConn && !n.currentClientState { if !managementConn && !signalConn && !n.currentClientState {
return stateDisconnected return ClientStateDisconnected
} }
if n.lastNotification == stateDisconnecting { if n.lastNotification == ClientStateDisconnecting {
return stateDisconnecting return ClientStateDisconnecting
} }
return stateConnecting return ClientStateConnecting
} }
func (n *notifier) peerListChanged(numOfPeers int) { func (n *notifier) peerListChanged(numOfPeers int) {
@@ -159,15 +206,19 @@ func (n *notifier) localAddressChanged(fqdn, address string) {
listener.OnAddressChanged(fqdn, address) listener.OnAddressChanged(fqdn, address)
} }
func notifyListener(l Listener, state int) { func notifyListener(l Listener, state ClientState) {
// legacy per-state callbacks; NoNetwork is delivered only via
// OnStateChanged below
switch state { switch state {
case stateDisconnected: case ClientStateDisconnected:
l.OnDisconnected() l.OnDisconnected()
case stateConnected: case ClientStateConnected:
l.OnConnected() l.OnConnected()
case stateConnecting: case ClientStateConnecting:
l.OnConnecting() l.OnConnecting()
case stateDisconnecting: case ClientStateDisconnecting:
l.OnDisconnecting() l.OnDisconnecting()
} }
l.OnStateChanged(state)
} }
@@ -0,0 +1,108 @@
package peer
import (
"sync"
"testing"
"time"
)
type recordingListener struct {
mu sync.Mutex
states []ClientState
onState func(ClientState)
}
func (l *recordingListener) OnStateChanged(state ClientState) {
l.mu.Lock()
l.states = append(l.states, state)
hook := l.onState
l.mu.Unlock()
if hook != nil {
hook(state)
}
}
func (l *recordingListener) last() (ClientState, bool) {
l.mu.Lock()
defer l.mu.Unlock()
if len(l.states) == 0 {
return 0, false
}
return l.states[len(l.states)-1], true
}
func (l *recordingListener) snapshot() []ClientState {
l.mu.Lock()
defer l.mu.Unlock()
return append([]ClientState(nil), l.states...)
}
func (l *recordingListener) OnConnected() {}
func (l *recordingListener) OnDisconnected() {}
func (l *recordingListener) OnConnecting() {}
func (l *recordingListener) OnDisconnecting() {}
func (l *recordingListener) OnAddressChanged(string, string) {}
func (l *recordingListener) OnPeersListChanged(int) {}
// TestNotifier_ConcurrentAvailabilityFlipOrdersPublication holds the first
// transition inside the listener callback and flips availability again from
// another goroutine while it is parked. The second flip must not publish
// ahead of the one in flight, otherwise the listener ends up on a state the
// notifier already superseded.
func TestNotifier_ConcurrentAvailabilityFlipOrdersPublication(t *testing.T) {
n := newNotifier()
n.currentClientState = true
n.lastNotification = ClientStateConnecting
entered := make(chan struct{})
release := make(chan struct{})
l := &recordingListener{}
l.onState = func(state ClientState) {
if state != ClientStateNoNetwork {
return
}
l.mu.Lock()
l.onState = nil
l.mu.Unlock()
close(entered)
<-release
}
n.listener = l
var wg sync.WaitGroup
wg.Add(1)
go func() {
defer wg.Done()
n.setNetworkAvailable(false)
}()
<-entered
flipped := make(chan struct{})
go func() {
defer close(flipped)
n.setNetworkAvailable(true)
}()
select {
case <-flipped:
t.Fatal("the online transition published while the offline one was " +
"still in flight; publication is not serialized")
case <-time.After(200 * time.Millisecond):
}
close(release)
<-flipped
wg.Wait()
got, ok := l.last()
if !ok {
t.Fatal("listener never observed a state")
}
if got != ClientStateConnecting {
t.Fatalf("listener holds %v after the network came back, want Connecting; sequence: %v",
got, l.snapshot())
}
}
+15 -12
View File
@@ -6,29 +6,32 @@ import (
) )
type mocListener struct { type mocListener struct {
lastState int lastState ClientState
wg sync.WaitGroup wg sync.WaitGroup
peersWg sync.WaitGroup peersWg sync.WaitGroup
peers int peers int
} }
func (l *mocListener) OnConnected() { func (l *mocListener) OnConnected() {
l.lastState = stateConnected l.lastState = ClientStateConnected
l.wg.Done() l.wg.Done()
} }
func (l *mocListener) OnDisconnected() { func (l *mocListener) OnDisconnected() {
l.lastState = stateDisconnected l.lastState = ClientStateDisconnected
l.wg.Done() l.wg.Done()
} }
func (l *mocListener) OnConnecting() { func (l *mocListener) OnConnecting() {
l.lastState = stateConnecting l.lastState = ClientStateConnecting
l.wg.Done() l.wg.Done()
} }
func (l *mocListener) OnDisconnecting() { func (l *mocListener) OnDisconnecting() {
l.lastState = stateDisconnecting l.lastState = ClientStateDisconnecting
l.wg.Done() l.wg.Done()
} }
func (l *mocListener) OnStateChanged(state ClientState) {
}
func (l *mocListener) OnAddressChanged(host, addr string) { func (l *mocListener) OnAddressChanged(host, addr string) {
} }
@@ -57,15 +60,15 @@ func Test_notifier_serverState(t *testing.T) {
type scenario struct { type scenario struct {
name string name string
expected int expected ClientState
mgmState bool mgmState bool
signalState bool signalState bool
} }
scenarios := []scenario{ scenarios := []scenario{
{"connected", stateConnected, true, true}, {"connected", ClientStateConnected, true, true},
{"mgm down", stateConnecting, false, true}, {"mgm down", ClientStateConnecting, false, true},
{"signal down", stateConnecting, true, false}, {"signal down", ClientStateConnecting, true, false},
{"disconnected", stateDisconnected, false, false}, {"disconnected", ClientStateDisconnected, false, false},
} }
for _, tt := range scenarios { for _, tt := range scenarios {
@@ -85,7 +88,7 @@ func Test_notifier_SetListener(t *testing.T) {
listener.setPeersWaiter() listener.setPeersWaiter()
n := newNotifier() n := newNotifier()
n.lastNotification = stateConnecting n.lastNotification = ClientStateConnecting
n.setListener(listener) n.setListener(listener)
listener.wait() listener.wait()
listener.waitPeers() listener.waitPeers()
@@ -99,7 +102,7 @@ func Test_notifier_RemoveListener(t *testing.T) {
listener.setWaiter() listener.setWaiter()
listener.setPeersWaiter() listener.setPeersWaiter()
n := newNotifier() n := newNotifier()
n.lastNotification = stateConnecting n.lastNotification = ClientStateConnecting
n.setListener(listener) n.setListener(listener)
// setListener replays cached state on a goroutine; wait for both the state // setListener replays cached state on a goroutine; wait for both the state
// and peers callbacks to finish so we don't race on listener.peers. // and peers callbacks to finish so we don't race on listener.peers.
+6
View File
@@ -1211,6 +1211,12 @@ func (d *Status) ClientTeardown() {
d.notifyStateChange() d.notifyStateChange()
} }
// SetNetworkAvailable records the OS-reported network availability; while
// unavailable, listeners see NoNetwork instead of Connecting.
func (d *Status) SetNetworkAvailable(available bool) {
d.notifier.setNetworkAvailable(available)
}
// SetConnectionListener set a listener to the notifier // SetConnectionListener set a listener to the notifier
func (d *Status) SetConnectionListener(listener Listener) { func (d *Status) SetConnectionListener(listener Listener) {
d.notifier.setListener(listener) d.notifier.setListener(listener)
+16 -5
View File
@@ -255,8 +255,8 @@ func (w *WorkerICE) connect(ctx context.Context, agent *icemaker.ThreadSafeAgent
return return
} }
w.log.Debugf("turn agent dial") w.log.Debugf("agent dial")
remoteConn, err := w.turnAgentDial(ctx, agent, remoteOfferAnswer) remoteConn, err := w.agentDial(ctx, agent, remoteOfferAnswer)
if err != nil { if err != nil {
w.log.Debugf("failed to dial the remote peer: %s", err) w.log.Debugf("failed to dial the remote peer: %s", err)
w.closeAgent(agent, w.agentDialerCancel) w.closeAgent(agent, w.agentDialerCancel)
@@ -389,6 +389,17 @@ func (w *WorkerICE) injectPortForwardedCandidate(srflxCandidate ice.Candidate) {
return return
} }
// A forwarded candidate only makes sense for an IPv4 mapping, which
// translates a port on the gateway's address. An IPv6 pinhole translates
// nothing: it unblocks the address ICE already gathers as a host candidate,
// so there is no second address to advertise. Injecting one here would also
// paste an IPv6 address onto whichever server-reflexive candidate arrived
// first, which is usually IPv4.
if mapping.ExternalIP != nil && mapping.ExternalIP.To4() == nil {
w.log.Debugf("skipping port-forwarded candidate: %s mapping is IPv6-only", mapping.NATType)
return
}
w.muxAgent.Lock() w.muxAgent.Lock()
if w.portForwardAttempted { if w.portForwardAttempted {
w.muxAgent.Unlock() w.muxAgent.Unlock()
@@ -517,8 +528,8 @@ func (w *WorkerICE) onConnectionStateChange(agent *icemaker.ThreadSafeAgent, dia
w.logSuccessfulPaths(agent) w.logSuccessfulPaths(agent)
return return
case ice.ConnectionStateFailed, ice.ConnectionStateDisconnected, ice.ConnectionStateClosed: case ice.ConnectionStateFailed, ice.ConnectionStateDisconnected, ice.ConnectionStateClosed:
// ice.ConnectionStateClosed happens when we recreate the agent. For the P2P to TURN switch important to // ice.ConnectionStateClosed happens when we recreate the agent. The P2P to relay switch requires
// notify the conn.onICEStateDisconnected changes to update the current used priority // notifying conn.onICEStateDisconnected so it can update the currently used priority.
sessionChanged := w.closeAgent(agent, dialerCancel) sessionChanged := w.closeAgent(agent, dialerCancel)
@@ -532,7 +543,7 @@ func (w *WorkerICE) onConnectionStateChange(agent *icemaker.ThreadSafeAgent, dia
} }
} }
func (w *WorkerICE) turnAgentDial(ctx context.Context, agent *icemaker.ThreadSafeAgent, remoteOfferAnswer *OfferAnswer) (*ice.Conn, error) { func (w *WorkerICE) agentDial(ctx context.Context, agent *icemaker.ThreadSafeAgent, remoteOfferAnswer *OfferAnswer) (*ice.Conn, error) {
if isController(w.config) { if isController(w.config) {
return agent.Dial(ctx, remoteOfferAnswer.IceCredentials.UFrag, remoteOfferAnswer.IceCredentials.Pwd) return agent.Dial(ctx, remoteOfferAnswer.IceCredentials.UFrag, remoteOfferAnswer.IceCredentials.Pwd)
} else { } else {
+25 -5
View File
@@ -10,10 +10,8 @@ import (
"sync" "sync"
"time" "time"
"github.com/libp2p/go-nat" "github.com/netbirdio/go-nat"
log "github.com/sirupsen/logrus" log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/portforward/pcp"
) )
const ( const (
@@ -168,6 +166,11 @@ func (m *Manager) setup(ctx context.Context) (nat.NAT, *Mapping, error) {
if err != nil { if err != nil {
return nil, nil, fmt.Errorf("create port mapping: %w", err) return nil, nil, fmt.Errorf("create port mapping: %w", err)
} }
// Only meaningful once a mapping has been attempted: that is what opens the
// pinhole and records its outcome.
logIPv6Pinhole(gateway)
return gateway, mapping, nil return gateway, mapping, nil
} }
@@ -265,7 +268,9 @@ func (m *Manager) checkHealthAndRecreate(ctx context.Context, gateway nat.NAT) b
return false return false
} }
pcpNAT, ok := gateway.(*pcp.NAT) // Assert on the interface, not on a concrete type: a dual-stack gateway is
// a wrapper around the IPv4 NAT, so a type assertion misses it.
checker, ok := gateway.(nat.HealthChecker)
if !ok { if !ok {
return false return false
} }
@@ -273,7 +278,7 @@ func (m *Manager) checkHealthAndRecreate(ctx context.Context, gateway nat.NAT) b
ctx, cancel := context.WithTimeout(ctx, 10*time.Second) ctx, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel() defer cancel()
epoch, serverRestarted, err := pcpNAT.CheckServerHealth(ctx) epoch, serverRestarted, err := checker.CheckServerHealth(ctx)
if err != nil { if err != nil {
log.Debugf("PCP health check failed: %v", err) log.Debugf("PCP health check failed: %v", err)
return false return false
@@ -340,3 +345,18 @@ func (m *Manager) startTearDown(ctx context.Context) {
func isPermanentLeaseRequired(err error) bool { func isPermanentLeaseRequired(err error) bool {
return err != nil && upnpErrPermanentLeaseOnly.MatchString(err.Error()) return err != nil && upnpErrPermanentLeaseOnly.MatchString(err.Error())
} }
// logIPv6Pinhole reports the outcome of the IPv6 pinhole. Pinholes are best
// effort and never fail a mapping on their own, so this is the only way to see
// whether one was actually opened.
func logIPv6Pinhole(gateway nat.NAT) {
reporter, ok := gateway.(nat.IPv6PinholeReporter)
if !ok {
return
}
if err := reporter.IPv6PinholeError(); err != nil {
log.Warnf("IPv6 pinhole: %v", err)
return
}
log.Infof("IPv6 pinhole open")
}

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