mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-28 09:39:05 +02:00
Merge branch 'main' into oauth-flow-fallback
This commit is contained in:
@@ -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"
|
||||||
|
|||||||
@@ -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
|
|
||||||
@@ -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."
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
name: UI Translations
|
||||||
|
|
||||||
|
on:
|
||||||
|
pull_request:
|
||||||
|
paths:
|
||||||
|
- "client/ui/i18n/locales/**"
|
||||||
|
- "client/ui/i18n/check-translations.mjs"
|
||||||
|
- ".github/workflows/ui-translations.yml"
|
||||||
|
push:
|
||||||
|
branches:
|
||||||
|
- main
|
||||||
|
paths:
|
||||||
|
- "client/ui/i18n/locales/**"
|
||||||
|
- "client/ui/i18n/check-translations.mjs"
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
concurrency:
|
||||||
|
group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }}
|
||||||
|
cancel-in-progress: true
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
check-translations:
|
||||||
|
name: Check translation key parity
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
timeout-minutes: 5
|
||||||
|
steps:
|
||||||
|
- name: Checkout repository
|
||||||
|
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||||
|
with:
|
||||||
|
persist-credentials: false
|
||||||
|
|
||||||
|
- name: Set up Node.js
|
||||||
|
uses: actions/setup-node@v4
|
||||||
|
with:
|
||||||
|
node-version: "22"
|
||||||
|
|
||||||
|
# English (en) is the source of truth for translation keys; every other
|
||||||
|
# locale declared in _index.json must carry the exact same key set.
|
||||||
|
- name: Check translation key parity
|
||||||
|
run: node client/ui/i18n/check-translations.mjs
|
||||||
@@ -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
@@ -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.
|
||||||
|
|||||||
@@ -130,7 +130,7 @@ In November 2022, NetBird joined the [StartUpSecure program](https://www.forschu
|
|||||||

|

|
||||||
|
|
||||||
### 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
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -26,6 +26,8 @@ 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/netstate"
|
||||||
|
"github.com/netbirdio/netbird/client/netsweep"
|
||||||
"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 +42,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 +82,13 @@ type Client struct {
|
|||||||
deviceName string
|
deviceName string
|
||||||
uiVersion string
|
uiVersion string
|
||||||
networkChangeListener listener.NetworkChangeListener
|
networkChangeListener listener.NetworkChangeListener
|
||||||
|
// netState outlives engine restarts: it mirrors the OS connectivity, not
|
||||||
|
// the engine lifecycle. Run and RunWithoutLogin inject it into each new
|
||||||
|
// ConnectClient, which distributes it to every reconnection loop.
|
||||||
|
netState *netstate.State
|
||||||
|
|
||||||
|
// sweeper also outlives engine restarts; NotifyNetworkChange sweeps it.
|
||||||
|
sweeper *netsweep.Sweeper
|
||||||
|
|
||||||
stateMu sync.RWMutex
|
stateMu sync.RWMutex
|
||||||
connectClient *internal.ConnectClient
|
connectClient *internal.ConnectClient
|
||||||
@@ -148,6 +152,7 @@ func NewClient(androidSDKVersion int, deviceName string, uiVersion string, tunAd
|
|||||||
execWorkaround(androidSDKVersion)
|
execWorkaround(androidSDKVersion)
|
||||||
|
|
||||||
net.SetAndroidProtectSocketFn(tunAdapter.ProtectSocket)
|
net.SetAndroidProtectSocketFn(tunAdapter.ProtectSocket)
|
||||||
|
system.SetIFaceDiscover(iFaceDiscover)
|
||||||
return &Client{
|
return &Client{
|
||||||
deviceName: deviceName,
|
deviceName: deviceName,
|
||||||
uiVersion: uiVersion,
|
uiVersion: uiVersion,
|
||||||
@@ -156,6 +161,8 @@ func NewClient(androidSDKVersion int, deviceName string, uiVersion string, tunAd
|
|||||||
recorder: peer.NewRecorder(""),
|
recorder: peer.NewRecorder(""),
|
||||||
ctxCancelLock: &sync.Mutex{},
|
ctxCancelLock: &sync.Mutex{},
|
||||||
networkChangeListener: networkChangeListener,
|
networkChangeListener: networkChangeListener,
|
||||||
|
netState: netstate.New(),
|
||||||
|
sweeper: netsweep.New(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -196,7 +203,8 @@ 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.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper))
|
||||||
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
|
||||||
// is authenticated again — release the latch Status() reports from. Clear
|
// is authenticated again — release the latch Status() reports from. Clear
|
||||||
@@ -237,7 +245,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.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper))
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -285,6 +294,24 @@ 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.
|
||||||
|
func (c *Client) SetNetworkAvailable(available bool) {
|
||||||
|
c.netState.Set(available)
|
||||||
|
c.recorder.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.sweeper.MarkNetworkChange()
|
||||||
|
log.Infof("network change: connections marked stale")
|
||||||
|
}
|
||||||
|
|
||||||
// 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
|
||||||
@@ -525,7 +552,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
|
||||||
|
|||||||
@@ -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))
|
||||||
|
}
|
||||||
+33
-23
@@ -191,39 +191,49 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// loginHintSetter is implemented by both concrete flows (PKCE and device code)
|
|
||||||
// but absent from the OAuthFlow interface, hence the assertion below — the same
|
|
||||||
// way internal/auth wires it in authenticateWithPKCEFlow.
|
|
||||||
type loginHintSetter interface {
|
|
||||||
SetLoginHint(hint string)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) {
|
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) {
|
||||||
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV)
|
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV, profileLoginHint(a.cfgPath))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
|
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// An empty hint is deliberate, not a fallback: a fresh or logged-out profile
|
return runOAuthFlow(a.ctx, oAuthFlow, urlOpener, nil)
|
||||||
// leaves the choice to the IdP, which is how accounts get switched.
|
}
|
||||||
if a.cfgPath != "" {
|
|
||||||
if hint := readProfileEmail(a.cfgPath); hint != "" {
|
// profileLoginHint returns the stored account email for the profile at cfgPath.
|
||||||
if setter, ok := oAuthFlow.(loginHintSetter); ok {
|
// An empty hint is deliberate, not a fallback: a fresh profile leaves the
|
||||||
setter.SetLoginHint(hint)
|
// choice to the IdP. Switching accounts is done by switching or removing
|
||||||
}
|
// profiles, not by logging out — logout keeps the email.
|
||||||
}
|
func profileLoginHint(cfgPath string) string {
|
||||||
|
if cfgPath == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return readProfileEmail(cfgPath)
|
||||||
|
}
|
||||||
|
|
||||||
|
// runOAuthFlow drives an already acquired OAuth flow to a token: requests the
|
||||||
|
// flow info, presents the verification URL through the opener and waits for
|
||||||
|
// the browser round-trip. Open is called synchronously — it is what marks the
|
||||||
|
// surface as opened on the client side, and a fast token's OnLoginSuccess is
|
||||||
|
// a no-op until it has, so the dismissal would be dropped rather than
|
||||||
|
// delayed. Openers must therefore not block: they post their UI work and
|
||||||
|
// return. onWaiting, when set, runs after the URL is shown, right before the
|
||||||
|
// blocking wait.
|
||||||
|
func runOAuthFlow(ctx context.Context, flow auth.OAuthFlow, urlOpener URLOpener, onWaiting func()) (*auth.TokenInfo, error) {
|
||||||
|
flowInfo, err := flow.RequestAuthInfo(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("request auth info: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO())
|
urlOpener.Open(flowInfo.VerificationURIComplete, flowInfo.UserCode)
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err)
|
if onWaiting != nil {
|
||||||
|
onWaiting()
|
||||||
}
|
}
|
||||||
|
|
||||||
go urlOpener.Open(flowInfo.VerificationURIComplete, flowInfo.UserCode)
|
tokenInfo, err := flow.WaitToken(ctx, flowInfo)
|
||||||
|
|
||||||
tokenInfo, err := oAuthFlow.WaitToken(a.ctx, flowInfo)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("waiting for browser login failed: %v", err)
|
return nil, fmt.Errorf("wait for token: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &tokenInfo, nil
|
return &tokenInfo, nil
|
||||||
|
|||||||
@@ -22,7 +22,8 @@ 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 or was logged out. See profile_state.go.
|
// completed an SSO login. Kept across logouts; cleared when the profile is
|
||||||
|
// removed. See profile_state.go.
|
||||||
Email string
|
Email string
|
||||||
IsActive bool
|
IsActive bool
|
||||||
}
|
}
|
||||||
@@ -200,11 +201,9 @@ func (pm *ProfileManager) LogoutProfile(id string) error {
|
|||||||
return fmt.Errorf("failed to save config: %w", err)
|
return fmt.Errorf("failed to save config: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Not fatal: a stale hint costs an account switch, not the logout itself.
|
// The stored account email is kept on purpose, matching the desktop and CLI
|
||||||
if err := removeProfileEmail(configPath); err != nil {
|
// logout semantics: the next login passes it as the login_hint so the IdP
|
||||||
log.Warnf("failed to clear stored account email for profile %s: %v", id, err)
|
// preselects the account. Removing the profile is what deletes it.
|
||||||
}
|
|
||||||
|
|
||||||
log.Infof("logged out from profile: %s", id)
|
log.Infof("logged out from profile: %s", id)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -224,11 +223,24 @@ func (pm *ProfileManager) RenameProfile(id string, newName string) error {
|
|||||||
|
|
||||||
// RemoveProfile deletes a profile
|
// RemoveProfile deletes a profile
|
||||||
func (pm *ProfileManager) RemoveProfile(id string) error {
|
func (pm *ProfileManager) RemoveProfile(id string) error {
|
||||||
|
configPath, err := pm.getProfileConfigPath(id)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
// Use ServiceManager (removes profile from profiles/ directory)
|
// Use ServiceManager (removes profile from profiles/ directory)
|
||||||
if err := pm.serviceMgr.RemoveProfile(profilemanager.ID(id), androidUsername); err != nil {
|
if err := pm.serviceMgr.RemoveProfile(profilemanager.ID(id), androidUsername); err != nil {
|
||||||
return fmt.Errorf("failed to remove profile: %w", err)
|
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)
|
log.Infof("removed profile: %s", id)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
//go:build android
|
||||||
|
|
||||||
|
package android
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
)
|
||||||
|
|
||||||
|
type prefsStore interface {
|
||||||
|
Get(namespace string, v any) (bool, error)
|
||||||
|
Put(namespace string, v any) error
|
||||||
|
}
|
||||||
|
|
||||||
|
type profilePrefs struct {
|
||||||
|
prefs *profilemanager.Prefs
|
||||||
|
}
|
||||||
|
|
||||||
|
func newProfilePrefs(configDir, profileID string) (*profilePrefs, error) {
|
||||||
|
if configDir == "" || profileID == "" {
|
||||||
|
return nil, fmt.Errorf("profile prefs require a config dir and profile ID")
|
||||||
|
}
|
||||||
|
pm := NewProfileManager(configDir)
|
||||||
|
prefs, err := pm.serviceMgr.ProfilePrefs(profilemanager.ID(profileID), androidUsername)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("resolve profile prefs: %w", err)
|
||||||
|
}
|
||||||
|
return &profilePrefs{prefs: prefs}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *profilePrefs) Get(namespace string, v any) (bool, error) {
|
||||||
|
return p.prefs.Get(namespace, v)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *profilePrefs) Put(namespace string, v any) error {
|
||||||
|
return p.prefs.Put(namespace, v)
|
||||||
|
}
|
||||||
@@ -90,10 +90,10 @@ func writeProfileEmail(configPath string, email string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// removeProfileEmail drops the stored account email. Called on logout: while the
|
// removeProfileEmail drops the stored account email. Called on profile removal,
|
||||||
// email is on disk it goes out as a login_hint, which would steer the next login
|
// not on logout: a logged-out profile keeps its email so the next login passes
|
||||||
// straight back into the account just logged out of. Mirrors the desktop UI's
|
// it as the login_hint, matching the desktop and CLI semantics. Mirrors the
|
||||||
// RemoveProfileState call.
|
// desktop UI's RemoveProfileState call.
|
||||||
func removeProfileEmail(configPath string) error {
|
func removeProfileEmail(configPath string) error {
|
||||||
accountPath, err := profileAccountPathFor(configPath)
|
accountPath, err := profileAccountPathFor(configPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -127,10 +127,10 @@ func TestWriteThenReadProfileEmail(t *testing.T) {
|
|||||||
t.Fatalf("remove: %v", err)
|
t.Fatalf("remove: %v", err)
|
||||||
}
|
}
|
||||||
if got := readProfileEmail(configPath); got != "" {
|
if got := readProfileEmail(configPath); got != "" {
|
||||||
t.Errorf("expected no email after logout, got %q", got)
|
t.Errorf("expected no email after removal, got %q", got)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Logout may run on a never-logged-in profile, so a second remove must pass.
|
// Removal may run on a never-logged-in profile, so a second remove must pass.
|
||||||
if err := removeProfileEmail(configPath); err != nil {
|
if err := removeProfileEmail(configPath); err != nil {
|
||||||
t.Fatalf("second remove should be a no-op: %v", err)
|
t.Fatalf("second remove should be a no-op: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,649 @@
|
|||||||
|
//go:build android
|
||||||
|
|
||||||
|
package android
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
gossh "golang.org/x/crypto/ssh"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
nbssh "github.com/netbirdio/netbird/client/ssh"
|
||||||
|
"github.com/netbirdio/netbird/client/ssh/detection"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
sshDialTimeout = 30 * time.Second
|
||||||
|
sshDetectionTimeout = 5 * time.Second
|
||||||
|
)
|
||||||
|
|
||||||
|
// PasswordRequiredMarker tells Java to prompt for a password and retry. It is
|
||||||
|
// a string because gomobile flattens errors to their message, so a sentinel
|
||||||
|
// value would not survive the binding.
|
||||||
|
const PasswordRequiredMarker = "netbird-ssh-password-required"
|
||||||
|
|
||||||
|
// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation,
|
||||||
|
// retry with TrustHostKey set. The presented fingerprint is appended after the
|
||||||
|
// marker so the prompt can display it and the retry can guard against a key
|
||||||
|
// that changed between the two connects. Only regular (non-NetBird) servers
|
||||||
|
// reach this: NetBird peers verify against the registry.
|
||||||
|
const HostKeyUnknownMarker = "netbird-ssh-hostkey-unknown"
|
||||||
|
|
||||||
|
var (
|
||||||
|
errPasswordRequired = errors.New(PasswordRequiredMarker)
|
||||||
|
errClientClosed = errors.New("ssh client closed")
|
||||||
|
)
|
||||||
|
|
||||||
|
// errHostKeyUnknown carries the presented fingerprint so Connect can build the
|
||||||
|
// marker message the Java side parses.
|
||||||
|
type errHostKeyUnknown struct {
|
||||||
|
fingerprint string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *errHostKeyUnknown) Error() string {
|
||||||
|
return HostKeyUnknownMarker + ":" + e.fingerprint
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSHTerminalListener receives SSH session events. It is implemented in Java.
|
||||||
|
//
|
||||||
|
// All callbacks are invoked from goroutines and may run concurrently with each
|
||||||
|
// other; the implementation must be safe to call from any thread.
|
||||||
|
type SSHTerminalListener interface {
|
||||||
|
OnConnected()
|
||||||
|
OnData(data []byte)
|
||||||
|
OnClose(reason string)
|
||||||
|
OnError(message string)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSHClient is a NetBird-aware SSH client exposed to Java via gomobile.
|
||||||
|
//
|
||||||
|
// It dials through the running NetBird tunnel and runs a standard SSH session
|
||||||
|
// on top with PTY enabled. Host-key verification uses the NetBird-provided
|
||||||
|
// peer SSH host keys, identical to the desktop client.
|
||||||
|
type SSHClient struct {
|
||||||
|
nb *Client
|
||||||
|
mu sync.Mutex
|
||||||
|
listener SSHTerminalListener
|
||||||
|
urlOpener URLOpener
|
||||||
|
|
||||||
|
sshClient *gossh.Client
|
||||||
|
session *gossh.Session
|
||||||
|
stdin io.WriteCloser
|
||||||
|
closed bool
|
||||||
|
|
||||||
|
// gen identifies the current connection attempt. Connect and Close bump it,
|
||||||
|
// so an in-flight dial or a reader left over from a previous connection
|
||||||
|
// finds itself stale and stays silent instead of publishing OnConnected or
|
||||||
|
// OnClose for a connection the caller already abandoned.
|
||||||
|
gen uint64
|
||||||
|
dialCancel context.CancelFunc
|
||||||
|
|
||||||
|
// knownHostsConfigDir and knownHostsProfile locate the TOFU store for
|
||||||
|
// regular SSH servers in the profile's preferences. Java supplies them,
|
||||||
|
// since an overlay IP is a different host under a different profile. Empty
|
||||||
|
// until set: without them a regular server cannot be verified and Connect
|
||||||
|
// refuses one.
|
||||||
|
knownHostsConfigDir string
|
||||||
|
knownHostsProfile string
|
||||||
|
// trustHostKey carries the fingerprint the user confirmed on a previous
|
||||||
|
// attempt, so the retry accepts exactly that key and persists it.
|
||||||
|
trustHostKey string
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSSHClient creates a new SSH client bound to the running NetBird Client.
|
||||||
|
func NewSSHClient(c *Client) *SSHClient {
|
||||||
|
return &SSHClient{nb: c}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetListener registers the Java listener. Must be called before Connect to
|
||||||
|
// receive any events.
|
||||||
|
func (s *SSHClient) SetListener(l SSHTerminalListener) {
|
||||||
|
s.mu.Lock()
|
||||||
|
s.listener = l
|
||||||
|
s.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetURLOpener registers the Java URL opener used to display the device-code
|
||||||
|
// authorization page in a Custom Tabs window when the target peer requires
|
||||||
|
// JWT authentication. Must be set before Connect to be effective.
|
||||||
|
func (s *SSHClient) SetURLOpener(opener URLOpener) {
|
||||||
|
s.mu.Lock()
|
||||||
|
s.urlOpener = opener
|
||||||
|
s.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetKnownHostsStore points the TOFU host-key store at a profile's preferences.
|
||||||
|
// Must be set before connecting to a regular SSH server; without it such a
|
||||||
|
// server cannot be verified and Connect refuses one.
|
||||||
|
func (s *SSHClient) SetKnownHostsStore(configDir, profileID string) {
|
||||||
|
s.mu.Lock()
|
||||||
|
s.knownHostsConfigDir = configDir
|
||||||
|
s.knownHostsProfile = profileID
|
||||||
|
s.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// TrustHostKey records the fingerprint the user confirmed for a regular server,
|
||||||
|
// so the next Connect accepts that exact key and adds it to the known-hosts
|
||||||
|
// store. Passing a fingerprint that no longer matches makes the connect fail
|
||||||
|
// rather than trust a key that changed since the prompt.
|
||||||
|
func (s *SSHClient) TrustHostKey(fingerprint string) {
|
||||||
|
s.mu.Lock()
|
||||||
|
s.trustHostKey = fingerprint
|
||||||
|
s.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Connect dials the SSH server through the NetBird tunnel and performs the
|
||||||
|
// SSH handshake. It auto-detects the server type via SSH banner inspection
|
||||||
|
// and selects the appropriate authentication path:
|
||||||
|
//
|
||||||
|
// - NetBird-SSH server requiring JWT: launches the OAuth 2.0 device-code
|
||||||
|
// flow, opens the verification URL through the registered URLOpener, and
|
||||||
|
// uses the resulting token as the SSH password. Host-key verification
|
||||||
|
// uses the NetBird peer registry.
|
||||||
|
// - NetBird-SSH server without JWT: authenticates with the NetBird SSH
|
||||||
|
// private key. Host-key verification uses the NetBird peer registry.
|
||||||
|
// - Regular SSH server (e.g. OpenSSH): authenticates with the NetBird key
|
||||||
|
// first (so a user-installed NetBird public key works), then falls back
|
||||||
|
// to the supplied password if non-empty. Host-key verification is
|
||||||
|
// trust-on-first-use against the per-profile known-hosts store.
|
||||||
|
//
|
||||||
|
// The password parameter is only consulted for regular SSH servers.
|
||||||
|
func (s *SSHClient) Connect(host string, port int, user, password string) error {
|
||||||
|
if port < 1 || port > 65535 {
|
||||||
|
return fmt.Errorf("invalid port: %d", port)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, cfgPath, cc := s.nb.authSnapshot()
|
||||||
|
if cc == nil {
|
||||||
|
return errors.New("netbird client not running")
|
||||||
|
}
|
||||||
|
if cfg == nil {
|
||||||
|
return errors.New("netbird config not loaded")
|
||||||
|
}
|
||||||
|
engine := cc.Engine()
|
||||||
|
if engine == nil {
|
||||||
|
return errors.New("netbird engine not available")
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
s.gen++
|
||||||
|
gen := s.gen
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
serverType := detectServerType(host, port)
|
||||||
|
log.Debugf("SSH server type: %s", serverType)
|
||||||
|
|
||||||
|
authMethods, hostKeyCallback, err := s.buildAuth(cfg, cfgPath, engine, serverType, password)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
clientConfig := &gossh.ClientConfig{
|
||||||
|
User: user,
|
||||||
|
Auth: authMethods,
|
||||||
|
HostKeyCallback: hostKeyCallback,
|
||||||
|
Timeout: sshDialTimeout,
|
||||||
|
}
|
||||||
|
err = s.dialAndHandshake(gen, host, port, clientConfig)
|
||||||
|
|
||||||
|
// An unknown host key is a prompt, not a failure: return the marker intact
|
||||||
|
// (rootCause would unwrap it) so Java can show the fingerprint and retry.
|
||||||
|
var unknownHost *errHostKeyUnknown
|
||||||
|
if errors.As(err, &unknownHost) {
|
||||||
|
return errors.New(unknownHost.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
// A regular server may still accept a password, so let the caller ask for
|
||||||
|
// one instead of failing. NetBird servers never use a password, so a
|
||||||
|
// failure there is genuine.
|
||||||
|
if err != nil && serverType != detection.ServerTypeNetBirdJWT &&
|
||||||
|
serverType != detection.ServerTypeNetBirdNoJWT && isAuthFailure(err) &&
|
||||||
|
passwordCouldHelp(err, password != "") {
|
||||||
|
return errPasswordRequired
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return rootCause(err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// StartSession requests a PTY and starts an interactive shell. Output from
|
||||||
|
// the session is forwarded to the listener via OnData.
|
||||||
|
func (s *SSHClient) StartSession(cols, rows int) error {
|
||||||
|
err := s.startSession(cols, rows)
|
||||||
|
if err != nil {
|
||||||
|
log.Infof("SSH: start session failed: %v", err)
|
||||||
|
return rootCause(err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write sends data to the SSH session stdin.
|
||||||
|
func (s *SSHClient) Write(data []byte) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
stdin := s.stdin
|
||||||
|
s.mu.Unlock()
|
||||||
|
if stdin == nil {
|
||||||
|
return errors.New("ssh session not started")
|
||||||
|
}
|
||||||
|
if _, err := stdin.Write(data); err != nil {
|
||||||
|
return fmt.Errorf("write stdin: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Resize updates the PTY window size.
|
||||||
|
func (s *SSHClient) Resize(cols, rows int) error {
|
||||||
|
s.mu.Lock()
|
||||||
|
session := s.session
|
||||||
|
s.mu.Unlock()
|
||||||
|
if session == nil {
|
||||||
|
return errors.New("ssh session not started")
|
||||||
|
}
|
||||||
|
return session.WindowChange(rows, cols)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset makes a closed client usable for another Connect: Close leaves the
|
||||||
|
// one-shot guard set, and clearing it lets the same client back a reconnect.
|
||||||
|
func (s *SSHClient) Reset() {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
s.closed = false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close terminates the SSH session and underlying connection. Safe to call
|
||||||
|
// multiple times.
|
||||||
|
func (s *SSHClient) Close() error {
|
||||||
|
s.mu.Lock()
|
||||||
|
s.gen++
|
||||||
|
if s.dialCancel != nil {
|
||||||
|
s.dialCancel()
|
||||||
|
s.dialCancel = nil
|
||||||
|
}
|
||||||
|
sshClient := s.sshClient
|
||||||
|
session := s.session
|
||||||
|
stdin := s.stdin
|
||||||
|
s.sshClient = nil
|
||||||
|
s.session = nil
|
||||||
|
s.stdin = nil
|
||||||
|
notify := !s.closed
|
||||||
|
s.closed = true
|
||||||
|
listener := s.listener
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
if stdin != nil {
|
||||||
|
if err := stdin.Close(); err != nil {
|
||||||
|
log.Debugf("ssh: stdin close: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if session != nil {
|
||||||
|
if err := session.Close(); err != nil && !errors.Is(err, io.EOF) {
|
||||||
|
log.Debugf("ssh: session close: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var firstErr error
|
||||||
|
if sshClient != nil {
|
||||||
|
if err := sshClient.Close(); err != nil {
|
||||||
|
firstErr = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if notify && listener != nil {
|
||||||
|
listener.OnClose("closed by client")
|
||||||
|
}
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *SSHClient) startSession(cols, rows int) error {
|
||||||
|
log.Debugf("SSH: starting session %dx%d", cols, rows)
|
||||||
|
s.mu.Lock()
|
||||||
|
sshClient := s.sshClient
|
||||||
|
gen := s.gen
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
if sshClient == nil {
|
||||||
|
return errors.New("ssh client not connected")
|
||||||
|
}
|
||||||
|
|
||||||
|
pty, err := nbssh.StartPTYSession(sshClient, cols, rows)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
if gen != s.gen {
|
||||||
|
s.mu.Unlock()
|
||||||
|
closeQuiet(pty.Session, "stale session")
|
||||||
|
return errClientClosed
|
||||||
|
}
|
||||||
|
s.session = pty.Session
|
||||||
|
s.stdin = pty.Stdin
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
readerDone := make(chan string, 2)
|
||||||
|
go func() { readerDone <- s.readLoop(pty.Stdout, "stdout") }()
|
||||||
|
go func() { readerDone <- s.readLoop(pty.Stderr, "stderr") }()
|
||||||
|
go func() {
|
||||||
|
reason := <-readerDone
|
||||||
|
if second := <-readerDone; reason == "" {
|
||||||
|
reason = second
|
||||||
|
}
|
||||||
|
s.notifyClose(gen, reason)
|
||||||
|
}()
|
||||||
|
log.Debug("SSH: session started, shell running")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *SSHClient) buildAuth(cfg *profilemanager.Config, cfgPath string, engine *internal.Engine,
|
||||||
|
serverType detection.ServerType, password string) ([]gossh.AuthMethod, gossh.HostKeyCallback, error) {
|
||||||
|
|
||||||
|
switch serverType {
|
||||||
|
case detection.ServerTypeNetBirdJWT:
|
||||||
|
token, err := s.requestJWTToken(cfg, cfgPath)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, fmt.Errorf("jwt: %w", err)
|
||||||
|
}
|
||||||
|
auths := []gossh.AuthMethod{gossh.Password(token)}
|
||||||
|
return auths, nbssh.CreateHostKeyCallback(nbssh.PeerKeyLookup(engine.GetPeerSSHKey)), nil
|
||||||
|
|
||||||
|
case detection.ServerTypeNetBirdNoJWT:
|
||||||
|
if cfg.SSHKey == "" {
|
||||||
|
return nil, nil, errors.New("no NetBird SSH key available")
|
||||||
|
}
|
||||||
|
signer, err := gossh.ParsePrivateKey([]byte(cfg.SSHKey))
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, fmt.Errorf("parse netbird ssh key: %w", err)
|
||||||
|
}
|
||||||
|
auths := []gossh.AuthMethod{gossh.PublicKeys(signer)}
|
||||||
|
return auths, nbssh.CreateHostKeyCallback(nbssh.PeerKeyLookup(engine.GetPeerSSHKey)), nil
|
||||||
|
|
||||||
|
case detection.ServerTypeRegular:
|
||||||
|
var auths []gossh.AuthMethod
|
||||||
|
if cfg.SSHKey != "" {
|
||||||
|
if signer, err := gossh.ParsePrivateKey([]byte(cfg.SSHKey)); err == nil {
|
||||||
|
auths = append(auths, gossh.PublicKeys(signer))
|
||||||
|
} else {
|
||||||
|
log.Debugf("ssh: parse netbird key for regular auth: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if password != "" {
|
||||||
|
pw := password
|
||||||
|
auths = append(auths, gossh.Password(pw))
|
||||||
|
auths = append(auths, gossh.KeyboardInteractive(func(_, _ string, questions []string, _ []bool) ([]string, error) {
|
||||||
|
answers := make([]string, len(questions))
|
||||||
|
for i := range questions {
|
||||||
|
answers[i] = pw
|
||||||
|
}
|
||||||
|
return answers, nil
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
if len(auths) == 0 {
|
||||||
|
// Nothing to offer at all: ask for a password rather than failing,
|
||||||
|
// so the caller can retry once the user supplies one.
|
||||||
|
return nil, nil, errPasswordRequired
|
||||||
|
}
|
||||||
|
callback, err := s.tofuHostKeyCallback()
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
return auths, callback, nil
|
||||||
|
|
||||||
|
default:
|
||||||
|
return nil, nil, fmt.Errorf("unsupported SSH server type: %v", serverType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// tofuHostKeyCallback verifies a regular server's host key against the
|
||||||
|
// per-profile known-hosts store. An unknown host returns errHostKeyUnknown so
|
||||||
|
// Java can show the fingerprint and, once confirmed, retry with the key
|
||||||
|
// trusted; a changed key is rejected outright, as OpenSSH does. When the user
|
||||||
|
// has confirmed a fingerprint, the callback accepts exactly that key and
|
||||||
|
// appends it to the store.
|
||||||
|
func (s *SSHClient) tofuHostKeyCallback() (gossh.HostKeyCallback, error) {
|
||||||
|
s.mu.Lock()
|
||||||
|
configDir := s.knownHostsConfigDir
|
||||||
|
profileID := s.knownHostsProfile
|
||||||
|
trusted := s.trustHostKey
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
if configDir == "" || profileID == "" {
|
||||||
|
return nil, errors.New("no known-hosts store configured for regular SSH")
|
||||||
|
}
|
||||||
|
|
||||||
|
store, err := openKnownHostsStore(configDir, profileID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("load known-hosts store: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return func(hostname string, remote net.Addr, key gossh.PublicKey) error {
|
||||||
|
verdict, err := store.verify(hostname, remote, key)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if verdict == hostKeyMatched {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if verdict == hostKeyChanged {
|
||||||
|
return fmt.Errorf("SSH host key changed for %s (possible attack)", hostname)
|
||||||
|
}
|
||||||
|
|
||||||
|
fingerprint := gossh.FingerprintSHA256(key)
|
||||||
|
if trusted == "" {
|
||||||
|
return &errHostKeyUnknown{fingerprint: fingerprint}
|
||||||
|
}
|
||||||
|
if trusted != fingerprint {
|
||||||
|
return fmt.Errorf("SSH host key changed since it was confirmed for %s", hostname)
|
||||||
|
}
|
||||||
|
if err := store.append(hostname, remote, key); err != nil {
|
||||||
|
return fmt.Errorf("persist trusted host key: %w", err)
|
||||||
|
}
|
||||||
|
// The confirmation is spent: now that the key is stored, a later
|
||||||
|
// reconnect must verify against the file, not re-accept this fingerprint.
|
||||||
|
s.mu.Lock()
|
||||||
|
s.trustHostKey = ""
|
||||||
|
s.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *SSHClient) requestJWTToken(cfg *profilemanager.Config, cfgPath string) (string, error) {
|
||||||
|
s.mu.Lock()
|
||||||
|
urlOpener := s.urlOpener
|
||||||
|
s.mu.Unlock()
|
||||||
|
if urlOpener == nil {
|
||||||
|
return "", errors.New("URL opener not configured for JWT auth")
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
flow, err := auth.NewOAuthFlow(ctx, cfg, false, true, profileLoginHint(cfgPath))
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("create oauth flow: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The status callback covers the browser round-trip, which would
|
||||||
|
// otherwise leave the terminal blank.
|
||||||
|
tokenInfo, err := runOAuthFlow(ctx, flow, urlOpener, func() {
|
||||||
|
s.notifyStatus("Waiting for browser authentication...")
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
token := tokenInfo.GetTokenToUse()
|
||||||
|
if token == "" {
|
||||||
|
return "", errors.New("empty token returned by IdP")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tells the client the browser round-trip is over so it can dismiss the
|
||||||
|
// surface it opened, the same way the login and session-extend flows do.
|
||||||
|
// Without it the Custom Tab stays in front of the terminal even though the
|
||||||
|
// token has already been collected.
|
||||||
|
urlOpener.OnLoginSuccess()
|
||||||
|
|
||||||
|
return token, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *SSHClient) dialAndHandshake(gen uint64, host string, port int, clientConfig *gossh.ClientConfig) error {
|
||||||
|
addr := net.JoinHostPort(host, strconv.Itoa(port))
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), sshDialTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
if gen != s.gen {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return errClientClosed
|
||||||
|
}
|
||||||
|
s.dialCancel = cancel
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
var dialer net.Dialer
|
||||||
|
conn, err := dialer.DialContext(ctx, "tcp", addr)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("dial %s: %w", addr, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
client, err := nbssh.Handshake(ctx, conn, addr, clientConfig)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
s.mu.Lock()
|
||||||
|
if gen != s.gen {
|
||||||
|
s.mu.Unlock()
|
||||||
|
closeQuiet(client, "stale ssh client")
|
||||||
|
return errClientClosed
|
||||||
|
}
|
||||||
|
s.sshClient = client
|
||||||
|
listener := s.listener
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
if listener != nil {
|
||||||
|
listener.OnConnected()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *SSHClient) readLoop(r io.Reader, name string) string {
|
||||||
|
buf := make([]byte, 4096)
|
||||||
|
for {
|
||||||
|
n, err := r.Read(buf)
|
||||||
|
if n > 0 {
|
||||||
|
s.mu.Lock()
|
||||||
|
listener := s.listener
|
||||||
|
s.mu.Unlock()
|
||||||
|
if listener != nil {
|
||||||
|
chunk := make([]byte, n)
|
||||||
|
copy(chunk, buf[:n])
|
||||||
|
listener.OnData(chunk)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
// EOF is a normal shell exit, so report it without a reason.
|
||||||
|
if errors.Is(err, io.EOF) {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
log.Debugf("ssh %s read: %v", name, err)
|
||||||
|
return rootCause(err).Error()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// notifyStatus writes a progress line to the terminal through the normal
|
||||||
|
// output path, so long steps are visible while nothing else is arriving.
|
||||||
|
func (s *SSHClient) notifyStatus(text string) {
|
||||||
|
s.mu.Lock()
|
||||||
|
listener := s.listener
|
||||||
|
s.mu.Unlock()
|
||||||
|
if listener != nil {
|
||||||
|
listener.OnData([]byte("\r\n\x1b[33m" + text + "\x1b[0m\r\n"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *SSHClient) notifyClose(gen uint64, reason string) {
|
||||||
|
s.mu.Lock()
|
||||||
|
if gen != s.gen || s.closed {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.closed = true
|
||||||
|
listener := s.listener
|
||||||
|
s.mu.Unlock()
|
||||||
|
if listener != nil {
|
||||||
|
listener.OnClose(reason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func closeQuiet(c io.Closer, label string) {
|
||||||
|
if c == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := c.Close(); err != nil && !errors.Is(err, io.EOF) {
|
||||||
|
log.Debugf("ssh: close %s: %v", label, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func detectServerType(host string, port int) detection.ServerType {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), sshDetectionTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
dialer := &net.Dialer{}
|
||||||
|
serverType, err := detection.DetectSSHServerType(ctx, dialer, host, port)
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("ssh: server detection failed: %v (assuming regular SSH)", err)
|
||||||
|
return detection.ServerTypeRegular
|
||||||
|
}
|
||||||
|
return serverType
|
||||||
|
}
|
||||||
|
|
||||||
|
// rootCause returns the innermost error of a %w chain, so the terminal shows
|
||||||
|
// "i/o timeout" rather than every layer that added context on the way up.
|
||||||
|
func rootCause(err error) error {
|
||||||
|
for {
|
||||||
|
// A joined error has no single root, so keep it as-is.
|
||||||
|
if _, ok := err.(interface{ Unwrap() []error }); ok {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
next := errors.Unwrap(err)
|
||||||
|
if next == nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
err = next
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// isAuthFailure distinguishes credential rejection from dial, timeout and
|
||||||
|
// host-key errors, which retrying with a password would not fix.
|
||||||
|
func isAuthFailure(err error) bool {
|
||||||
|
if errors.Is(err, errPasswordRequired) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
var partial *gossh.PartialSuccessError
|
||||||
|
if errors.As(err, &partial) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return strings.Contains(err.Error(), "unable to authenticate")
|
||||||
|
}
|
||||||
|
|
||||||
|
// passwordCouldHelp reports whether prompting for a password again can change
|
||||||
|
// the outcome. gossh lists a method under "attempted methods" only when the
|
||||||
|
// server offered it, so a supplied password that was never attempted means the
|
||||||
|
// server does not accept passwords and the real error should surface instead.
|
||||||
|
func passwordCouldHelp(err error, passwordOffered bool) bool {
|
||||||
|
if !passwordOffered {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
msg := err.Error()
|
||||||
|
return strings.Contains(msg, "password") || strings.Contains(msg, "keyboard-interactive")
|
||||||
|
}
|
||||||
@@ -0,0 +1,168 @@
|
|||||||
|
//go:build android
|
||||||
|
|
||||||
|
package android
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
gossh "golang.org/x/crypto/ssh"
|
||||||
|
"golang.org/x/crypto/ssh/knownhosts"
|
||||||
|
)
|
||||||
|
|
||||||
|
const knownHostsNamespace = "ssh"
|
||||||
|
|
||||||
|
const (
|
||||||
|
hostKeyUnknown hostKeyVerdict = iota
|
||||||
|
hostKeyMatched
|
||||||
|
hostKeyChanged
|
||||||
|
)
|
||||||
|
|
||||||
|
var knownHostsMu sync.Mutex
|
||||||
|
|
||||||
|
type hostKeyVerdict uint8
|
||||||
|
|
||||||
|
type knownHostsSection struct {
|
||||||
|
KnownHosts []string `json:"knownHosts"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type knownHostsStore struct {
|
||||||
|
prefs prefsStore
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveKnownHost deletes every known-hosts entry for host:port from the
|
||||||
|
// profile's store, so a host trusted for a session that is being deleted does
|
||||||
|
// not linger. Java calls this only once no session targets that host, so a
|
||||||
|
// shared host stays trusted. A missing entry is not an error: the goal state
|
||||||
|
// is "absent".
|
||||||
|
func RemoveKnownHost(configDir, profileID, host string, port int) error {
|
||||||
|
store, err := openKnownHostsStore(configDir, profileID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return store.removeHost(host, port)
|
||||||
|
}
|
||||||
|
|
||||||
|
func openKnownHostsStore(configDir, profileID string) (*knownHostsStore, error) {
|
||||||
|
prefs, err := newProfilePrefs(configDir, profileID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &knownHostsStore{prefs: prefs}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (st *knownHostsStore) verify(hostname string, remote net.Addr, key gossh.PublicKey) (hostKeyVerdict, error) {
|
||||||
|
lines, err := st.lines()
|
||||||
|
if err != nil {
|
||||||
|
return hostKeyUnknown, err
|
||||||
|
}
|
||||||
|
targets := knownHostsTargets(hostname, remote)
|
||||||
|
|
||||||
|
verdict := hostKeyUnknown
|
||||||
|
for _, line := range lines {
|
||||||
|
pubKey, ok := knownHostsLineKey(line, targets)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if pubKey.Type() == key.Type() && bytes.Equal(pubKey.Marshal(), key.Marshal()) {
|
||||||
|
return hostKeyMatched, nil
|
||||||
|
}
|
||||||
|
verdict = hostKeyChanged
|
||||||
|
}
|
||||||
|
return verdict, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (st *knownHostsStore) append(hostname string, remote net.Addr, key gossh.PublicKey) error {
|
||||||
|
line := knownhosts.Line(knownHostsTargets(hostname, remote), key)
|
||||||
|
|
||||||
|
knownHostsMu.Lock()
|
||||||
|
defer knownHostsMu.Unlock()
|
||||||
|
|
||||||
|
lines, err := st.lines()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return st.prefs.Put(knownHostsNamespace, knownHostsSection{KnownHosts: append(lines, line)})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (st *knownHostsStore) removeHost(host string, port int) error {
|
||||||
|
target := knownhosts.Normalize(net.JoinHostPort(host, strconv.Itoa(port)))
|
||||||
|
|
||||||
|
knownHostsMu.Lock()
|
||||||
|
defer knownHostsMu.Unlock()
|
||||||
|
|
||||||
|
lines, err := st.lines()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
kept := make([]string, 0, len(lines))
|
||||||
|
for _, line := range lines {
|
||||||
|
if knownHostsLineMatches(line, target) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
kept = append(kept, line)
|
||||||
|
}
|
||||||
|
if len(kept) == len(lines) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return st.prefs.Put(knownHostsNamespace, knownHostsSection{KnownHosts: kept})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (st *knownHostsStore) lines() ([]string, error) {
|
||||||
|
var section knownHostsSection
|
||||||
|
if _, err := st.prefs.Get(knownHostsNamespace, §ion); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return section.KnownHosts, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func knownHostsTargets(hostname string, remote net.Addr) []string {
|
||||||
|
targets := []string{knownhosts.Normalize(hostname)}
|
||||||
|
if remote != nil {
|
||||||
|
if normalized := knownhosts.Normalize(remote.String()); normalized != targets[0] {
|
||||||
|
targets = append(targets, normalized)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return targets
|
||||||
|
}
|
||||||
|
|
||||||
|
func knownHostsLineKey(line string, targets []string) (gossh.PublicKey, bool) {
|
||||||
|
trimmed := strings.TrimSpace(line)
|
||||||
|
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
_, hosts, pubKey, _, _, err := gossh.ParseKnownHosts([]byte(trimmed))
|
||||||
|
if err != nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
for _, host := range hosts {
|
||||||
|
for _, target := range targets {
|
||||||
|
if host == target {
|
||||||
|
return pubKey, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// knownHostsLineMatches reports whether a known-hosts line's address list
|
||||||
|
// contains the normalized target. Comment and blank lines never match.
|
||||||
|
func knownHostsLineMatches(line, target string) bool {
|
||||||
|
trimmed := strings.TrimSpace(line)
|
||||||
|
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
fields := strings.Fields(trimmed)
|
||||||
|
if len(fields) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, addr := range strings.Split(fields[0], ",") {
|
||||||
|
if addr == target {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
@@ -0,0 +1,104 @@
|
|||||||
|
//go:build android
|
||||||
|
|
||||||
|
package android
|
||||||
|
|
||||||
|
const (
|
||||||
|
sshSessionsNamespace = "ssh-sessions"
|
||||||
|
maxStoredSSHSessions = 50
|
||||||
|
)
|
||||||
|
|
||||||
|
type sshSessionRecord struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Host string `json:"host"`
|
||||||
|
Port int `json:"port"`
|
||||||
|
User string `json:"user"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type sshSessionsSection struct {
|
||||||
|
Sessions []sshSessionRecord `json:"sessions"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSHSessionEntry is one stored SSH session, without any credential.
|
||||||
|
type SSHSessionEntry struct {
|
||||||
|
ID string
|
||||||
|
Host string
|
||||||
|
Port int
|
||||||
|
User string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSHSessionArray wraps stored SSH sessions for gomobile compatibility.
|
||||||
|
type SSHSessionArray struct {
|
||||||
|
items []*SSHSessionEntry
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSSHSessionArray creates an empty session array to fill via Add.
|
||||||
|
func NewSSHSessionArray() *SSHSessionArray {
|
||||||
|
return &SSHSessionArray{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add appends a session entry, oldest first.
|
||||||
|
func (a *SSHSessionArray) Add(id, host string, port int, user string) {
|
||||||
|
a.items = append(a.items, &SSHSessionEntry{ID: id, Host: host, Port: port, User: user})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Length returns the number of entries.
|
||||||
|
func (a *SSHSessionArray) Length() int {
|
||||||
|
return len(a.items)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get returns the entry at index i, or nil when out of range.
|
||||||
|
func (a *SSHSessionArray) Get(i int) *SSHSessionEntry {
|
||||||
|
if i < 0 || i >= len(a.items) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return a.items[i]
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSHSessionStore reads and writes a profile's stored SSH sessions.
|
||||||
|
type SSHSessionStore struct {
|
||||||
|
prefs prefsStore
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSSHSessionStore opens the session store of the given profile.
|
||||||
|
func NewSSHSessionStore(configDir, profileID string) (*SSHSessionStore, error) {
|
||||||
|
prefs, err := newProfilePrefs(configDir, profileID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &SSHSessionStore{prefs: prefs}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load returns the stored sessions, oldest first.
|
||||||
|
func (s *SSHSessionStore) Load() (*SSHSessionArray, error) {
|
||||||
|
var section sshSessionsSection
|
||||||
|
if _, err := s.prefs.Get(sshSessionsNamespace, §ion); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
out := NewSSHSessionArray()
|
||||||
|
for _, record := range section.Sessions {
|
||||||
|
if record.ID == "" || record.Host == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
out.Add(record.ID, record.Host, record.Port, record.User)
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Save replaces the stored sessions, keeping only the newest entries when the
|
||||||
|
// list exceeds the storage cap.
|
||||||
|
func (s *SSHSessionStore) Save(sessions *SSHSessionArray) error {
|
||||||
|
var items []*SSHSessionEntry
|
||||||
|
if sessions != nil {
|
||||||
|
items = sessions.items
|
||||||
|
}
|
||||||
|
if len(items) > maxStoredSSHSessions {
|
||||||
|
items = items[len(items)-maxStoredSSHSessions:]
|
||||||
|
}
|
||||||
|
|
||||||
|
records := make([]sshSessionRecord, 0, len(items))
|
||||||
|
for _, item := range items {
|
||||||
|
records = append(records, sshSessionRecord{ID: item.ID, Host: item.Host, Port: item.Port, User: item.User})
|
||||||
|
}
|
||||||
|
return s.prefs.Put(sshSessionsNamespace, sshSessionsSection{Sessions: records})
|
||||||
|
}
|
||||||
@@ -305,6 +305,12 @@ func (a *Anonymizer) AnonymizeDomain(domain string) string {
|
|||||||
return domain
|
return domain
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// A reverse zone names an address prefix, so it follows the address rules,
|
||||||
|
// which also keeps its digit labels intact.
|
||||||
|
if zone, ok := a.anonymizeReverseZone(baseDomain); ok {
|
||||||
|
return withTrailingDot(zone, hasDot)
|
||||||
|
}
|
||||||
|
|
||||||
if suffix := protectedSuffix(baseDomain); suffix != "" {
|
if suffix := protectedSuffix(baseDomain); suffix != "" {
|
||||||
if a.level < LevelStrict || baseDomain == suffix || suffix == infraDomain {
|
if a.level < LevelStrict || baseDomain == suffix || suffix == infraDomain {
|
||||||
return domain
|
return domain
|
||||||
@@ -405,6 +411,10 @@ func (a *Anonymizer) AnonymizeString(str string) string {
|
|||||||
ipv4Regex := regexp.MustCompile(`\b(?:[0-9]{1,3}\.){3}[0-9]{1,3}\b`)
|
ipv4Regex := regexp.MustCompile(`\b(?:[0-9]{1,3}\.){3}[0-9]{1,3}\b`)
|
||||||
ipv6Regex := regexp.MustCompile(`\b([0-9a-fA-F:]+:+[0-9a-fA-F]{0,4})(?:%[0-9a-zA-Z]+)?(?:\/[0-9]{1,3})?(?::[0-9]{1,5})?\b`)
|
ipv6Regex := regexp.MustCompile(`\b([0-9a-fA-F:]+:+[0-9a-fA-F]{0,4})(?:%[0-9a-zA-Z]+)?(?:\/[0-9]{1,3})?(?::[0-9]{1,5})?\b`)
|
||||||
|
|
||||||
|
// Reverse zones go first and are then held out of the passes below: their
|
||||||
|
// labels are digits, which the address patterns would otherwise consume.
|
||||||
|
str, restoreZones := a.replaceReverseZones(str)
|
||||||
|
|
||||||
str = ipv4Regex.ReplaceAllStringFunc(str, a.AnonymizeIPString)
|
str = ipv4Regex.ReplaceAllStringFunc(str, a.AnonymizeIPString)
|
||||||
str = ipv6Regex.ReplaceAllStringFunc(str, a.AnonymizeIPString)
|
str = ipv6Regex.ReplaceAllStringFunc(str, a.AnonymizeIPString)
|
||||||
|
|
||||||
@@ -425,7 +435,7 @@ func (a *Anonymizer) AnonymizeString(str string) string {
|
|||||||
str = wgKeyRegex.ReplaceAllStringFunc(str, a.AnonymizeWGKey)
|
str = wgKeyRegex.ReplaceAllStringFunc(str, a.AnonymizeWGKey)
|
||||||
}
|
}
|
||||||
|
|
||||||
return str
|
return restoreZones(str)
|
||||||
}
|
}
|
||||||
|
|
||||||
// sortedDomains returns the domain mappings longest-first, so a full-FQDN
|
// sortedDomains returns the domain mappings longest-first, so a full-FQDN
|
||||||
|
|||||||
@@ -0,0 +1,174 @@
|
|||||||
|
package anonymize
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/hex"
|
||||||
|
"net/netip"
|
||||||
|
"regexp"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
reverseZoneSuffixV4 = ".in-addr.arpa"
|
||||||
|
reverseZoneSuffixV6 = ".ip6.arpa"
|
||||||
|
|
||||||
|
v6Nibbles = 32
|
||||||
|
v4Octets = 4
|
||||||
|
)
|
||||||
|
|
||||||
|
// reverseZoneRegexes match a reverse zone or a full reverse name in free text.
|
||||||
|
// They are applied before the address passes of AnonymizeString, whose IPv4
|
||||||
|
// pattern would otherwise consume the digit labels of a zone and replace parts
|
||||||
|
// of it with unrelated addresses.
|
||||||
|
var reverseZoneRegexes = []*regexp.Regexp{
|
||||||
|
regexp.MustCompile(`(?:[0-9]{1,3}\.){1,4}in-addr\.arpa\b`),
|
||||||
|
regexp.MustCompile(`(?:[0-9a-fA-F]\.){1,32}ip6\.arpa\b`),
|
||||||
|
}
|
||||||
|
|
||||||
|
// anonymizeReverseZone maps a reverse zone to the zone of the anonymized form
|
||||||
|
// of the prefix it encodes, so it follows the address rules rather than the
|
||||||
|
// domain ones: the zone of an address that is preserved is preserved too, and
|
||||||
|
// the zone of one that is replaced names the replacement. This keeps a reverse
|
||||||
|
// zone recognizable as such, and consistent with the addresses it belongs to
|
||||||
|
// elsewhere in the same output. It reports false for anything that is not a
|
||||||
|
// reverse zone.
|
||||||
|
func (a *Anonymizer) anonymizeReverseZone(domain string) (string, bool) {
|
||||||
|
prefix, labelCount, suffix, ok := parseReverseZone(domain)
|
||||||
|
if !ok {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
anonymized := a.AnonymizeIP(prefix)
|
||||||
|
if anonymized == prefix {
|
||||||
|
return domain, true
|
||||||
|
}
|
||||||
|
|
||||||
|
return reverseZoneName(anonymized, labelCount) + suffix, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// replaceReverseZones anonymizes every reverse zone in str and swaps each one
|
||||||
|
// for a placeholder, returning a function that puts the anonymized zones back.
|
||||||
|
// The placeholders carry no dots, digits or colons, so no later pass matches
|
||||||
|
// them.
|
||||||
|
func (a *Anonymizer) replaceReverseZones(str string) (string, func(string) string) {
|
||||||
|
var zones []string
|
||||||
|
|
||||||
|
for _, re := range reverseZoneRegexes {
|
||||||
|
str = re.ReplaceAllStringFunc(str, func(match string) string {
|
||||||
|
zone, ok := a.anonymizeReverseZone(match)
|
||||||
|
if !ok {
|
||||||
|
return match
|
||||||
|
}
|
||||||
|
|
||||||
|
zones = append(zones, zone)
|
||||||
|
return reverseZonePlaceholder(len(zones) - 1)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(zones) == 0 {
|
||||||
|
return str, func(s string) string { return s }
|
||||||
|
}
|
||||||
|
|
||||||
|
return str, func(s string) string {
|
||||||
|
for i, zone := range zones {
|
||||||
|
s = strings.ReplaceAll(s, reverseZonePlaceholder(i), zone)
|
||||||
|
}
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func reverseZonePlaceholder(index int) string {
|
||||||
|
return "\x00reversezone" + strconv.Itoa(index) + "\x00"
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseReverseZone turns a reverse zone into the address of the prefix its
|
||||||
|
// labels spell backwards, padding the absent low-order part with zeroes, and
|
||||||
|
// returns the label count and zone suffix so the name can be rebuilt.
|
||||||
|
func parseReverseZone(domain string) (netip.Addr, int, string, bool) {
|
||||||
|
lower := strings.ToLower(domain)
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case strings.HasSuffix(lower, reverseZoneSuffixV4):
|
||||||
|
labels := strings.Split(strings.TrimSuffix(lower, reverseZoneSuffixV4), ".")
|
||||||
|
addr, ok := reverseZoneAddrV4(labels)
|
||||||
|
return addr, len(labels), reverseZoneSuffixV4, ok
|
||||||
|
case strings.HasSuffix(lower, reverseZoneSuffixV6):
|
||||||
|
labels := strings.Split(strings.TrimSuffix(lower, reverseZoneSuffixV6), ".")
|
||||||
|
addr, ok := reverseZoneAddrV6(labels)
|
||||||
|
return addr, len(labels), reverseZoneSuffixV6, ok
|
||||||
|
default:
|
||||||
|
return netip.Addr{}, 0, "", false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func reverseZoneAddrV4(labels []string) (netip.Addr, bool) {
|
||||||
|
if len(labels) == 0 || len(labels) > v4Octets {
|
||||||
|
return netip.Addr{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
var octets [v4Octets]byte
|
||||||
|
for i, label := range labels {
|
||||||
|
octet, err := strconv.ParseUint(label, 10, 8)
|
||||||
|
if err != nil {
|
||||||
|
return netip.Addr{}, false
|
||||||
|
}
|
||||||
|
octets[len(labels)-1-i] = byte(octet)
|
||||||
|
}
|
||||||
|
|
||||||
|
return netip.AddrFrom4(octets), true
|
||||||
|
}
|
||||||
|
|
||||||
|
func reverseZoneAddrV6(labels []string) (netip.Addr, bool) {
|
||||||
|
if len(labels) == 0 || len(labels) > v6Nibbles {
|
||||||
|
return netip.Addr{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
nibbles := make([]byte, 0, v6Nibbles)
|
||||||
|
for i := len(labels) - 1; i >= 0; i-- {
|
||||||
|
if len(labels[i]) != 1 || !isHexDigit(labels[i][0]) {
|
||||||
|
return netip.Addr{}, false
|
||||||
|
}
|
||||||
|
nibbles = append(nibbles, labels[i][0])
|
||||||
|
}
|
||||||
|
for len(nibbles) < v6Nibbles {
|
||||||
|
nibbles = append(nibbles, '0')
|
||||||
|
}
|
||||||
|
|
||||||
|
var groups []string
|
||||||
|
for i := 0; i < len(nibbles); i += 4 {
|
||||||
|
groups = append(groups, string(nibbles[i:i+4]))
|
||||||
|
}
|
||||||
|
|
||||||
|
addr, err := netip.ParseAddr(strings.Join(groups, ":"))
|
||||||
|
if err != nil {
|
||||||
|
return netip.Addr{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
return addr, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// reverseZoneName spells the first labelCount labels of addr backwards, the
|
||||||
|
// inverse of parseReverseZone, without the zone suffix.
|
||||||
|
func reverseZoneName(addr netip.Addr, labelCount int) string {
|
||||||
|
labels := make([]string, 0, labelCount)
|
||||||
|
|
||||||
|
if addr.Is4() {
|
||||||
|
octets := addr.As4()
|
||||||
|
for i := labelCount - 1; i >= 0; i-- {
|
||||||
|
labels = append(labels, strconv.Itoa(int(octets[i])))
|
||||||
|
}
|
||||||
|
return strings.Join(labels, ".")
|
||||||
|
}
|
||||||
|
|
||||||
|
address := addr.As16()
|
||||||
|
nibbles := hex.EncodeToString(address[:])
|
||||||
|
for i := labelCount - 1; i >= 0; i-- {
|
||||||
|
labels = append(labels, string(nibbles[i]))
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(labels, ".")
|
||||||
|
}
|
||||||
|
|
||||||
|
func isHexDigit(c byte) bool {
|
||||||
|
return c >= '0' && c <= '9' || c >= 'a' && c <= 'f' || c >= 'A' && c <= 'F'
|
||||||
|
}
|
||||||
@@ -0,0 +1,171 @@
|
|||||||
|
package anonymize
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newLeveledAnonymizer(level Level) *Anonymizer {
|
||||||
|
a := NewAnonymizer(DefaultAddresses())
|
||||||
|
a.SetLevel(level)
|
||||||
|
return a
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestAnonymizeDomainReverseZone covers reverse zones going through the address
|
||||||
|
// rules instead of the domain ones, so a zone stays a zone and an address that
|
||||||
|
// is preserved keeps the zone that names it.
|
||||||
|
func TestAnonymizeDomainReverseZone(t *testing.T) {
|
||||||
|
// 100.64.0.0/10 is the overlay range, which is CGNAT: preserved at the
|
||||||
|
// default level and replaced from the internal pool at the strict one
|
||||||
|
const overlayZone = "64.100.in-addr.arpa"
|
||||||
|
|
||||||
|
t.Run("overlay zone preserved at the default level", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
assert.Equal(t, overlayZone, a.AnonymizeDomain(overlayZone), "should keep the zone of a preserved address")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("private zone preserved at the default level", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
assert.Equal(t, "168.192.in-addr.arpa", a.AnonymizeDomain("168.192.in-addr.arpa"), "should keep the zone of a private address")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("overlay zone replaced at the strict level", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelStrict)
|
||||||
|
|
||||||
|
got := a.AnonymizeDomain(overlayZone)
|
||||||
|
require.True(t, strings.HasSuffix(got, reverseZoneSuffixV4), "should stay a reverse zone, got %q", got)
|
||||||
|
assert.NotEqual(t, overlayZone, got, "should replace the encoded prefix")
|
||||||
|
assert.Len(t, strings.Split(strings.TrimSuffix(got, reverseZoneSuffixV4), "."), 2,
|
||||||
|
"should keep the label count, got %q", got)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("public zone replaced at the default level", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
|
||||||
|
got := a.AnonymizeDomain("113.0.203.in-addr.arpa")
|
||||||
|
require.True(t, strings.HasSuffix(got, reverseZoneSuffixV4), "should stay a reverse zone, got %q", got)
|
||||||
|
assert.NotEqual(t, "113.0.203.in-addr.arpa", got, "should replace a public prefix")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("zone of an address keeps that address mapping", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
|
||||||
|
anonymizedAddr := a.AnonymizeIPString("203.0.113.7")
|
||||||
|
got := a.AnonymizeDomain("7.113.0.203.in-addr.arpa")
|
||||||
|
|
||||||
|
octets := strings.Split(anonymizedAddr, ".")
|
||||||
|
want := octets[3] + "." + octets[2] + "." + octets[1] + "." + octets[0] + reverseZoneSuffixV4
|
||||||
|
assert.Equal(t, want, got, "should name the same replacement as the address itself")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("ipv6 nibble labels stay single digits", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
|
||||||
|
zone := "0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.2.0.0.0" + reverseZoneSuffixV6
|
||||||
|
got := a.AnonymizeDomain(zone)
|
||||||
|
|
||||||
|
require.True(t, strings.HasSuffix(got, reverseZoneSuffixV6), "should stay a reverse zone, got %q", got)
|
||||||
|
labels := strings.Split(strings.TrimSuffix(got, reverseZoneSuffixV6), ".")
|
||||||
|
assert.Len(t, labels, 28, "should keep every nibble label, got %q", got)
|
||||||
|
for _, label := range labels {
|
||||||
|
assert.Len(t, label, 1, "nibble label %q should stay a single digit", label)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("trailing dot is kept", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
assert.Equal(t, "64.100.in-addr.arpa.", a.AnonymizeDomain("64.100.in-addr.arpa."), "should keep the trailing dot")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a domain that only looks like a zone is anonymized as a domain", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
|
||||||
|
got := a.AnonymizeDomain("not-a-zone.in-addr.arpa")
|
||||||
|
assert.NotContains(t, got, "in-addr.arpa", "should fall back to domain anonymization")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestAnonymizeStringReverseZone verifies that a zone inside free text, such as
|
||||||
|
// a DNS log line, is not chewed up by the address passes. The IPv4 pattern
|
||||||
|
// matches any run of dotted digits, which a reverse zone is made of.
|
||||||
|
func TestAnonymizeStringReverseZone(t *testing.T) {
|
||||||
|
t.Run("ipv6 zone survives the address passes", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
|
||||||
|
zone := "0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.2.0.0.0" + reverseZoneSuffixV6
|
||||||
|
got := a.AnonymizeString("question: domain=" + zone + " type=PTR")
|
||||||
|
|
||||||
|
assert.Contains(t, got, "type=PTR", "should keep the rest of the line")
|
||||||
|
assert.NotContains(t, got, "198.51.100", "should not rewrite nibble labels as an address")
|
||||||
|
|
||||||
|
labels := strings.Split(strings.TrimSuffix(strings.TrimPrefix(got, "question: domain="), reverseZoneSuffixV6+" type=PTR"), ".")
|
||||||
|
assert.Len(t, labels, 28, "should keep every nibble label, got %q", got)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("preserved ipv4 zone is untouched", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
|
||||||
|
line := "reverse zone 64.100.in-addr.arpa registered"
|
||||||
|
assert.Equal(t, line, a.AnonymizeString(line), "should keep the zone of a preserved address")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("public ipv4 zone is replaced consistently", func(t *testing.T) {
|
||||||
|
a := newLeveledAnonymizer(LevelDefault)
|
||||||
|
|
||||||
|
got := a.AnonymizeString("zone 113.0.203.in-addr.arpa and address 203.0.113.7")
|
||||||
|
assert.NotContains(t, got, "113.0.203.in-addr.arpa", "should replace the zone")
|
||||||
|
assert.NotContains(t, got, "203.0.113.7", "should replace the address")
|
||||||
|
assert.Contains(t, got, reverseZoneSuffixV4, "should keep the zone suffix")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseReverseZone(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
zone string
|
||||||
|
addr string
|
||||||
|
labels int
|
||||||
|
}{
|
||||||
|
{name: "v4 two labels", zone: "0.100" + reverseZoneSuffixV4, addr: "100.0.0.0", labels: 2},
|
||||||
|
{name: "v4 three labels", zone: "1.168.192" + reverseZoneSuffixV4, addr: "192.168.1.0", labels: 3},
|
||||||
|
{name: "v4 full address", zone: "7.113.0.203" + reverseZoneSuffixV4, addr: "203.0.113.7", labels: 4},
|
||||||
|
{
|
||||||
|
name: "v6 prefix",
|
||||||
|
zone: "0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.0.2.0.0.0" + reverseZoneSuffixV6,
|
||||||
|
addr: "2::",
|
||||||
|
labels: 28,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
addr, labels, suffix, ok := parseReverseZone(tc.zone)
|
||||||
|
require.True(t, ok, "should decode the reverse zone")
|
||||||
|
assert.Equal(t, tc.addr, addr.String(), "should decode to the encoded prefix")
|
||||||
|
assert.Equal(t, tc.labels, labels, "should count the labels")
|
||||||
|
assert.Equal(t, tc.zone, reverseZoneName(addr, labels)+suffix, "should re-encode to the original zone")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseReverseZoneRejectsNonZones(t *testing.T) {
|
||||||
|
tests := []string{
|
||||||
|
"example.com",
|
||||||
|
"in-addr.arpa",
|
||||||
|
"x.100" + reverseZoneSuffixV4,
|
||||||
|
"256" + reverseZoneSuffixV4,
|
||||||
|
"1.2.3.4.5" + reverseZoneSuffixV4,
|
||||||
|
"ab" + reverseZoneSuffixV6,
|
||||||
|
"g" + reverseZoneSuffixV6,
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, zone := range tests {
|
||||||
|
t.Run(zone, func(t *testing.T) {
|
||||||
|
_, _, _, ok := parseReverseZone(zone)
|
||||||
|
assert.False(t, ok, "should reject %q", zone)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
+18
-7
@@ -21,7 +21,7 @@ import (
|
|||||||
"github.com/netbirdio/netbird/client/internal/auth"
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
"github.com/netbirdio/netbird/client/internal/peer"
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
sshcommon "github.com/netbirdio/netbird/client/ssh"
|
nbssh "github.com/netbirdio/netbird/client/ssh"
|
||||||
"github.com/netbirdio/netbird/client/system"
|
"github.com/netbirdio/netbird/client/system"
|
||||||
"github.com/netbirdio/netbird/shared/management/domain"
|
"github.com/netbirdio/netbird/shared/management/domain"
|
||||||
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
||||||
@@ -91,6 +91,13 @@ type Options struct {
|
|||||||
// 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.
|
||||||
@@ -220,6 +227,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)
|
||||||
}
|
}
|
||||||
@@ -521,12 +537,7 @@ func (c *Client) VerifySSHHostKey(peerAddress string, key []byte) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
storedKey, found := engine.GetPeerSSHKey(peerAddress)
|
return nbssh.PeerKeyLookup(engine.GetPeerSSHKey).VerifySSHHostKey(peerAddress, key)
|
||||||
if !found {
|
|
||||||
return sshcommon.ErrPeerNotFound
|
|
||||||
}
|
|
||||||
|
|
||||||
return sshcommon.VerifyHostKey(storedKey, key, peerAddress)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetPerformance retunes a running Client. Only PreallocatedBuffersPerPool
|
// SetPerformance retunes a running Client. Only PreallocatedBuffersPerPool
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -16,28 +16,47 @@ 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/netsweep"
|
||||||
)
|
)
|
||||||
|
|
||||||
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 *netsweep.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
|
||||||
|
}
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package grpc
|
|||||||
import (
|
import (
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/netsweep"
|
||||||
"github.com/netbirdio/netbird/util/wsproxy/client"
|
"github.com/netbirdio/netbird/util/wsproxy/client"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -11,3 +12,8 @@ import (
|
|||||||
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(_ *netsweep.Sweeper) grpc.DialOption {
|
||||||
|
return grpc.EmptyDialOption{}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,49 @@
|
|||||||
|
package grpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/cenkalti/backoff/v4"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/netstate"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 netState never fires, leaving plain backoff.Retry
|
||||||
|
// behavior.
|
||||||
|
func Retry(ctx context.Context, operation backoff.Operation, bo backoff.BackOff, netState *netstate.State) 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
|
||||||
|
}
|
||||||
|
|
||||||
|
timer := time.NewTimer(next)
|
||||||
|
select {
|
||||||
|
case <-timer.C:
|
||||||
|
case <-netState.Changed():
|
||||||
|
timer.Stop()
|
||||||
|
case <-ctx.Done():
|
||||||
|
timer.Stop()
|
||||||
|
return ctx.Err()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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/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")
|
||||||
|
}
|
||||||
@@ -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"
|
||||||
|
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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() {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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).
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -147,7 +147,7 @@ func (a *Auth) IsSSOSupported(ctx context.Context) (bool, error) {
|
|||||||
|
|
||||||
// GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection
|
// GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection
|
||||||
// This avoids creating a new connection to the management server
|
// This avoids creating a new connection to the management server
|
||||||
func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool) (OAuthFlow, error) {
|
func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool, hint string) (OAuthFlow, error) {
|
||||||
var flow OAuthFlow
|
var flow OAuthFlow
|
||||||
|
|
||||||
// the connection is owned by a and outlives this call, so a later fallback reuses it
|
// the connection is owned by a and outlives this call, so a later fallback reuses it
|
||||||
@@ -157,7 +157,7 @@ func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool) (OAuthFlo
|
|||||||
|
|
||||||
err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
|
err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
|
||||||
var err error
|
var err error
|
||||||
flow, err = oauthFlowWithFallback(a, client, flowOrder(forceDeviceAuth, true), "", newAuth)
|
flow, err = oauthFlowWithFallback(a, client, flowOrder(forceDeviceAuth, true), hint, newAuth)
|
||||||
|
|
||||||
if IsSSOUnavailable(err) {
|
if IsSSOUnavailable(err) {
|
||||||
return backoff.Permanent(err)
|
return backoff.Permanent(err)
|
||||||
|
|||||||
@@ -38,6 +38,8 @@ 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/netstate"
|
||||||
|
"github.com/netbirdio/netbird/client/netsweep"
|
||||||
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"
|
||||||
@@ -70,18 +72,42 @@ type ConnectClient struct {
|
|||||||
updateManager *updater.Manager
|
updateManager *updater.Manager
|
||||||
|
|
||||||
persistSyncResponse bool
|
persistSyncResponse bool
|
||||||
|
|
||||||
|
// netState gates every reconnection loop on OS-reported network
|
||||||
|
// availability. Nil (the default) disables gating; mobile platforms
|
||||||
|
// inject it via WithNetworkState.
|
||||||
|
netState *netstate.State
|
||||||
|
|
||||||
|
// sweeper cuts the management, signal and relay connections on network
|
||||||
|
// change; nil disables it.
|
||||||
|
sweeper *netsweep.Sweeper
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConnectClientOption configures optional ConnectClient behavior.
|
||||||
|
type ConnectClientOption func(*ConnectClient)
|
||||||
|
|
||||||
|
// WithNetworkState injects the OS network availability state that gates every
|
||||||
|
// reconnection loop; without it gating is disabled.
|
||||||
|
func WithNetworkState(netState *netstate.State) ConnectClientOption {
|
||||||
|
return func(c *ConnectClient) { c.netState = netState }
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithSweeper injects the network change sweeper.
|
||||||
|
func WithSweeper(sweeper *netsweep.Sweeper) ConnectClientOption {
|
||||||
|
return func(c *ConnectClient) { c.sweeper = sweeper }
|
||||||
}
|
}
|
||||||
|
|
||||||
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{}),
|
||||||
@@ -89,6 +115,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) {
|
||||||
@@ -274,6 +304,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.netState.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)
|
||||||
@@ -285,7 +322,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.WithNetworkState(c.netState), mgm.WithSweeper(c.sweeper))
|
||||||
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
|
||||||
@@ -360,7 +398,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.netState, c.sweeper)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Error(err)
|
log.Error(err)
|
||||||
return wrapErr(err)
|
return wrapErr(err)
|
||||||
@@ -396,7 +434,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.WithNetworkState(c.netState), relayClient.WithSweeper(c.sweeper))
|
||||||
c.statusRecorder.SetRelayMgr(relayManager)
|
c.statusRecorder.SetRelayMgr(relayManager)
|
||||||
if len(relayURLs) > 0 {
|
if len(relayURLs) > 0 {
|
||||||
if token != nil {
|
if token != nil {
|
||||||
@@ -424,6 +463,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
|
|||||||
UpdateManager: c.updateManager,
|
UpdateManager: c.updateManager,
|
||||||
ClientMetrics: c.clientMetrics,
|
ClientMetrics: c.clientMetrics,
|
||||||
MetricsCtx: c.ctx,
|
MetricsCtx: c.ctx,
|
||||||
|
NetState: c.netState,
|
||||||
}, mobileDependency)
|
}, mobileDependency)
|
||||||
engine.SetSyncResponsePersistence(c.persistSyncResponse)
|
engine.SetSyncResponsePersistence(c.persistSyncResponse)
|
||||||
c.engine = engine
|
c.engine = engine
|
||||||
@@ -480,6 +520,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)
|
||||||
@@ -673,7 +723,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, netState *netstate.State, sweeper *netsweep.Sweeper) (*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
|
||||||
@@ -681,7 +731,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.WithNetworkState(netState), signal.WithSweeper(sweeper))
|
||||||
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)
|
||||||
|
|||||||
@@ -51,6 +51,7 @@ nftables.txt: Anonymized nftables rules with packet counters across all families
|
|||||||
sysctls.txt: Forwarding, reverse-path filter, source-validation, and conntrack accounting sysctl values that the NetBird client may read or modify, if --system-info flag was provided (Linux only).
|
sysctls.txt: Forwarding, reverse-path filter, source-validation, and conntrack accounting sysctl values that the NetBird client may read or modify, if --system-info flag was provided (Linux only).
|
||||||
resolv.conf: DNS resolver configuration from /etc/resolv.conf (Unix systems only), if --system-info flag was provided.
|
resolv.conf: DNS resolver configuration from /etc/resolv.conf (Unix systems only), if --system-info flag was provided.
|
||||||
scutil_dns.txt: DNS configuration from scutil --dns (macOS only), if --system-info flag was provided.
|
scutil_dns.txt: DNS configuration from scutil --dns (macOS only), if --system-info flag was provided.
|
||||||
|
dns_windows.txt: Anonymized NRPT rules and policy table in effect, DNS client policy, and per-interface and per-adapter DNS configuration (Windows only), if --system-info flag was provided.
|
||||||
resolved_domains.txt: Anonymized resolved domain IP addresses from the status recorder.
|
resolved_domains.txt: Anonymized resolved domain IP addresses from the status recorder.
|
||||||
config.txt: Anonymized configuration information of the NetBird client.
|
config.txt: Anonymized configuration information of the NetBird client.
|
||||||
network_map.json: Anonymized sync response containing peer configurations, routes, DNS settings, and firewall rules.
|
network_map.json: Anonymized sync response containing peer configurations, routes, DNS settings, and firewall rules.
|
||||||
@@ -237,6 +238,13 @@ scutil_dns.txt (macOS only):
|
|||||||
- Shows DNS configuration for all network interfaces
|
- Shows DNS configuration for all network interfaces
|
||||||
- Includes search domains, nameservers, and DNS resolver settings
|
- Includes search domains, nameservers, and DNS resolver settings
|
||||||
- All IP addresses and domain names are anonymized
|
- All IP addresses and domain names are anonymized
|
||||||
|
|
||||||
|
dns_windows.txt (Windows only):
|
||||||
|
- Lists the NRPT rules of both policy stores, the local one and the group policy one, marking the rules the client created
|
||||||
|
- Follows them with the policy table the resolver has loaded, which differs from the rules while a change has not been picked up yet
|
||||||
|
- Includes the DNS client group policy, the global TCP/IP and Dnscache parameters, and the DNS values of every interface that has any
|
||||||
|
- Ends with the resolver configuration in effect per adapter, from GetAdaptersAddresses
|
||||||
|
- All IP addresses and domain names are anonymized
|
||||||
`
|
`
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
//go:build !unix
|
//go:build !unix && !windows
|
||||||
|
|
||||||
package debug
|
package debug
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,443 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package debug
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/hex"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
|
"strings"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
"golang.org/x/sys/windows/registry"
|
||||||
|
|
||||||
|
nbdns "github.com/netbirdio/netbird/client/internal/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
const dnsInfoFileName = "dns_windows.txt"
|
||||||
|
|
||||||
|
const (
|
||||||
|
gpoDNSClientRoot = `SOFTWARE\Policies\Microsoft\Windows NT\DNSClient`
|
||||||
|
tcpipParamsPath = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters`
|
||||||
|
dnscacheParams = `SYSTEM\CurrentControlSet\Services\Dnscache\Parameters`
|
||||||
|
)
|
||||||
|
|
||||||
|
// interfaceDNSValues are the per-interface values that decide how a name is
|
||||||
|
// resolved and registered. Everything the DNS host manager writes is in here,
|
||||||
|
// so a bundle shows both what we set and what it replaced.
|
||||||
|
var interfaceDNSValues = []string{
|
||||||
|
"NameServer",
|
||||||
|
"DhcpNameServer",
|
||||||
|
"Domain",
|
||||||
|
"DhcpDomain",
|
||||||
|
"SearchList",
|
||||||
|
"RegistrationEnabled",
|
||||||
|
"DisableDynamicUpdate",
|
||||||
|
"MaxNumberOfAddressesToRegister",
|
||||||
|
"EnableDHCP",
|
||||||
|
}
|
||||||
|
|
||||||
|
// addDNSInfo collects and adds DNS configuration information to the archive
|
||||||
|
func (g *BundleGenerator) addDNSInfo() error {
|
||||||
|
if err := g.addFileToZip(strings.NewReader(g.collectDNSInfo()), dnsInfoFileName); err != nil {
|
||||||
|
return fmt.Errorf("add DNS info to zip: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// collectDNSInfo renders the report. Everything below it reaches the platform
|
||||||
|
// through COM and through lazily resolved procedures, which panic when a
|
||||||
|
// procedure is missing rather than returning an error, and a debug bundle is not
|
||||||
|
// allowed to take the daemon down. The panic is contained here, and whatever was
|
||||||
|
// collected before it is kept and reported with it.
|
||||||
|
func (g *BundleGenerator) collectDNSInfo() (content string) {
|
||||||
|
var sb strings.Builder
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
log.Errorf("collecting Windows DNS configuration panicked: %v", r)
|
||||||
|
fmt.Fprintf(&sb, "\nerror: collection stopped: %v\n", r)
|
||||||
|
}
|
||||||
|
content = sb.String()
|
||||||
|
}()
|
||||||
|
|
||||||
|
sb.WriteString("Windows DNS configuration\n")
|
||||||
|
sb.WriteString("=========================\n")
|
||||||
|
|
||||||
|
adapters, adaptersErr := adapterAddresses()
|
||||||
|
|
||||||
|
g.writeNRPTRules(&sb, "NRPT rules, local policy store", nbdns.DNSPolicyConfigRoot)
|
||||||
|
g.writeNRPTRules(&sb, "NRPT rules, group policy store", nbdns.GPODNSPolicyConfigRoot)
|
||||||
|
g.writeEffectiveNRPTPolicies(&sb)
|
||||||
|
g.writeRegistryKey(&sb, "DNS client group policy", gpoDNSClientRoot)
|
||||||
|
g.writeRegistryKey(&sb, "Global TCP/IP parameters", tcpipParamsPath)
|
||||||
|
g.writeRegistryKey(&sb, "Dnscache parameters", dnscacheParams)
|
||||||
|
g.writeInterfaceDNS(&sb, "Per-interface DNS, IPv4", nbdns.InterfaceConfigPath, adapterNames(adapters))
|
||||||
|
g.writeInterfaceDNS(&sb, "Per-interface DNS, IPv6", nbdns.InterfaceConfigPathV6, adapterNames(adapters))
|
||||||
|
g.writeAdapterDNS(&sb, adapters, adaptersErr)
|
||||||
|
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeNRPTRules lists every rule in a policy store, ours and any other
|
||||||
|
// product's, since a foreign rule for the same namespace decides resolution
|
||||||
|
// just as ours does. Rules the client wrote are marked.
|
||||||
|
func (g *BundleGenerator) writeNRPTRules(sb *strings.Builder, title, root string) {
|
||||||
|
writeSection(sb, title, root)
|
||||||
|
|
||||||
|
names, err := subKeyNames(root)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Fprintf(sb, "error: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(names) == 0 {
|
||||||
|
sb.WriteString("no rules\n")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, name := range names {
|
||||||
|
owner := ""
|
||||||
|
if strings.HasPrefix(strings.ToLower(name), strings.ToLower(nbdns.NRPTKeyPrefix)) {
|
||||||
|
owner = " (netbird)"
|
||||||
|
}
|
||||||
|
fmt.Fprintf(sb, "%s%s\n", name, owner)
|
||||||
|
g.writeValues(sb, root+`\`+name, nil, " ")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeEffectiveNRPTPolicies reports the table the resolver answers from, which
|
||||||
|
// the registry cannot show: a rule is written before it is loaded, and it keeps
|
||||||
|
// being enforced after its key is gone until the resolver reloads its policy.
|
||||||
|
func (g *BundleGenerator) writeEffectiveNRPTPolicies(sb *strings.Builder) {
|
||||||
|
writeSection(sb, "NRPT policy table in effect", nrptPolicyClass+"."+nrptPolicyMethod+" in "+nrptPolicyNamespace)
|
||||||
|
|
||||||
|
entries, err := effectiveNRPTPolicies()
|
||||||
|
if err != nil {
|
||||||
|
fmt.Fprintf(sb, "error: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(entries) == 0 {
|
||||||
|
sb.WriteString("no policies\n")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, entry := range entries {
|
||||||
|
fmt.Fprintf(sb, "%s\n", g.anonymizeValue("Namespace", entry.namespace))
|
||||||
|
for _, value := range entry.values {
|
||||||
|
fmt.Fprintf(sb, " %s: %s\n", value.name, g.anonymizeValue(value.name, value.value))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeInterfaceDNS reports the DNS values of every interface that has any, so
|
||||||
|
// the netbird interface can be compared against the physical ones. The registry
|
||||||
|
// keys the values by GUID, so each is named from the adapter list; a GUID with
|
||||||
|
// no adapter is a leftover key of an interface that no longer exists.
|
||||||
|
func (g *BundleGenerator) writeInterfaceDNS(sb *strings.Builder, title, root string, names map[string]string) {
|
||||||
|
writeSection(sb, title, root)
|
||||||
|
|
||||||
|
guids, err := subKeyNames(root)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Fprintf(sb, "error: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var reported int
|
||||||
|
for _, guid := range guids {
|
||||||
|
var iface strings.Builder
|
||||||
|
g.writeValues(&iface, root+`\`+guid, interfaceDNSValues, " ")
|
||||||
|
if iface.Len() == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
name, ok := names[strings.ToLower(guid)]
|
||||||
|
if !ok {
|
||||||
|
name = "no adapter with this GUID"
|
||||||
|
}
|
||||||
|
|
||||||
|
reported++
|
||||||
|
fmt.Fprintf(sb, "%s (%s)\n%s", guid, name, iface.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
if reported == 0 {
|
||||||
|
sb.WriteString("no interface holds DNS values\n")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeRegistryKey reports the values of a single key, without its subkeys.
|
||||||
|
func (g *BundleGenerator) writeRegistryKey(sb *strings.Builder, title, path string) {
|
||||||
|
writeSection(sb, title, path)
|
||||||
|
|
||||||
|
var values strings.Builder
|
||||||
|
g.writeValues(&values, path, nil, "")
|
||||||
|
if values.Len() == 0 {
|
||||||
|
sb.WriteString("no values\n")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
sb.WriteString(values.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeValues renders the values of a key. A nil names list reports every
|
||||||
|
// value, otherwise only those named and present.
|
||||||
|
func (g *BundleGenerator) writeValues(sb *strings.Builder, path string, names []string, indent string) {
|
||||||
|
k, err := registry.OpenKey(registry.LOCAL_MACHINE, path, registry.QUERY_VALUE)
|
||||||
|
switch {
|
||||||
|
case errors.Is(err, registry.ErrNotExist), errors.Is(err, windows.ERROR_PATH_NOT_FOUND):
|
||||||
|
// an absent key is the normal state for the GPO store and for
|
||||||
|
// interfaces without DNS settings
|
||||||
|
log.Debugf("HKEY_LOCAL_MACHINE\\%s does not exist", path)
|
||||||
|
return
|
||||||
|
case err != nil:
|
||||||
|
fmt.Fprintf(sb, "%serror: open HKEY_LOCAL_MACHINE\\%s: %v\n", indent, path, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer closeKey(k)
|
||||||
|
|
||||||
|
if names == nil {
|
||||||
|
names, err = k.ReadValueNames(-1)
|
||||||
|
if err != nil {
|
||||||
|
fmt.Fprintf(sb, "%serror: read value names: %v\n", indent, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, name := range names {
|
||||||
|
value, err := readRegistryValue(k, name)
|
||||||
|
switch {
|
||||||
|
case errors.Is(err, registry.ErrNotExist):
|
||||||
|
// the caller asks for a fixed set of values, most of which a
|
||||||
|
// given interface does not carry
|
||||||
|
continue
|
||||||
|
case err != nil:
|
||||||
|
// report rather than omit: a value that is there but cannot be
|
||||||
|
// read reads as unset otherwise
|
||||||
|
fmt.Fprintf(sb, "%s%s: error: %v\n", indent, name, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(sb, "%s%s: %s\n", indent, name, g.anonymizeValue(name, value))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// anonymizeValue redacts a registry value according to what its name says it
|
||||||
|
// holds. Domains and addresses are handled per entry rather than by the string
|
||||||
|
// pass: the pass only replaces domains something else in the bundle already
|
||||||
|
// seeded, and its address regex would eat the digit labels of a reverse zone.
|
||||||
|
func (g *BundleGenerator) anonymizeValue(name, value string) string {
|
||||||
|
if !g.anonymize || value == "" {
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case holdsDomains(name):
|
||||||
|
return joinValueEntries(splitValueEntries(value), g.anonymizeDomain)
|
||||||
|
case holdsAddresses(name):
|
||||||
|
return joinValueEntries(splitValueEntries(value), g.anonymizer.AnonymizeIPString)
|
||||||
|
default:
|
||||||
|
return g.anonymizer.AnonymizeString(value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// holdsDomains reports whether a value name holds domains: the domain list of
|
||||||
|
// an NRPT rule (Name) or of the policy table (Namespace), a search list, the
|
||||||
|
// DNS suffix values of the TCP/IP and policy keys, which all end in "Domain"
|
||||||
|
// (Domain, DhcpDomain, NV Domain, ICSDomain), and a proxy host name.
|
||||||
|
func holdsDomains(name string) bool {
|
||||||
|
lower := strings.ToLower(name)
|
||||||
|
return lower == "name" || lower == "namespace" || lower == "searchlist" ||
|
||||||
|
strings.HasSuffix(lower, "domain") || strings.HasSuffix(lower, "proxyname")
|
||||||
|
}
|
||||||
|
|
||||||
|
// holdsAddresses reports whether a value name holds DNS server addresses
|
||||||
|
// (NameServer, DhcpNameServer, GenericDNSServers, NameServers).
|
||||||
|
func holdsAddresses(name string) bool {
|
||||||
|
lower := strings.ToLower(name)
|
||||||
|
return strings.Contains(lower, "nameserver") || strings.Contains(lower, "dnsserver")
|
||||||
|
}
|
||||||
|
|
||||||
|
// adapterNames maps adapter GUIDs, as the registry keys the interfaces, to the
|
||||||
|
// names an operator sees.
|
||||||
|
func adapterNames(adapters []*windows.IpAdapterAddresses) map[string]string {
|
||||||
|
names := make(map[string]string, len(adapters))
|
||||||
|
for _, adapter := range adapters {
|
||||||
|
guid := windows.BytePtrToString(adapter.AdapterName)
|
||||||
|
names[strings.ToLower(guid)] = windows.UTF16PtrToString(adapter.FriendlyName)
|
||||||
|
}
|
||||||
|
return names
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeAdapterDNS reports the resolver configuration in effect per adapter,
|
||||||
|
// which is what the resolver uses for a name no NRPT rule matches.
|
||||||
|
func (g *BundleGenerator) writeAdapterDNS(sb *strings.Builder, adapters []*windows.IpAdapterAddresses, err error) {
|
||||||
|
writeSection(sb, "Adapter DNS configuration", "GetAdaptersAddresses")
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
fmt.Fprintf(sb, "error: %v\n", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, adapter := range adapters {
|
||||||
|
name := windows.UTF16PtrToString(adapter.FriendlyName)
|
||||||
|
suffix := g.anonymizeDomain(windows.UTF16PtrToString(adapter.DnsSuffix))
|
||||||
|
|
||||||
|
fmt.Fprintf(sb, "%s (index %d, oper status %d)\n", name, adapter.IfIndex, adapter.OperStatus)
|
||||||
|
fmt.Fprintf(sb, " DNS suffix: %s\n", suffix)
|
||||||
|
|
||||||
|
var servers []string
|
||||||
|
for server := adapter.FirstDnsServerAddress; server != nil; server = server.Next {
|
||||||
|
addr, ok := netip.AddrFromSlice(server.Address.IP())
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
addr = addr.Unmap()
|
||||||
|
if g.anonymize {
|
||||||
|
addr = g.anonymizer.AnonymizeIP(addr)
|
||||||
|
}
|
||||||
|
servers = append(servers, addr.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Fprintf(sb, " DNS servers: %s\n", strings.Join(servers, ", "))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// anonymizeDomain anonymizes a single domain, keeping the leading dot an NRPT
|
||||||
|
// match domain carries.
|
||||||
|
func (g *BundleGenerator) anonymizeDomain(entry string) string {
|
||||||
|
if !g.anonymize {
|
||||||
|
return entry
|
||||||
|
}
|
||||||
|
|
||||||
|
domain, dot := strings.CutPrefix(entry, ".")
|
||||||
|
if domain == "" {
|
||||||
|
return entry
|
||||||
|
}
|
||||||
|
|
||||||
|
anonymized := g.anonymizer.AnonymizeDomain(domain)
|
||||||
|
if dot {
|
||||||
|
anonymized = "." + anonymized
|
||||||
|
}
|
||||||
|
return anonymized
|
||||||
|
}
|
||||||
|
|
||||||
|
// splitValueEntries splits a registry value that holds a list. The separator
|
||||||
|
// differs per value: a REG_MULTI_SZ arrives joined with ", ", a SearchList is
|
||||||
|
// comma separated and a NameServer may use commas or spaces.
|
||||||
|
func splitValueEntries(value string) []string {
|
||||||
|
return strings.FieldsFunc(value, func(r rune) bool {
|
||||||
|
return r == ',' || r == ';' || r == ' ' || r == '\t'
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func joinValueEntries(entries []string, anonymize func(string) string) string {
|
||||||
|
for i, entry := range entries {
|
||||||
|
entries[i] = anonymize(entry)
|
||||||
|
}
|
||||||
|
return strings.Join(entries, ", ")
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeSection(sb *strings.Builder, title, source string) {
|
||||||
|
fmt.Fprintf(sb, "\n%s\n%s\n%s\n", title, strings.Repeat("-", len(title)), source)
|
||||||
|
}
|
||||||
|
|
||||||
|
func subKeyNames(root string) ([]string, error) {
|
||||||
|
k, err := registry.OpenKey(registry.LOCAL_MACHINE, root, registry.ENUMERATE_SUB_KEYS)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", root, err)
|
||||||
|
}
|
||||||
|
defer closeKey(k)
|
||||||
|
|
||||||
|
names, err := k.ReadSubKeyNames(-1)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("read subkey names: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return names, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// readRegistryValue renders a value as text regardless of its type, so an
|
||||||
|
// unexpected type in a policy key still shows up instead of being dropped.
|
||||||
|
func readRegistryValue(k registry.Key, name string) (string, error) {
|
||||||
|
_, valueType, err := k.GetValue(name, nil)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("get value %s: %w", name, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
switch valueType {
|
||||||
|
case registry.SZ, registry.EXPAND_SZ:
|
||||||
|
value, _, err := k.GetStringValue(name)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("get string value %s: %w", name, err)
|
||||||
|
}
|
||||||
|
return value, nil
|
||||||
|
case registry.MULTI_SZ:
|
||||||
|
values, _, err := k.GetStringsValue(name)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("get strings value %s: %w", name, err)
|
||||||
|
}
|
||||||
|
return strings.Join(values, ", "), nil
|
||||||
|
case registry.DWORD, registry.QWORD:
|
||||||
|
value, _, err := k.GetIntegerValue(name)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("get integer value %s: %w", name, err)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%d (0x%x)", value, value), nil
|
||||||
|
case registry.BINARY:
|
||||||
|
value, _, err := k.GetBinaryValue(name)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("get binary value %s: %w", name, err)
|
||||||
|
}
|
||||||
|
return hex.EncodeToString(value), nil
|
||||||
|
default:
|
||||||
|
return fmt.Sprintf("<unhandled registry type %d>", valueType), nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// adapterAddresses returns the adapter list including DNS servers. The call
|
||||||
|
// reports the size it needs, so grow the buffer and retry until it fits.
|
||||||
|
func adapterAddresses() (adapters []*windows.IpAdapterAddresses, err error) {
|
||||||
|
// GetAdaptersAddresses is resolved on first use and panics when it is
|
||||||
|
// missing, so this reports it as an error and leaves the rest of the
|
||||||
|
// report intact.
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
adapters, err = nil, fmt.Errorf("GetAdaptersAddresses: %v", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
const flags = windows.GAA_FLAG_SKIP_ANYCAST | windows.GAA_FLAG_SKIP_MULTICAST
|
||||||
|
|
||||||
|
size := uint32(15000)
|
||||||
|
for range 3 {
|
||||||
|
buf := make([]byte, size)
|
||||||
|
first := (*windows.IpAdapterAddresses)(unsafe.Pointer(&buf[0]))
|
||||||
|
|
||||||
|
err := windows.GetAdaptersAddresses(windows.AF_UNSPEC, flags, 0, first, &size)
|
||||||
|
if errors.Is(err, windows.ERROR_BUFFER_OVERFLOW) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("GetAdaptersAddresses: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for adapter := first; adapter != nil; adapter = adapter.Next {
|
||||||
|
adapters = append(adapters, adapter)
|
||||||
|
}
|
||||||
|
return adapters, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("GetAdaptersAddresses: buffer kept growing")
|
||||||
|
}
|
||||||
|
|
||||||
|
func closeKey(k registry.Key) {
|
||||||
|
if err := k.Close(); err != nil {
|
||||||
|
log.Debugf("close registry key: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,146 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package debug
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/anonymize"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newDNSValueGenerator(level anonymize.Level) *BundleGenerator {
|
||||||
|
anonymizer := anonymize.NewAnonymizer(anonymize.DefaultAddresses())
|
||||||
|
anonymizer.SetLevel(level)
|
||||||
|
|
||||||
|
return &BundleGenerator{
|
||||||
|
anonymize: true,
|
||||||
|
anonymizeLevel: level,
|
||||||
|
anonymizer: anonymizer,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestAnonymizeValueByName covers the value kinds of the DNS registry keys. The
|
||||||
|
// names decide the treatment, because the string pass alone replaces only
|
||||||
|
// domains another part of the bundle already seeded.
|
||||||
|
func TestAnonymizeValueByName(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
valueName string
|
||||||
|
value string
|
||||||
|
assert func(t *testing.T, got string)
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "NRPT match domains keep the leading dot",
|
||||||
|
valueName: "Name",
|
||||||
|
value: ".internal.example.com, .corp.example.org",
|
||||||
|
assert: func(t *testing.T, got string) {
|
||||||
|
t.Helper()
|
||||||
|
for _, entry := range strings.Split(got, ", ") {
|
||||||
|
assert.True(t, strings.HasPrefix(entry, "."), "entry %q should keep its leading dot", entry)
|
||||||
|
assert.NotContains(t, entry, "example", "entry %q should not keep the original domain", entry)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "any value name ending in Domain is treated as a domain",
|
||||||
|
valueName: "ICSDomain",
|
||||||
|
value: "mshome.net",
|
||||||
|
assert: func(t *testing.T, got string) {
|
||||||
|
t.Helper()
|
||||||
|
assert.NotContains(t, got, "mshome", "should anonymize a domain suffix value")
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "search list is a comma separated domain list",
|
||||||
|
valueName: "SearchList",
|
||||||
|
value: "corp.example.com,branch.example.com",
|
||||||
|
assert: func(t *testing.T, got string) {
|
||||||
|
t.Helper()
|
||||||
|
assert.NotContains(t, got, "example", "should anonymize every search domain")
|
||||||
|
assert.Len(t, strings.Split(got, ", "), 2, "should keep both search domains")
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "name servers are anonymized as addresses",
|
||||||
|
valueName: "DhcpNameServer",
|
||||||
|
value: "203.0.113.10 8.8.8.8",
|
||||||
|
assert: func(t *testing.T, got string) {
|
||||||
|
t.Helper()
|
||||||
|
assert.NotContains(t, got, "203.0.113.10", "should anonymize a public resolver address")
|
||||||
|
// well-known resolvers stay readable at every level
|
||||||
|
assert.Contains(t, got, "8.8.8.8", "should keep a well-known resolver address")
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "opaque values are left to the string pass",
|
||||||
|
valueName: "DataBasePath",
|
||||||
|
value: `%SystemRoot%\System32\drivers\etc`,
|
||||||
|
assert: func(t *testing.T, got string) {
|
||||||
|
t.Helper()
|
||||||
|
assert.Equal(t, `%SystemRoot%\System32\drivers\etc`, got, "should not alter a path")
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
g := newDNSValueGenerator(anonymize.LevelDefault)
|
||||||
|
tc.assert(t, g.anonymizeValue(tc.valueName, tc.value))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestParseNRPTPolicyTable parses the MOF text of the policy table out
|
||||||
|
// parameters, as the provider on a client with one NRPT rule renders it.
|
||||||
|
func TestParseNRPTPolicyTable(t *testing.T) {
|
||||||
|
const text = `[abstract]
|
||||||
|
class __PARAMETERS
|
||||||
|
{
|
||||||
|
[Out, EmbeddedInstance("DnsClientPolicyConfiguration"): ToSubClass, ID(2): DisableOverride ToInstance] DnsClientPolicyConfiguration cmdletOutput[] = {
|
||||||
|
instance of DnsClientPolicyConfiguration
|
||||||
|
{
|
||||||
|
DirectAccessProxyType = "NoProxy";
|
||||||
|
DirectAccessQueryIPsecRequired = FALSE;
|
||||||
|
NameEncoding = "Utf8WithoutMapping";
|
||||||
|
Namespace = ".0.100.in-addr.arpa";
|
||||||
|
},
|
||||||
|
instance of DnsClientPolicyConfiguration
|
||||||
|
{
|
||||||
|
DirectAccessProxyType = "NoProxy";
|
||||||
|
NameEncoding = "Utf8WithoutMapping";
|
||||||
|
NameServers = {"100.0.255.254", "100.0.255.253"};
|
||||||
|
Namespace = ".nb.internal";
|
||||||
|
}};
|
||||||
|
[in] boolean Effective;
|
||||||
|
[out] uint32 ReturnValue = 0;
|
||||||
|
};
|
||||||
|
`
|
||||||
|
|
||||||
|
entries := parseNRPTPolicyTable(text)
|
||||||
|
require.Len(t, entries, 2, "should parse both embedded instances")
|
||||||
|
|
||||||
|
assert.Equal(t, ".0.100.in-addr.arpa", entries[0].namespace, "should read the namespace of the first instance")
|
||||||
|
assert.Equal(t, ".nb.internal", entries[1].namespace, "should read the namespace of the second instance")
|
||||||
|
|
||||||
|
assert.Equal(t, []registryValue{
|
||||||
|
{name: "DirectAccessProxyType", value: "NoProxy"},
|
||||||
|
{name: "DirectAccessQueryIPsecRequired", value: "FALSE"},
|
||||||
|
{name: "NameEncoding", value: "Utf8WithoutMapping"},
|
||||||
|
}, entries[0].values, "should keep the remaining values in order")
|
||||||
|
|
||||||
|
assert.Contains(t, entries[1].values, registryValue{name: "NameServers", value: "100.0.255.254, 100.0.255.253"},
|
||||||
|
"should flatten a MOF array")
|
||||||
|
|
||||||
|
for _, value := range entries[1].values {
|
||||||
|
assert.NotContains(t, value.name, "ReturnValue", "should not read the class level parameters as values")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseNRPTPolicyTableEmpty(t *testing.T) {
|
||||||
|
assert.Empty(t, parseNRPTPolicyTable(""), "should parse no entries from empty text")
|
||||||
|
assert.Empty(t, parseNRPTPolicyTable("class __PARAMETERS\n{\n};\n"), "should parse no entries from a table with no instances")
|
||||||
|
}
|
||||||
@@ -0,0 +1,317 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package debug
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/go-ole/go-ole"
|
||||||
|
"github.com/go-ole/go-ole/oleutil"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// The NRPT policy table is reachable through the CIM class that backs
|
||||||
|
// Get-DnsClientNrptPolicy. Unlike the rules in the registry, the table is
|
||||||
|
// what the resolver currently has loaded, which is the only way to tell an
|
||||||
|
// applied rule from one that is merely written, in either direction.
|
||||||
|
nrptPolicyNamespace = `root\Microsoft\Windows\DNS`
|
||||||
|
nrptPolicyClass = "PS_DnsClientNrptPolicy"
|
||||||
|
nrptPolicyMethod = "Get"
|
||||||
|
|
||||||
|
// The class has no instances, so the table comes from the out parameters
|
||||||
|
// of a static method call, rendered as MOF text: the embedded instances
|
||||||
|
// arrive as a safe array of objects, which cannot be read back through the
|
||||||
|
// COM bindings, and the text form carries all of them.
|
||||||
|
nrptPolicyInstanceKeyword = "instance of DnsClientPolicyConfiguration"
|
||||||
|
|
||||||
|
nrptPolicyTimeout = 15 * time.Second
|
||||||
|
)
|
||||||
|
|
||||||
|
// COM initialization results that leave the calling thread usable: S_FALSE for
|
||||||
|
// a thread this process already initialized, RPC_E_CHANGED_MODE for one that
|
||||||
|
// belongs to another apartment.
|
||||||
|
const (
|
||||||
|
sFalse = 0x00000001
|
||||||
|
rpcEChangedMode = 0x80010106
|
||||||
|
)
|
||||||
|
|
||||||
|
// nrptQueryInFlight admits one read of the policy table at a time. A provider
|
||||||
|
// that stops answering keeps its goroutine and the OS thread that goroutine
|
||||||
|
// pinned, so a later bundle reports that instead of pinning another one.
|
||||||
|
var nrptQueryInFlight = make(chan struct{}, 1)
|
||||||
|
|
||||||
|
// nrptPolicyEntry is one namespace of the effective policy table, holding the
|
||||||
|
// values of an embedded DnsClientPolicyConfiguration instance in the order the
|
||||||
|
// provider reported them.
|
||||||
|
type nrptPolicyEntry struct {
|
||||||
|
namespace string
|
||||||
|
values []registryValue
|
||||||
|
}
|
||||||
|
|
||||||
|
// registryValue is a name and its rendered value, shared by the registry and
|
||||||
|
// policy table readers so both anonymize by value name the same way.
|
||||||
|
type registryValue struct {
|
||||||
|
name string
|
||||||
|
value string
|
||||||
|
}
|
||||||
|
|
||||||
|
// effectiveNRPTPolicies reads the effective NRPT table. The call is bounded
|
||||||
|
// because a WMI provider can block indefinitely and a debug bundle must not.
|
||||||
|
func effectiveNRPTPolicies() ([]nrptPolicyEntry, error) {
|
||||||
|
type result struct {
|
||||||
|
text string
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case nrptQueryInFlight <- struct{}{}:
|
||||||
|
default:
|
||||||
|
return nil, errors.New("an earlier read of the policy table has not returned")
|
||||||
|
}
|
||||||
|
|
||||||
|
done := make(chan result, 1)
|
||||||
|
go func() {
|
||||||
|
// the slot is released here rather than by the caller, so a read that
|
||||||
|
// outlives the timeout holds it until the provider answers
|
||||||
|
defer func() { <-nrptQueryInFlight }()
|
||||||
|
|
||||||
|
text, err := nrptPolicyTableText()
|
||||||
|
done <- result{text: text, err: err}
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case res := <-done:
|
||||||
|
if res.err != nil {
|
||||||
|
return nil, res.err
|
||||||
|
}
|
||||||
|
return parseNRPTPolicyTable(res.text), nil
|
||||||
|
case <-time.After(nrptPolicyTimeout):
|
||||||
|
return nil, errors.New("read of the policy table timed out")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// nrptPolicyTableText calls the policy table method and returns the MOF text of
|
||||||
|
// its out parameters.
|
||||||
|
func nrptPolicyTableText() (text string, err error) {
|
||||||
|
// COM is per thread, and the collection is short lived, so the thread is
|
||||||
|
// pinned for the duration rather than initialized for the process.
|
||||||
|
runtime.LockOSThread()
|
||||||
|
defer runtime.UnlockOSThread()
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
// The COM call chain is dynamically typed, so a provider that answers
|
||||||
|
// with an unexpected shape must not take the daemon down with it.
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
err = fmt.Errorf("read NRPT policy table: %v", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
owns, err := coInitialize()
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
if owns {
|
||||||
|
defer ole.CoUninitialize()
|
||||||
|
}
|
||||||
|
|
||||||
|
locator, err := oleutil.CreateObject("WbemScripting.SWbemLocator")
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("create WMI locator: %w", err)
|
||||||
|
}
|
||||||
|
defer locator.Release()
|
||||||
|
|
||||||
|
dispatch, err := locator.QueryInterface(ole.IID_IDispatch)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("query WMI locator interface: %w", err)
|
||||||
|
}
|
||||||
|
defer dispatch.Release()
|
||||||
|
|
||||||
|
service, err := dispatchCall(dispatch, "ConnectServer", nil, nrptPolicyNamespace)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("connect to %s: %w", nrptPolicyNamespace, err)
|
||||||
|
}
|
||||||
|
defer service.Release()
|
||||||
|
|
||||||
|
inParams, err := spawnMethodInParams(service)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer inParams.Release()
|
||||||
|
|
||||||
|
// The effective table is the merge of the local and the group policy
|
||||||
|
// store, which is what the resolver answers from.
|
||||||
|
if _, err := oleutil.PutProperty(inParams, "Effective", true); err != nil {
|
||||||
|
return "", fmt.Errorf("set Effective parameter: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
outParams, err := dispatchCall(service, "ExecMethod", nrptPolicyClass, nrptPolicyMethod, inParams)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("call %s.%s: %w", nrptPolicyClass, nrptPolicyMethod, err)
|
||||||
|
}
|
||||||
|
defer outParams.Release()
|
||||||
|
|
||||||
|
textVariant, err := oleutil.CallMethod(outParams, "GetObjectText_")
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("render policy table: %w", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := textVariant.Clear(); err != nil {
|
||||||
|
log.Debugf("clear policy table variant: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return textVariant.ToString(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// spawnMethodInParams builds the in parameters instance the method needs. The
|
||||||
|
// provider rejects the call without one, even when every parameter is optional.
|
||||||
|
func spawnMethodInParams(service *ole.IDispatch) (*ole.IDispatch, error) {
|
||||||
|
class, err := dispatchCall(service, "Get", nrptPolicyClass)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("get class %s: %w", nrptPolicyClass, err)
|
||||||
|
}
|
||||||
|
defer class.Release()
|
||||||
|
|
||||||
|
methods, err := dispatchProperty(class, "Methods_")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("get class methods: %w", err)
|
||||||
|
}
|
||||||
|
defer methods.Release()
|
||||||
|
|
||||||
|
method, err := dispatchCall(methods, "Item", nrptPolicyMethod)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("get method %s: %w", nrptPolicyMethod, err)
|
||||||
|
}
|
||||||
|
defer method.Release()
|
||||||
|
|
||||||
|
params, err := dispatchProperty(method, "InParameters")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("get method parameters: %w", err)
|
||||||
|
}
|
||||||
|
defer params.Release()
|
||||||
|
|
||||||
|
inParams, err := dispatchCall(params, "SpawnInstance_")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("spawn parameter instance: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return inParams, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseNRPTPolicyTable pulls the embedded instances out of the MOF text. Each
|
||||||
|
// instance is a namespace of the table, with one name and value per line.
|
||||||
|
func parseNRPTPolicyTable(text string) []nrptPolicyEntry {
|
||||||
|
var entries []nrptPolicyEntry
|
||||||
|
var current *nrptPolicyEntry
|
||||||
|
|
||||||
|
for _, line := range strings.Split(text, "\n") {
|
||||||
|
line = strings.TrimSpace(strings.TrimSuffix(strings.TrimSpace(line), ";"))
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case strings.HasPrefix(line, nrptPolicyInstanceKeyword):
|
||||||
|
entries = append(entries, nrptPolicyEntry{})
|
||||||
|
current = &entries[len(entries)-1]
|
||||||
|
continue
|
||||||
|
case strings.HasPrefix(line, "}"):
|
||||||
|
// closes an instance, and the array with the last one, so the
|
||||||
|
// class level parameters that follow are not read as values
|
||||||
|
current = nil
|
||||||
|
continue
|
||||||
|
case current == nil, line == "{":
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
name, value, ok := strings.Cut(line, " = ")
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
value = unquoteMOFValue(value)
|
||||||
|
if name == "Namespace" {
|
||||||
|
current.namespace = value
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
current.values = append(current.values, registryValue{name: name, value: value})
|
||||||
|
}
|
||||||
|
|
||||||
|
return entries
|
||||||
|
}
|
||||||
|
|
||||||
|
// unquoteMOFValue renders a MOF scalar or array as plain text: "a" becomes a,
|
||||||
|
// and {"a", "b"} becomes a, b.
|
||||||
|
func unquoteMOFValue(value string) string {
|
||||||
|
value = strings.TrimSpace(value)
|
||||||
|
|
||||||
|
if inner, ok := strings.CutPrefix(value, "{"); ok {
|
||||||
|
value = strings.TrimSuffix(inner, "}")
|
||||||
|
|
||||||
|
entries := strings.Split(value, ",")
|
||||||
|
for i, entry := range entries {
|
||||||
|
entries[i] = strings.Trim(strings.TrimSpace(entry), `"`)
|
||||||
|
}
|
||||||
|
return strings.Join(entries, ", ")
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Trim(value, `"`)
|
||||||
|
}
|
||||||
|
|
||||||
|
// coInitialize prepares the calling thread for COM and reports whether this
|
||||||
|
// call owns the initialization, which decides whether it may be balanced with
|
||||||
|
// CoUninitialize. S_FALSE took a reference on a thread this process had already
|
||||||
|
// initialized and so has to be released, while RPC_E_CHANGED_MODE took none:
|
||||||
|
// the thread belongs to another apartment, which is usable but is not ours to
|
||||||
|
// uninitialize.
|
||||||
|
func coInitialize() (bool, error) {
|
||||||
|
err := ole.CoInitializeEx(0, ole.COINIT_MULTITHREADED)
|
||||||
|
if err == nil {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var oleErr *ole.OleError
|
||||||
|
if errors.As(err, &oleErr) {
|
||||||
|
switch oleErr.Code() {
|
||||||
|
case sFalse:
|
||||||
|
return true, nil
|
||||||
|
case rpcEChangedMode:
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return false, fmt.Errorf("initialize COM: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// dispatchCall calls a COM method that returns an object.
|
||||||
|
func dispatchCall(dispatch *ole.IDispatch, method string, params ...any) (*ole.IDispatch, error) {
|
||||||
|
variant, err := oleutil.CallMethod(dispatch, method, params...)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
object := variant.ToIDispatch()
|
||||||
|
if object == nil {
|
||||||
|
return nil, fmt.Errorf("%s returned no object", method)
|
||||||
|
}
|
||||||
|
|
||||||
|
return object, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// dispatchProperty reads a COM property that holds an object.
|
||||||
|
func dispatchProperty(dispatch *ole.IDispatch, property string) (*ole.IDispatch, error) {
|
||||||
|
variant, err := oleutil.GetProperty(dispatch, property)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
object := variant.ToIDispatch()
|
||||||
|
if object == nil {
|
||||||
|
return nil, fmt.Errorf("property %s holds no object", property)
|
||||||
|
}
|
||||||
|
|
||||||
|
return object, nil
|
||||||
|
}
|
||||||
@@ -31,10 +31,30 @@ var (
|
|||||||
dnsFlushResolverCacheFn = dnsapi.NewProc("DnsFlushResolverCache")
|
dnsFlushResolverCacheFn = dnsapi.NewProc("DnsFlushResolverCache")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Registry locations of the host DNS configuration this package programs,
|
||||||
|
// exported so a diagnostic reader reports the same locations that are written.
|
||||||
const (
|
const (
|
||||||
dnsPolicyConfigMatchPath = `SYSTEM\CurrentControlSet\Services\Dnscache\Parameters\DnsPolicyConfig\NetBird-Match`
|
// NRPTKeyPrefix starts the name of every NRPT rule key this client creates.
|
||||||
gpoDnsPolicyRoot = `SOFTWARE\Policies\Microsoft\Windows NT\DNSClient\DnsPolicyConfig`
|
// Older versions used different layouts under the same prefix: a single
|
||||||
gpoDnsPolicyConfigMatchPath = gpoDnsPolicyRoot + `\NetBird-Match`
|
// unsuffixed key, then one key per domain, now one key per batch of domains.
|
||||||
|
NRPTKeyPrefix = "NetBird-Match"
|
||||||
|
|
||||||
|
// DNSPolicyConfigRoot holds the NRPT rules of the local policy store.
|
||||||
|
DNSPolicyConfigRoot = `SYSTEM\CurrentControlSet\Services\Dnscache\Parameters\DnsPolicyConfig`
|
||||||
|
|
||||||
|
// GPODNSPolicyConfigRoot holds the NRPT rules of the group policy store,
|
||||||
|
// which takes precedence over the local one when it is present.
|
||||||
|
GPODNSPolicyConfigRoot = `SOFTWARE\Policies\Microsoft\Windows NT\DNSClient\DnsPolicyConfig`
|
||||||
|
|
||||||
|
// InterfaceConfigPath and InterfaceConfigPathV6 hold the per-interface DNS
|
||||||
|
// settings, keyed by interface GUID, in separate hives per address family.
|
||||||
|
InterfaceConfigPath = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces`
|
||||||
|
InterfaceConfigPathV6 = `SYSTEM\CurrentControlSet\Services\Tcpip6\Parameters\Interfaces`
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
dnsPolicyConfigMatchPath = DNSPolicyConfigRoot + `\` + NRPTKeyPrefix
|
||||||
|
gpoDnsPolicyConfigMatchPath = GPODNSPolicyConfigRoot + `\` + NRPTKeyPrefix
|
||||||
|
|
||||||
dnsPolicyConfigVersionKey = "Version"
|
dnsPolicyConfigVersionKey = "Version"
|
||||||
dnsPolicyConfigVersionValue = 2
|
dnsPolicyConfigVersionValue = 2
|
||||||
@@ -45,8 +65,6 @@ const (
|
|||||||
|
|
||||||
nrptMaxDomainsPerRule = 50
|
nrptMaxDomainsPerRule = 50
|
||||||
|
|
||||||
interfaceConfigPath = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces`
|
|
||||||
interfaceConfigPathV6 = `SYSTEM\CurrentControlSet\Services\Tcpip6\Parameters\Interfaces`
|
|
||||||
interfaceConfigNameServerKey = "NameServer"
|
interfaceConfigNameServerKey = "NameServer"
|
||||||
interfaceConfigDhcpNameSrvKey = "DhcpNameServer"
|
interfaceConfigDhcpNameSrvKey = "DhcpNameServer"
|
||||||
interfaceConfigSearchListKey = "SearchList"
|
interfaceConfigSearchListKey = "SearchList"
|
||||||
@@ -73,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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -84,7 +101,7 @@ func newHostManager(wgInterface WGIface) (*registryConfigurator, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var useGPO bool
|
var useGPO bool
|
||||||
k, err := registry.OpenKey(registry.LOCAL_MACHINE, gpoDnsPolicyRoot, registry.QUERY_VALUE)
|
k, err := registry.OpenKey(registry.LOCAL_MACHINE, GPODNSPolicyConfigRoot, registry.QUERY_VALUE)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Debugf("failed to open GPO DNS policy root: %v", err)
|
log.Debugf("failed to open GPO DNS policy root: %v", err)
|
||||||
} else {
|
} else {
|
||||||
@@ -123,7 +140,7 @@ func (r *registryConfigurator) captureOriginalNameservers() ([]netip.Addr, error
|
|||||||
seen := make(map[netip.Addr]struct{})
|
seen := make(map[netip.Addr]struct{})
|
||||||
var out []netip.Addr
|
var out []netip.Addr
|
||||||
var merr *multierror.Error
|
var merr *multierror.Error
|
||||||
for _, root := range []string{interfaceConfigPath, interfaceConfigPathV6} {
|
for _, root := range []string{InterfaceConfigPath, InterfaceConfigPathV6} {
|
||||||
addrs, err := r.captureFromTcpipRoot(root)
|
addrs, err := r.captureFromTcpipRoot(root)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
merr = multierror.Append(merr, fmt.Errorf("%s: %w", root, err))
|
merr = multierror.Append(merr, fmt.Errorf("%s: %w", root, err))
|
||||||
@@ -306,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)
|
||||||
@@ -329,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)
|
||||||
}
|
}
|
||||||
@@ -346,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
|
||||||
|
|
||||||
@@ -363,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 {
|
||||||
@@ -385,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 {
|
||||||
@@ -450,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
|
||||||
}
|
}
|
||||||
@@ -496,7 +505,7 @@ func (r *registryConfigurator) deleteInterfaceRegistryKeyProperty(propertyKey st
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *registryConfigurator) getInterfaceRegistryKey() (registry.Key, error) {
|
func (r *registryConfigurator) getInterfaceRegistryKey() (registry.Key, error) {
|
||||||
regKeyPath := interfaceConfigPath + "\\" + r.guid
|
regKeyPath := InterfaceConfigPath + "\\" + r.guid
|
||||||
regKey, err := registry.OpenKey(registry.LOCAL_MACHINE, regKeyPath, registry.SET_VALUE)
|
regKey, err := registry.OpenKey(registry.LOCAL_MACHINE, regKeyPath, registry.SET_VALUE)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return regKey, fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", regKeyPath, err)
|
return regKey, fmt.Errorf("open HKEY_LOCAL_MACHINE\\%s: %w", regKeyPath, err)
|
||||||
@@ -518,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))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -554,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 {
|
||||||
@@ -585,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")
|
||||||
|
|||||||
@@ -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 := ®istryConfigurator{}
|
||||||
|
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 := ®istryConfigurator{}
|
||||||
|
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 := ®istryConfigurator{}
|
||||||
cfg := ®istryConfigurator{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
|
||||||
|
|||||||
@@ -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 := ®istryConfigurator{
|
manager := ®istryConfigurator{
|
||||||
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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -59,6 +59,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/netstate"
|
||||||
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"
|
||||||
@@ -181,6 +182,9 @@ type EngineServices struct {
|
|||||||
UpdateManager *updater.Manager
|
UpdateManager *updater.Manager
|
||||||
ClientMetrics *metrics.ClientMetrics
|
ClientMetrics *metrics.ClientMetrics
|
||||||
MetricsCtx context.Context
|
MetricsCtx context.Context
|
||||||
|
// NetState gates the reconnection loops on OS-reported network
|
||||||
|
// availability; nil disables gating.
|
||||||
|
NetState *netstate.State
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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.
|
||||||
@@ -204,6 +208,10 @@ type Engine struct {
|
|||||||
config *EngineConfig
|
config *EngineConfig
|
||||||
mobileDep MobileDependency
|
mobileDep MobileDependency
|
||||||
|
|
||||||
|
// netState gates the peer reconnection guards on OS-reported network
|
||||||
|
// availability; nil disables gating.
|
||||||
|
netState *netstate.State
|
||||||
|
|
||||||
// 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
|
||||||
@@ -337,6 +345,7 @@ func NewEngine(
|
|||||||
syncMsgMux: &sync.Mutex{},
|
syncMsgMux: &sync.Mutex{},
|
||||||
config: config,
|
config: config,
|
||||||
mobileDep: mobileDep,
|
mobileDep: mobileDep,
|
||||||
|
netState: services.NetState,
|
||||||
STUNs: []*stun.URI{},
|
STUNs: []*stun.URI{},
|
||||||
TURNs: []*stun.URI{},
|
TURNs: []*stun.URI{},
|
||||||
networkSerial: 0,
|
networkSerial: 0,
|
||||||
@@ -1893,7 +1902,8 @@ func (e *Engine) createPeerConn(pubKey string, allowedIPs []netip.Prefix, agentV
|
|||||||
Addr: e.getRosenpassAddr(),
|
Addr: e.getRosenpassAddr(),
|
||||||
PermissiveMode: e.config.RosenpassPermissive,
|
PermissiveMode: e.config.RosenpassPermissive,
|
||||||
},
|
},
|
||||||
ICEConfig: e.createICEConfig(),
|
ICEConfig: e.createICEConfig(),
|
||||||
|
NetworkState: e.netState,
|
||||||
}
|
}
|
||||||
|
|
||||||
serviceDependencies := peer.ServiceDependencies{
|
serviceDependencies := peer.ServiceDependencies{
|
||||||
@@ -2562,7 +2572,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
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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/netstate"
|
||||||
"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
|
||||||
|
|
||||||
|
// NetworkState gates the reconnection guard on OS-reported network
|
||||||
|
// availability; nil disables gating.
|
||||||
|
NetworkState *netstate.State
|
||||||
}
|
}
|
||||||
|
|
||||||
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.NetworkState)
|
||||||
|
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ import (
|
|||||||
|
|
||||||
"github.com/cenkalti/backoff/v4"
|
"github.com/cenkalti/backoff/v4"
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/netstate"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConnStatus represents the connection state as seen by the guard.
|
// ConnStatus represents the connection state as seen by the guard.
|
||||||
@@ -31,20 +33,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
|
||||||
|
// netState gates reconnect attempts on OS-reported network availability;
|
||||||
|
// nil disables gating.
|
||||||
|
netState *netstate.State
|
||||||
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 netState
|
||||||
|
// disables network availability gating.
|
||||||
|
func NewGuard(log *log.Entry, isConnectedFn connStatusFunc, timeout time.Duration, srWatcher *SRWatcher, netState *netstate.State) *Guard {
|
||||||
return &Guard{
|
return &Guard{
|
||||||
log: log,
|
log: log,
|
||||||
isConnectedOnAllWay: isConnectedFn,
|
isConnectedOnAllWay: isConnectedFn,
|
||||||
timeout: timeout,
|
timeout: timeout,
|
||||||
srWatcher: srWatcher,
|
srWatcher: srWatcher,
|
||||||
|
netState: netState,
|
||||||
relayedConnDisconnected: make(chan struct{}, 1),
|
relayedConnDisconnected: make(chan struct{}, 1),
|
||||||
iCEConnDisconnected: make(chan struct{}, 1),
|
iCEConnDisconnected: make(chan struct{}, 1),
|
||||||
}
|
}
|
||||||
@@ -96,9 +104,16 @@ 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()
|
||||||
|
|
||||||
|
netChanged := g.netState.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.netState.IsOnline() {
|
||||||
|
continue
|
||||||
|
}
|
||||||
switch g.isConnectedOnAllWay() {
|
switch g.isConnectedOnAllWay() {
|
||||||
case ConnStatusConnected:
|
case ConnStatusConnected:
|
||||||
// all good, nothing to do
|
// all good, nothing to do
|
||||||
@@ -135,6 +150,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.netState.Changed()
|
||||||
|
if !g.netState.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/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
|
||||||
|
}
|
||||||
@@ -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:
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 ¬ifier{}
|
return ¬ifier{
|
||||||
|
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())
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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.
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,408 +0,0 @@
|
|||||||
package pcp
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"crypto/rand"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"net/netip"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
defaultTimeout = 3 * time.Second
|
|
||||||
responseBufferSize = 128
|
|
||||||
|
|
||||||
// RFC 6887 Section 8.1.1 retry timing
|
|
||||||
initialRetryDelay = 3 * time.Second
|
|
||||||
maxRetryDelay = 1024 * time.Second
|
|
||||||
maxRetries = 4 // 3s + 6s + 12s + 24s = 45s total worst case
|
|
||||||
)
|
|
||||||
|
|
||||||
// Client is a PCP protocol client.
|
|
||||||
// All methods are safe for concurrent use.
|
|
||||||
type Client struct {
|
|
||||||
gateway netip.Addr
|
|
||||||
timeout time.Duration
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
// localIP caches the resolved local IP address.
|
|
||||||
localIP netip.Addr
|
|
||||||
// lastEpoch is the last observed server epoch value.
|
|
||||||
lastEpoch uint32
|
|
||||||
// epochTime tracks when lastEpoch was received for state loss detection.
|
|
||||||
epochTime time.Time
|
|
||||||
// externalIP caches the external IP from the last successful MAP response.
|
|
||||||
externalIP netip.Addr
|
|
||||||
// epochStateLost is set when epoch indicates server restart.
|
|
||||||
epochStateLost bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewClient creates a new PCP client for the gateway at the given IP.
|
|
||||||
func NewClient(gateway net.IP) *Client {
|
|
||||||
addr, ok := netip.AddrFromSlice(gateway)
|
|
||||||
if !ok {
|
|
||||||
log.Debugf("invalid gateway IP: %v", gateway)
|
|
||||||
}
|
|
||||||
return &Client{
|
|
||||||
gateway: addr.Unmap(),
|
|
||||||
timeout: defaultTimeout,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewClientWithTimeout creates a new PCP client with a custom timeout.
|
|
||||||
func NewClientWithTimeout(gateway net.IP, timeout time.Duration) *Client {
|
|
||||||
addr, ok := netip.AddrFromSlice(gateway)
|
|
||||||
if !ok {
|
|
||||||
log.Debugf("invalid gateway IP: %v", gateway)
|
|
||||||
}
|
|
||||||
return &Client{
|
|
||||||
gateway: addr.Unmap(),
|
|
||||||
timeout: timeout,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetLocalIP sets the local IP address to use in PCP requests.
|
|
||||||
func (c *Client) SetLocalIP(ip net.IP) {
|
|
||||||
addr, ok := netip.AddrFromSlice(ip)
|
|
||||||
if !ok {
|
|
||||||
log.Debugf("invalid local IP: %v", ip)
|
|
||||||
}
|
|
||||||
c.mu.Lock()
|
|
||||||
c.localIP = addr.Unmap()
|
|
||||||
c.mu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Gateway returns the gateway IP address.
|
|
||||||
func (c *Client) Gateway() net.IP {
|
|
||||||
return c.gateway.AsSlice()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Announce sends a PCP ANNOUNCE request to discover PCP support.
|
|
||||||
// Returns the server's epoch time on success.
|
|
||||||
func (c *Client) Announce(ctx context.Context) (epoch uint32, err error) {
|
|
||||||
localIP, err := c.getLocalIP()
|
|
||||||
if err != nil {
|
|
||||||
return 0, fmt.Errorf("get local IP: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
req := buildAnnounceRequest(localIP)
|
|
||||||
resp, err := c.sendRequest(ctx, req)
|
|
||||||
if err != nil {
|
|
||||||
return 0, fmt.Errorf("send announce: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
parsed, err := parseResponse(resp)
|
|
||||||
if err != nil {
|
|
||||||
return 0, fmt.Errorf("parse announce response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if parsed.ResultCode != ResultSuccess {
|
|
||||||
return 0, fmt.Errorf("PCP ANNOUNCE failed: %s", ResultCodeString(parsed.ResultCode))
|
|
||||||
}
|
|
||||||
|
|
||||||
c.mu.Lock()
|
|
||||||
if c.updateEpochLocked(parsed.Epoch) {
|
|
||||||
log.Warnf("PCP server epoch indicates state loss - mappings may need refresh")
|
|
||||||
}
|
|
||||||
c.mu.Unlock()
|
|
||||||
return parsed.Epoch, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddPortMapping requests a port mapping from the PCP server.
|
|
||||||
func (c *Client) AddPortMapping(ctx context.Context, protocol string, internalPort int, lifetime time.Duration) (*MapResponse, error) {
|
|
||||||
return c.addPortMappingWithHint(ctx, protocol, internalPort, internalPort, netip.Addr{}, lifetime)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddPortMappingWithHint requests a port mapping with suggested external port and IP.
|
|
||||||
// Use lifetime <= 0 to delete a mapping.
|
|
||||||
func (c *Client) AddPortMappingWithHint(ctx context.Context, protocol string, internalPort, suggestedExtPort int, suggestedExtIP net.IP, lifetime time.Duration) (*MapResponse, error) {
|
|
||||||
var extIP netip.Addr
|
|
||||||
if suggestedExtIP != nil {
|
|
||||||
var ok bool
|
|
||||||
extIP, ok = netip.AddrFromSlice(suggestedExtIP)
|
|
||||||
if !ok {
|
|
||||||
log.Debugf("invalid suggested external IP: %v", suggestedExtIP)
|
|
||||||
}
|
|
||||||
extIP = extIP.Unmap()
|
|
||||||
}
|
|
||||||
return c.addPortMappingWithHint(ctx, protocol, internalPort, suggestedExtPort, extIP, lifetime)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Client) addPortMappingWithHint(ctx context.Context, protocol string, internalPort, suggestedExtPort int, suggestedExtIP netip.Addr, lifetime time.Duration) (*MapResponse, error) {
|
|
||||||
localIP, err := c.getLocalIP()
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("get local IP: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
proto, err := protocolNumber(protocol)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("parse protocol: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var nonce [12]byte
|
|
||||||
if _, err := rand.Read(nonce[:]); err != nil {
|
|
||||||
return nil, fmt.Errorf("generate nonce: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Convert lifetime to seconds. Lifetime 0 means delete, so only apply
|
|
||||||
// default for positive durations that round to 0 seconds.
|
|
||||||
var lifetimeSec uint32
|
|
||||||
if lifetime > 0 {
|
|
||||||
lifetimeSec = uint32(lifetime.Seconds())
|
|
||||||
if lifetimeSec == 0 {
|
|
||||||
lifetimeSec = DefaultLifetime
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
req := buildMapRequest(localIP, nonce, proto, uint16(internalPort), uint16(suggestedExtPort), suggestedExtIP, lifetimeSec)
|
|
||||||
|
|
||||||
resp, err := c.sendRequest(ctx, req)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("send map request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
mapResp, err := parseMapResponse(resp)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("parse map response: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if mapResp.Nonce != nonce {
|
|
||||||
return nil, fmt.Errorf("nonce mismatch in response")
|
|
||||||
}
|
|
||||||
|
|
||||||
if mapResp.Protocol != proto {
|
|
||||||
return nil, fmt.Errorf("protocol mismatch: requested %d, got %d", proto, mapResp.Protocol)
|
|
||||||
}
|
|
||||||
if mapResp.InternalPort != uint16(internalPort) {
|
|
||||||
return nil, fmt.Errorf("internal port mismatch: requested %d, got %d", internalPort, mapResp.InternalPort)
|
|
||||||
}
|
|
||||||
|
|
||||||
if mapResp.ResultCode != ResultSuccess {
|
|
||||||
return nil, &Error{
|
|
||||||
Code: mapResp.ResultCode,
|
|
||||||
Message: ResultCodeString(mapResp.ResultCode),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
c.mu.Lock()
|
|
||||||
if c.updateEpochLocked(mapResp.Epoch) {
|
|
||||||
log.Warnf("PCP server epoch indicates state loss - mappings may need refresh")
|
|
||||||
}
|
|
||||||
c.cacheExternalIPLocked(mapResp.ExternalIP)
|
|
||||||
c.mu.Unlock()
|
|
||||||
return mapResp, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeletePortMapping removes a port mapping by requesting zero lifetime.
|
|
||||||
func (c *Client) DeletePortMapping(ctx context.Context, protocol string, internalPort int) error {
|
|
||||||
if _, err := c.addPortMappingWithHint(ctx, protocol, internalPort, 0, netip.Addr{}, 0); err != nil {
|
|
||||||
var pcpErr *Error
|
|
||||||
if errors.As(err, &pcpErr) && pcpErr.Code == ResultNotAuthorized {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return fmt.Errorf("delete mapping: %w", err)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetExternalAddress returns the external IP address.
|
|
||||||
// First checks for a cached value from previous MAP responses.
|
|
||||||
// If not cached, creates a short-lived mapping to discover the external IP.
|
|
||||||
func (c *Client) GetExternalAddress(ctx context.Context) (net.IP, error) {
|
|
||||||
c.mu.Lock()
|
|
||||||
if c.externalIP.IsValid() {
|
|
||||||
ip := c.externalIP.AsSlice()
|
|
||||||
c.mu.Unlock()
|
|
||||||
return ip, nil
|
|
||||||
}
|
|
||||||
c.mu.Unlock()
|
|
||||||
|
|
||||||
// Use an ephemeral port in the dynamic range (49152-65535).
|
|
||||||
// Port 0 is not valid with UDP/TCP protocols per RFC 6887.
|
|
||||||
ephemeralPort := 49152 + int(uint16(time.Now().UnixNano()))%(65535-49152)
|
|
||||||
|
|
||||||
// Use minimal lifetime (1 second) for discovery.
|
|
||||||
resp, err := c.AddPortMapping(ctx, "udp", ephemeralPort, time.Second)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("create temporary mapping: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := c.DeletePortMapping(ctx, "udp", ephemeralPort); err != nil {
|
|
||||||
log.Debugf("cleanup temporary PCP mapping: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return resp.ExternalIP.AsSlice(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// LastEpoch returns the last observed server epoch value.
|
|
||||||
// A decrease in epoch indicates the server may have restarted and mappings may be lost.
|
|
||||||
func (c *Client) LastEpoch() uint32 {
|
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
return c.lastEpoch
|
|
||||||
}
|
|
||||||
|
|
||||||
// EpochStateLost returns true if epoch state loss was detected and clears the flag.
|
|
||||||
func (c *Client) EpochStateLost() bool {
|
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
lost := c.epochStateLost
|
|
||||||
c.epochStateLost = false
|
|
||||||
return lost
|
|
||||||
}
|
|
||||||
|
|
||||||
// updateEpoch updates the epoch tracking and detects potential state loss.
|
|
||||||
// Returns true if state loss was detected (server likely restarted).
|
|
||||||
// Caller must hold c.mu.
|
|
||||||
func (c *Client) updateEpochLocked(newEpoch uint32) bool {
|
|
||||||
now := time.Now()
|
|
||||||
stateLost := false
|
|
||||||
|
|
||||||
// RFC 6887 Section 8.5: Detect invalid epoch indicating server state loss.
|
|
||||||
// client_delta = time since last response
|
|
||||||
// server_delta = epoch change since last response
|
|
||||||
// Invalid if: client_delta+2 < server_delta - server_delta/16
|
|
||||||
// OR: server_delta+2 < client_delta - client_delta/16
|
|
||||||
// The +2 handles quantization, /16 (6.25%) handles clock drift.
|
|
||||||
if !c.epochTime.IsZero() && c.lastEpoch > 0 {
|
|
||||||
clientDelta := uint32(now.Sub(c.epochTime).Seconds())
|
|
||||||
serverDelta := newEpoch - c.lastEpoch
|
|
||||||
|
|
||||||
// Check for epoch going backwards or jumping unexpectedly.
|
|
||||||
// Subtraction is safe: serverDelta/16 is always <= serverDelta.
|
|
||||||
if clientDelta+2 < serverDelta-(serverDelta/16) ||
|
|
||||||
serverDelta+2 < clientDelta-(clientDelta/16) {
|
|
||||||
stateLost = true
|
|
||||||
c.epochStateLost = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
c.lastEpoch = newEpoch
|
|
||||||
c.epochTime = now
|
|
||||||
return stateLost
|
|
||||||
}
|
|
||||||
|
|
||||||
// cacheExternalIP stores the external IP from a successful MAP response.
|
|
||||||
// Caller must hold c.mu.
|
|
||||||
func (c *Client) cacheExternalIPLocked(ip netip.Addr) {
|
|
||||||
if ip.IsValid() && !ip.IsUnspecified() {
|
|
||||||
c.externalIP = ip
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendRequest sends a PCP request with retries per RFC 6887 Section 8.1.1.
|
|
||||||
func (c *Client) sendRequest(ctx context.Context, req []byte) ([]byte, error) {
|
|
||||||
addr := &net.UDPAddr{IP: c.gateway.AsSlice(), Port: Port}
|
|
||||||
|
|
||||||
var lastErr error
|
|
||||||
delay := initialRetryDelay
|
|
||||||
|
|
||||||
for range maxRetries {
|
|
||||||
resp, err := c.sendOnce(ctx, addr, req)
|
|
||||||
if err == nil {
|
|
||||||
return resp, nil
|
|
||||||
}
|
|
||||||
lastErr = err
|
|
||||||
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
return nil, ctx.Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
// RFC 6887 Section 8.1.1: RT = (1 + RAND) * MIN(2 * RTprev, MRT)
|
|
||||||
// RAND is random between -0.1 and +0.1
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return nil, ctx.Err()
|
|
||||||
case <-time.After(retryDelayWithJitter(delay)):
|
|
||||||
}
|
|
||||||
delay = min(delay*2, maxRetryDelay)
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil, fmt.Errorf("PCP request failed after %d retries: %w", maxRetries, lastErr)
|
|
||||||
}
|
|
||||||
|
|
||||||
// retryDelayWithJitter applies RFC 6887 jitter: multiply by (1 + RAND) where RAND is [-0.1, +0.1].
|
|
||||||
func retryDelayWithJitter(d time.Duration) time.Duration {
|
|
||||||
var b [1]byte
|
|
||||||
_, _ = rand.Read(b[:])
|
|
||||||
// Convert byte to range [-0.1, +0.1]: (b/255 * 0.2) - 0.1
|
|
||||||
jitter := (float64(b[0])/255.0)*0.2 - 0.1
|
|
||||||
return time.Duration(float64(d) * (1 + jitter))
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Client) sendOnce(ctx context.Context, addr *net.UDPAddr, req []byte) ([]byte, error) {
|
|
||||||
// Use ListenUDP instead of DialUDP to validate response source address per RFC 6887 §8.3.
|
|
||||||
conn, err := net.ListenUDP("udp", nil)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("listen: %w", err)
|
|
||||||
}
|
|
||||||
defer func() {
|
|
||||||
if err := conn.Close(); err != nil {
|
|
||||||
log.Debugf("close UDP connection: %v", err)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
timeout := c.timeout
|
|
||||||
if deadline, ok := ctx.Deadline(); ok {
|
|
||||||
if remaining := time.Until(deadline); remaining < timeout {
|
|
||||||
timeout = remaining
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := conn.SetDeadline(time.Now().Add(timeout)); err != nil {
|
|
||||||
return nil, fmt.Errorf("set deadline: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := conn.WriteToUDP(req, addr); err != nil {
|
|
||||||
return nil, fmt.Errorf("write: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
resp := make([]byte, responseBufferSize)
|
|
||||||
n, from, err := conn.ReadFromUDP(resp)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("read: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// RFC 6887 §8.3: Validate response came from expected PCP server.
|
|
||||||
if !from.IP.Equal(addr.IP) {
|
|
||||||
return nil, fmt.Errorf("response from unexpected source %s (expected %s)", from.IP, addr.IP)
|
|
||||||
}
|
|
||||||
|
|
||||||
return resp[:n], nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Client) getLocalIP() (netip.Addr, error) {
|
|
||||||
c.mu.Lock()
|
|
||||||
defer c.mu.Unlock()
|
|
||||||
|
|
||||||
if !c.localIP.IsValid() {
|
|
||||||
return netip.Addr{}, fmt.Errorf("local IP not set for gateway %s", c.gateway)
|
|
||||||
}
|
|
||||||
return c.localIP, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func protocolNumber(protocol string) (uint8, error) {
|
|
||||||
switch protocol {
|
|
||||||
case "udp", "UDP":
|
|
||||||
return ProtoUDP, nil
|
|
||||||
case "tcp", "TCP":
|
|
||||||
return ProtoTCP, nil
|
|
||||||
default:
|
|
||||||
return 0, fmt.Errorf("unsupported protocol: %s", protocol)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Error represents a PCP error response.
|
|
||||||
type Error struct {
|
|
||||||
Code uint8
|
|
||||||
Message string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *Error) Error() string {
|
|
||||||
return fmt.Sprintf("PCP error: %s (%d)", e.Message, e.Code)
|
|
||||||
}
|
|
||||||
@@ -1,187 +0,0 @@
|
|||||||
package pcp
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"net"
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestAddrConversion(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
addr netip.Addr
|
|
||||||
}{
|
|
||||||
{"IPv4", netip.MustParseAddr("192.168.1.100")},
|
|
||||||
{"IPv4 loopback", netip.MustParseAddr("127.0.0.1")},
|
|
||||||
{"IPv6", netip.MustParseAddr("2001:db8::1")},
|
|
||||||
{"IPv6 loopback", netip.MustParseAddr("::1")},
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
b16 := addrTo16(tt.addr)
|
|
||||||
|
|
||||||
recovered := addrFrom16(b16)
|
|
||||||
assert.Equal(t, tt.addr, recovered, "address should round-trip")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildAnnounceRequest(t *testing.T) {
|
|
||||||
clientIP := netip.MustParseAddr("192.168.1.100")
|
|
||||||
req := buildAnnounceRequest(clientIP)
|
|
||||||
|
|
||||||
require.Len(t, req, headerSize)
|
|
||||||
assert.Equal(t, byte(Version), req[0], "version")
|
|
||||||
assert.Equal(t, byte(OpAnnounce), req[1], "opcode")
|
|
||||||
|
|
||||||
// Check client IP is properly encoded as IPv4-mapped IPv6
|
|
||||||
assert.Equal(t, byte(0xff), req[18], "IPv4-mapped prefix byte 10")
|
|
||||||
assert.Equal(t, byte(0xff), req[19], "IPv4-mapped prefix byte 11")
|
|
||||||
assert.Equal(t, byte(192), req[20], "IP octet 1")
|
|
||||||
assert.Equal(t, byte(168), req[21], "IP octet 2")
|
|
||||||
assert.Equal(t, byte(1), req[22], "IP octet 3")
|
|
||||||
assert.Equal(t, byte(100), req[23], "IP octet 4")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildMapRequest(t *testing.T) {
|
|
||||||
clientIP := netip.MustParseAddr("192.168.1.100")
|
|
||||||
nonce := [12]byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12}
|
|
||||||
req := buildMapRequest(clientIP, nonce, ProtoUDP, 51820, 51820, netip.Addr{}, 3600)
|
|
||||||
|
|
||||||
require.Len(t, req, mapRequestSize)
|
|
||||||
assert.Equal(t, byte(Version), req[0], "version")
|
|
||||||
assert.Equal(t, byte(OpMap), req[1], "opcode")
|
|
||||||
|
|
||||||
// Lifetime at bytes 4-7
|
|
||||||
assert.Equal(t, uint32(3600), (uint32(req[4])<<24)|(uint32(req[5])<<16)|(uint32(req[6])<<8)|uint32(req[7]), "lifetime")
|
|
||||||
|
|
||||||
// Nonce at bytes 24-35
|
|
||||||
assert.Equal(t, nonce[:], req[24:36], "nonce")
|
|
||||||
|
|
||||||
// Protocol at byte 36
|
|
||||||
assert.Equal(t, byte(ProtoUDP), req[36], "protocol")
|
|
||||||
|
|
||||||
// Internal port at bytes 40-41
|
|
||||||
assert.Equal(t, uint16(51820), (uint16(req[40])<<8)|uint16(req[41]), "internal port")
|
|
||||||
|
|
||||||
// External port at bytes 42-43
|
|
||||||
assert.Equal(t, uint16(51820), (uint16(req[42])<<8)|uint16(req[43]), "external port")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseResponse(t *testing.T) {
|
|
||||||
// Construct a valid ANNOUNCE response
|
|
||||||
resp := make([]byte, headerSize)
|
|
||||||
resp[0] = Version
|
|
||||||
resp[1] = OpAnnounce | OpReply
|
|
||||||
// Result code = 0 (success)
|
|
||||||
// Lifetime = 0
|
|
||||||
// Epoch = 12345
|
|
||||||
resp[8] = 0
|
|
||||||
resp[9] = 0
|
|
||||||
resp[10] = 0x30
|
|
||||||
resp[11] = 0x39
|
|
||||||
|
|
||||||
parsed, err := parseResponse(resp)
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint8(Version), parsed.Version)
|
|
||||||
assert.Equal(t, uint8(OpAnnounce|OpReply), parsed.Opcode)
|
|
||||||
assert.Equal(t, uint8(ResultSuccess), parsed.ResultCode)
|
|
||||||
assert.Equal(t, uint32(12345), parsed.Epoch)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestParseResponseErrors(t *testing.T) {
|
|
||||||
t.Run("too short", func(t *testing.T) {
|
|
||||||
_, err := parseResponse([]byte{1, 2, 3})
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("wrong version", func(t *testing.T) {
|
|
||||||
resp := make([]byte, headerSize)
|
|
||||||
resp[0] = 1 // Wrong version
|
|
||||||
resp[1] = OpReply
|
|
||||||
_, err := parseResponse(resp)
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("missing reply bit", func(t *testing.T) {
|
|
||||||
resp := make([]byte, headerSize)
|
|
||||||
resp[0] = Version
|
|
||||||
resp[1] = OpAnnounce // Missing OpReply bit
|
|
||||||
_, err := parseResponse(resp)
|
|
||||||
assert.Error(t, err)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestResultCodeString(t *testing.T) {
|
|
||||||
assert.Equal(t, "SUCCESS", ResultCodeString(ResultSuccess))
|
|
||||||
assert.Equal(t, "NOT_AUTHORIZED", ResultCodeString(ResultNotAuthorized))
|
|
||||||
assert.Equal(t, "ADDRESS_MISMATCH", ResultCodeString(ResultAddressMismatch))
|
|
||||||
assert.Contains(t, ResultCodeString(255), "UNKNOWN")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestProtocolNumber(t *testing.T) {
|
|
||||||
proto, err := protocolNumber("udp")
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint8(ProtoUDP), proto)
|
|
||||||
|
|
||||||
proto, err = protocolNumber("tcp")
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint8(ProtoTCP), proto)
|
|
||||||
|
|
||||||
proto, err = protocolNumber("UDP")
|
|
||||||
require.NoError(t, err)
|
|
||||||
assert.Equal(t, uint8(ProtoUDP), proto)
|
|
||||||
|
|
||||||
_, err = protocolNumber("icmp")
|
|
||||||
assert.Error(t, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClientCreation(t *testing.T) {
|
|
||||||
gateway := netip.MustParseAddr("192.168.1.1").AsSlice()
|
|
||||||
|
|
||||||
client := NewClient(gateway)
|
|
||||||
assert.Equal(t, net.IP(gateway), client.Gateway())
|
|
||||||
assert.Equal(t, defaultTimeout, client.timeout)
|
|
||||||
|
|
||||||
clientWithTimeout := NewClientWithTimeout(gateway, 5*time.Second)
|
|
||||||
assert.Equal(t, 5*time.Second, clientWithTimeout.timeout)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNATType(t *testing.T) {
|
|
||||||
n := NewNAT(netip.MustParseAddr("192.168.1.1").AsSlice(), netip.MustParseAddr("192.168.1.100").AsSlice())
|
|
||||||
assert.Equal(t, "PCP", n.Type())
|
|
||||||
}
|
|
||||||
|
|
||||||
// Integration test - skipped unless PCP_TEST_GATEWAY env is set
|
|
||||||
func TestClientIntegration(t *testing.T) {
|
|
||||||
t.Skip("Integration test - run manually with PCP_TEST_GATEWAY=<gateway-ip>")
|
|
||||||
|
|
||||||
gateway := netip.MustParseAddr("10.0.1.1").AsSlice() // Change to your test gateway
|
|
||||||
localIP := netip.MustParseAddr("10.0.1.100").AsSlice() // Change to your local IP
|
|
||||||
|
|
||||||
client := NewClient(gateway)
|
|
||||||
client.SetLocalIP(localIP)
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
// Test ANNOUNCE
|
|
||||||
epoch, err := client.Announce(ctx)
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Logf("Server epoch: %d", epoch)
|
|
||||||
|
|
||||||
// Test MAP
|
|
||||||
resp, err := client.AddPortMapping(ctx, "udp", 51820, 1*time.Hour)
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Logf("Mapping: internal=%d external=%d externalIP=%s",
|
|
||||||
resp.InternalPort, resp.ExternalPort, resp.ExternalIP)
|
|
||||||
|
|
||||||
// Cleanup
|
|
||||||
err = client.DeletePortMapping(ctx, "udp", 51820)
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
|
||||||
@@ -1,222 +0,0 @@
|
|||||||
package pcp
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"fmt"
|
|
||||||
"net"
|
|
||||||
"net/netip"
|
|
||||||
"runtime"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
|
||||||
|
|
||||||
"github.com/libp2p/go-nat"
|
|
||||||
"github.com/libp2p/go-netroute"
|
|
||||||
)
|
|
||||||
|
|
||||||
var _ nat.NAT = (*NAT)(nil)
|
|
||||||
|
|
||||||
// NAT implements the go-nat NAT interface using PCP.
|
|
||||||
// Supports dual-stack (IPv4 and IPv6) when available.
|
|
||||||
// All methods are safe for concurrent use.
|
|
||||||
//
|
|
||||||
// TODO: IPv6 pinholes use the local IPv6 address. If the address changes
|
|
||||||
// (e.g., due to SLAAC rotation or network change), the pinhole becomes stale
|
|
||||||
// and needs to be recreated with the new address.
|
|
||||||
type NAT struct {
|
|
||||||
client *Client
|
|
||||||
|
|
||||||
mu sync.RWMutex
|
|
||||||
// client6 is the IPv6 PCP client, nil if IPv6 is unavailable.
|
|
||||||
client6 *Client
|
|
||||||
// localIP6 caches the local IPv6 address used for PCP requests.
|
|
||||||
localIP6 netip.Addr
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewNAT creates a new NAT instance backed by PCP.
|
|
||||||
func NewNAT(gateway, localIP net.IP) *NAT {
|
|
||||||
client := NewClient(gateway)
|
|
||||||
client.SetLocalIP(localIP)
|
|
||||||
return &NAT{
|
|
||||||
client: client,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Type returns "PCP" as the NAT type.
|
|
||||||
func (n *NAT) Type() string {
|
|
||||||
return "PCP"
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetDeviceAddress returns the gateway IP address.
|
|
||||||
func (n *NAT) GetDeviceAddress() (net.IP, error) {
|
|
||||||
return n.client.Gateway(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetExternalAddress returns the external IP address.
|
|
||||||
func (n *NAT) GetExternalAddress() (net.IP, error) {
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
return n.client.GetExternalAddress(ctx)
|
|
||||||
}
|
|
||||||
|
|
||||||
// GetInternalAddress returns the local IP address used to communicate with the gateway.
|
|
||||||
func (n *NAT) GetInternalAddress() (net.IP, error) {
|
|
||||||
addr, err := n.client.getLocalIP()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return addr.AsSlice(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddPortMapping creates a port mapping on both IPv4 and IPv6 (if available).
|
|
||||||
func (n *NAT) AddPortMapping(ctx context.Context, protocol string, internalPort int, _ string, timeout time.Duration) (int, error) {
|
|
||||||
resp, err := n.client.AddPortMapping(ctx, protocol, internalPort, timeout)
|
|
||||||
if err != nil {
|
|
||||||
return 0, fmt.Errorf("add mapping: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
n.mu.RLock()
|
|
||||||
client6 := n.client6
|
|
||||||
localIP6 := n.localIP6
|
|
||||||
n.mu.RUnlock()
|
|
||||||
|
|
||||||
if client6 == nil {
|
|
||||||
return int(resp.ExternalPort), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, err := client6.AddPortMapping(ctx, protocol, internalPort, timeout); err != nil {
|
|
||||||
log.Warnf("IPv6 PCP mapping failed (continuing with IPv4): %v", err)
|
|
||||||
return int(resp.ExternalPort), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Infof("created IPv6 PCP pinhole: %s:%d", localIP6, internalPort)
|
|
||||||
return int(resp.ExternalPort), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// DeletePortMapping removes a port mapping from both IPv4 and IPv6.
|
|
||||||
func (n *NAT) DeletePortMapping(ctx context.Context, protocol string, internalPort int) error {
|
|
||||||
err := n.client.DeletePortMapping(ctx, protocol, internalPort)
|
|
||||||
|
|
||||||
n.mu.RLock()
|
|
||||||
client6 := n.client6
|
|
||||||
n.mu.RUnlock()
|
|
||||||
|
|
||||||
if client6 != nil {
|
|
||||||
if err6 := client6.DeletePortMapping(ctx, protocol, internalPort); err6 != nil {
|
|
||||||
log.Warnf("IPv6 PCP delete mapping failed: %v", err6)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("delete mapping: %w", err)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// CheckServerHealth sends an ANNOUNCE to verify the server is still responsive.
|
|
||||||
// Returns the current epoch and whether the server may have restarted (epoch state loss detected).
|
|
||||||
func (n *NAT) CheckServerHealth(ctx context.Context) (epoch uint32, serverRestarted bool, err error) {
|
|
||||||
epoch, err = n.client.Announce(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return 0, false, fmt.Errorf("announce: %w", err)
|
|
||||||
}
|
|
||||||
return epoch, n.client.EpochStateLost(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// DiscoverPCP attempts to discover a PCP-capable gateway.
|
|
||||||
// Returns a NAT interface if PCP is supported, or an error otherwise.
|
|
||||||
// Discovers both IPv4 and IPv6 gateways when available.
|
|
||||||
func DiscoverPCP(ctx context.Context) (nat.NAT, error) {
|
|
||||||
gateway, localIP, err := getDefaultGateway()
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("get default gateway: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
client := NewClient(gateway)
|
|
||||||
client.SetLocalIP(localIP)
|
|
||||||
if _, err := client.Announce(ctx); err != nil {
|
|
||||||
return nil, fmt.Errorf("PCP announce: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
result := &NAT{client: client}
|
|
||||||
discoverIPv6(ctx, result)
|
|
||||||
|
|
||||||
return result, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func discoverIPv6(ctx context.Context, result *NAT) {
|
|
||||||
gateway6, localIP6, err := getDefaultGateway6()
|
|
||||||
if err != nil {
|
|
||||||
log.Debugf("IPv6 gateway discovery failed: %v", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
client6 := NewClient(gateway6)
|
|
||||||
client6.SetLocalIP(localIP6)
|
|
||||||
if _, err := client6.Announce(ctx); err != nil {
|
|
||||||
log.Debugf("PCP IPv6 announce failed: %v", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
addr, ok := netip.AddrFromSlice(localIP6)
|
|
||||||
if !ok {
|
|
||||||
log.Debugf("invalid IPv6 local IP: %v", localIP6)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
result.mu.Lock()
|
|
||||||
result.client6 = client6
|
|
||||||
result.localIP6 = addr
|
|
||||||
result.mu.Unlock()
|
|
||||||
log.Debugf("PCP IPv6 gateway discovered: %s (local: %s)", gateway6, localIP6)
|
|
||||||
}
|
|
||||||
|
|
||||||
// getDefaultGateway returns the default IPv4 gateway and local IP using the system routing table.
|
|
||||||
func getDefaultGateway() (gateway net.IP, localIP net.IP, err error) {
|
|
||||||
router, err := netroute.New()
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
dst := net.IPv4zero
|
|
||||||
if runtime.GOOS == "linux" || runtime.GOOS == "android" {
|
|
||||||
// go-netroute v0.4.0 rejects unspecified destinations client-side on Linux/Android.
|
|
||||||
// TODO: on android/ios, use platform APIs (ConnectivityManager.getLinkProperties /
|
|
||||||
// NWPathMonitor) when netlink-based lookup is restricted or unavailable.
|
|
||||||
dst = net.IPv4(0, 0, 0, 1)
|
|
||||||
}
|
|
||||||
_, gateway, localIP, err = router.Route(dst)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if gateway == nil {
|
|
||||||
return nil, nil, nat.ErrNoNATFound
|
|
||||||
}
|
|
||||||
|
|
||||||
return gateway, localIP, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// getDefaultGateway6 returns the default IPv6 gateway IP address using the system routing table.
|
|
||||||
func getDefaultGateway6() (gateway net.IP, localIP net.IP, err error) {
|
|
||||||
router, err := netroute.New()
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
dst := net.IPv6zero
|
|
||||||
if runtime.GOOS == "linux" || runtime.GOOS == "android" {
|
|
||||||
// ::2
|
|
||||||
dst = net.IP{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}
|
|
||||||
}
|
|
||||||
_, gateway, localIP, err = router.Route(dst)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if gateway == nil {
|
|
||||||
return nil, nil, nat.ErrNoNATFound
|
|
||||||
}
|
|
||||||
|
|
||||||
return gateway, localIP, nil
|
|
||||||
}
|
|
||||||
@@ -1,225 +0,0 @@
|
|||||||
// Package pcp implements the Port Control Protocol (RFC 6887).
|
|
||||||
//
|
|
||||||
// # Implemented Features
|
|
||||||
//
|
|
||||||
// - ANNOUNCE opcode: Discovers PCP server support
|
|
||||||
// - MAP opcode: Creates/deletes port mappings (IPv4 NAT) and firewall pinholes (IPv6)
|
|
||||||
// - Dual-stack: Simultaneous IPv4 and IPv6 support via separate clients
|
|
||||||
// - Nonce validation: Prevents response spoofing
|
|
||||||
// - Epoch tracking: Detects server restarts per Section 8.5
|
|
||||||
// - RFC-compliant retry timing: 3s initial, exponential backoff to 1024s max (Section 8.1.1)
|
|
||||||
//
|
|
||||||
// # Not Implemented
|
|
||||||
//
|
|
||||||
// - PEER opcode: For outbound peer connections (not needed for inbound NAT traversal)
|
|
||||||
// - THIRD_PARTY option: For managing mappings on behalf of other devices
|
|
||||||
// - PREFER_FAILURE option: Requires exact external port or fail (IPv4 NAT only, not needed for IPv6 pinholing)
|
|
||||||
// - FILTER option: To restrict remote peer addresses
|
|
||||||
//
|
|
||||||
// These optional features are omitted because the primary use case is simple
|
|
||||||
// port forwarding for WireGuard, which only requires MAP with default behavior.
|
|
||||||
package pcp
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
|
||||||
"net/netip"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
// Version is the PCP protocol version (RFC 6887).
|
|
||||||
Version = 2
|
|
||||||
|
|
||||||
// Port is the standard PCP server port.
|
|
||||||
Port = 5351
|
|
||||||
|
|
||||||
// DefaultLifetime is the default requested mapping lifetime in seconds.
|
|
||||||
DefaultLifetime = 7200 // 2 hours
|
|
||||||
|
|
||||||
// Header sizes
|
|
||||||
headerSize = 24
|
|
||||||
mapPayloadSize = 36
|
|
||||||
mapRequestSize = headerSize + mapPayloadSize // 60 bytes
|
|
||||||
)
|
|
||||||
|
|
||||||
// Opcodes
|
|
||||||
const (
|
|
||||||
OpAnnounce = 0
|
|
||||||
OpMap = 1
|
|
||||||
OpPeer = 2
|
|
||||||
OpReply = 0x80 // OR'd with opcode in responses
|
|
||||||
)
|
|
||||||
|
|
||||||
// Protocol numbers for MAP requests
|
|
||||||
const (
|
|
||||||
ProtoUDP = 17
|
|
||||||
ProtoTCP = 6
|
|
||||||
)
|
|
||||||
|
|
||||||
// Result codes (RFC 6887 Section 7.4)
|
|
||||||
const (
|
|
||||||
ResultSuccess = 0
|
|
||||||
ResultUnsuppVersion = 1
|
|
||||||
ResultNotAuthorized = 2
|
|
||||||
ResultMalformedRequest = 3
|
|
||||||
ResultUnsuppOpcode = 4
|
|
||||||
ResultUnsuppOption = 5
|
|
||||||
ResultMalformedOption = 6
|
|
||||||
ResultNetworkFailure = 7
|
|
||||||
ResultNoResources = 8
|
|
||||||
ResultUnsuppProtocol = 9
|
|
||||||
ResultUserExQuota = 10
|
|
||||||
ResultCannotProvideExt = 11
|
|
||||||
ResultAddressMismatch = 12
|
|
||||||
ResultExcessiveRemotePeers = 13
|
|
||||||
)
|
|
||||||
|
|
||||||
// ResultCodeString returns a human-readable string for a result code.
|
|
||||||
func ResultCodeString(code uint8) string {
|
|
||||||
switch code {
|
|
||||||
case ResultSuccess:
|
|
||||||
return "SUCCESS"
|
|
||||||
case ResultUnsuppVersion:
|
|
||||||
return "UNSUPP_VERSION"
|
|
||||||
case ResultNotAuthorized:
|
|
||||||
return "NOT_AUTHORIZED"
|
|
||||||
case ResultMalformedRequest:
|
|
||||||
return "MALFORMED_REQUEST"
|
|
||||||
case ResultUnsuppOpcode:
|
|
||||||
return "UNSUPP_OPCODE"
|
|
||||||
case ResultUnsuppOption:
|
|
||||||
return "UNSUPP_OPTION"
|
|
||||||
case ResultMalformedOption:
|
|
||||||
return "MALFORMED_OPTION"
|
|
||||||
case ResultNetworkFailure:
|
|
||||||
return "NETWORK_FAILURE"
|
|
||||||
case ResultNoResources:
|
|
||||||
return "NO_RESOURCES"
|
|
||||||
case ResultUnsuppProtocol:
|
|
||||||
return "UNSUPP_PROTOCOL"
|
|
||||||
case ResultUserExQuota:
|
|
||||||
return "USER_EX_QUOTA"
|
|
||||||
case ResultCannotProvideExt:
|
|
||||||
return "CANNOT_PROVIDE_EXTERNAL"
|
|
||||||
case ResultAddressMismatch:
|
|
||||||
return "ADDRESS_MISMATCH"
|
|
||||||
case ResultExcessiveRemotePeers:
|
|
||||||
return "EXCESSIVE_REMOTE_PEERS"
|
|
||||||
default:
|
|
||||||
return fmt.Sprintf("UNKNOWN(%d)", code)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Response represents a parsed PCP response header.
|
|
||||||
type Response struct {
|
|
||||||
Version uint8
|
|
||||||
Opcode uint8
|
|
||||||
ResultCode uint8
|
|
||||||
Lifetime uint32
|
|
||||||
Epoch uint32
|
|
||||||
}
|
|
||||||
|
|
||||||
// MapResponse contains the full response to a MAP request.
|
|
||||||
type MapResponse struct {
|
|
||||||
Response
|
|
||||||
Nonce [12]byte
|
|
||||||
Protocol uint8
|
|
||||||
InternalPort uint16
|
|
||||||
ExternalPort uint16
|
|
||||||
ExternalIP netip.Addr
|
|
||||||
}
|
|
||||||
|
|
||||||
// addrTo16 converts an address to its 16-byte IPv4-mapped IPv6 representation.
|
|
||||||
func addrTo16(addr netip.Addr) [16]byte {
|
|
||||||
if addr.Is4() {
|
|
||||||
return netip.AddrFrom4(addr.As4()).As16()
|
|
||||||
}
|
|
||||||
return addr.As16()
|
|
||||||
}
|
|
||||||
|
|
||||||
// addrFrom16 extracts an address from a 16-byte representation, unmapping IPv4.
|
|
||||||
func addrFrom16(b [16]byte) netip.Addr {
|
|
||||||
return netip.AddrFrom16(b).Unmap()
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildAnnounceRequest creates a PCP ANNOUNCE request packet.
|
|
||||||
func buildAnnounceRequest(clientIP netip.Addr) []byte {
|
|
||||||
req := make([]byte, headerSize)
|
|
||||||
req[0] = Version
|
|
||||||
req[1] = OpAnnounce
|
|
||||||
mapped := addrTo16(clientIP)
|
|
||||||
copy(req[8:24], mapped[:])
|
|
||||||
return req
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildMapRequest creates a PCP MAP request packet.
|
|
||||||
func buildMapRequest(clientIP netip.Addr, nonce [12]byte, protocol uint8, internalPort, suggestedExtPort uint16, suggestedExtIP netip.Addr, lifetime uint32) []byte {
|
|
||||||
req := make([]byte, mapRequestSize)
|
|
||||||
|
|
||||||
// Header
|
|
||||||
req[0] = Version
|
|
||||||
req[1] = OpMap
|
|
||||||
binary.BigEndian.PutUint32(req[4:8], lifetime)
|
|
||||||
mapped := addrTo16(clientIP)
|
|
||||||
copy(req[8:24], mapped[:])
|
|
||||||
|
|
||||||
// MAP payload
|
|
||||||
copy(req[24:36], nonce[:])
|
|
||||||
req[36] = protocol
|
|
||||||
binary.BigEndian.PutUint16(req[40:42], internalPort)
|
|
||||||
binary.BigEndian.PutUint16(req[42:44], suggestedExtPort)
|
|
||||||
if suggestedExtIP.IsValid() {
|
|
||||||
extMapped := addrTo16(suggestedExtIP)
|
|
||||||
copy(req[44:60], extMapped[:])
|
|
||||||
}
|
|
||||||
|
|
||||||
return req
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseResponse parses the common PCP response header.
|
|
||||||
func parseResponse(data []byte) (*Response, error) {
|
|
||||||
if len(data) < headerSize {
|
|
||||||
return nil, fmt.Errorf("response too short: %d bytes", len(data))
|
|
||||||
}
|
|
||||||
|
|
||||||
resp := &Response{
|
|
||||||
Version: data[0],
|
|
||||||
Opcode: data[1],
|
|
||||||
ResultCode: data[3], // Byte 2 is reserved, byte 3 is result code (RFC 6887 §7.2)
|
|
||||||
Lifetime: binary.BigEndian.Uint32(data[4:8]),
|
|
||||||
Epoch: binary.BigEndian.Uint32(data[8:12]),
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.Version != Version {
|
|
||||||
return nil, fmt.Errorf("unsupported PCP version: %d", resp.Version)
|
|
||||||
}
|
|
||||||
|
|
||||||
if resp.Opcode&OpReply == 0 {
|
|
||||||
return nil, fmt.Errorf("response missing reply bit: opcode=0x%02x", resp.Opcode)
|
|
||||||
}
|
|
||||||
|
|
||||||
return resp, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseMapResponse parses a complete MAP response.
|
|
||||||
func parseMapResponse(data []byte) (*MapResponse, error) {
|
|
||||||
if len(data) < mapRequestSize {
|
|
||||||
return nil, fmt.Errorf("MAP response too short: %d bytes", len(data))
|
|
||||||
}
|
|
||||||
|
|
||||||
resp, err := parseResponse(data)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("parse header: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
mapResp := &MapResponse{
|
|
||||||
Response: *resp,
|
|
||||||
Protocol: data[36],
|
|
||||||
InternalPort: binary.BigEndian.Uint16(data[40:42]),
|
|
||||||
ExternalPort: binary.BigEndian.Uint16(data[42:44]),
|
|
||||||
ExternalIP: addrFrom16([16]byte(data[44:60])),
|
|
||||||
}
|
|
||||||
copy(mapResp.Nonce[:], data[24:36])
|
|
||||||
|
|
||||||
return mapResp, nil
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,116 @@
|
|||||||
|
//go:build !js
|
||||||
|
|
||||||
|
package portforward
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/netbirdio/go-nat"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"github.com/sirupsen/logrus/hooks/test"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// mockPinholeNAT is a gateway that also reports an IPv6 pinhole outcome, the
|
||||||
|
// shape a dual-stack gateway has.
|
||||||
|
type mockPinholeNAT struct {
|
||||||
|
*mockNAT
|
||||||
|
pinholeErr error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockPinholeNAT) IPv6PinholeError() error {
|
||||||
|
return m.pinholeErr
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSetupLogsPinholeOutcome(t *testing.T) {
|
||||||
|
pinholeErr := errors.New("pcp ipv6: NOT_AUTHORIZED")
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
pinholeErr error
|
||||||
|
mappingErr error
|
||||||
|
wantLevel log.Level
|
||||||
|
wantText string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "an open pinhole is reported",
|
||||||
|
wantLevel: log.InfoLevel,
|
||||||
|
wantText: "IPv6 pinhole open",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "a failed pinhole is reported without failing the mapping",
|
||||||
|
// The IPv4 mapping is what the caller asked for, so the pinhole
|
||||||
|
// failure surfaces only in the log.
|
||||||
|
pinholeErr: pinholeErr,
|
||||||
|
wantLevel: log.WarnLevel,
|
||||||
|
wantText: pinholeErr.Error(),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
gateway := &mockPinholeNAT{mockNAT: newMockNAT(), pinholeErr: tt.pinholeErr}
|
||||||
|
hook := stubGatewayDiscovery(t, gateway)
|
||||||
|
|
||||||
|
m := NewManager()
|
||||||
|
m.wgPort = 51820
|
||||||
|
|
||||||
|
_, mapping, err := m.setup(context.Background())
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, mapping)
|
||||||
|
|
||||||
|
entry := findEntry(hook, tt.wantText)
|
||||||
|
require.NotNil(t, entry, "no log entry mentioning %q", tt.wantText)
|
||||||
|
assert.Equal(t, tt.wantLevel, entry.Level)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("a failed mapping reports no pinhole outcome", func(t *testing.T) {
|
||||||
|
// Nothing opened the pinhole, so whatever it currently reports says
|
||||||
|
// nothing about this attempt.
|
||||||
|
gateway := &mockPinholeNAT{mockNAT: newMockNAT()}
|
||||||
|
gateway.addMappingErr = errors.New("gateway refused")
|
||||||
|
hook := stubGatewayDiscovery(t, gateway)
|
||||||
|
|
||||||
|
m := NewManager()
|
||||||
|
m.wgPort = 51820
|
||||||
|
|
||||||
|
_, _, err := m.setup(context.Background())
|
||||||
|
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Nil(t, findEntry(hook, "IPv6 pinhole"))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// stubGatewayDiscovery makes discovery return gateway and captures log output.
|
||||||
|
func stubGatewayDiscovery(t *testing.T, gateway nat.NAT) *test.Hook {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
orig := discoverGateway
|
||||||
|
discoverGateway = func(context.Context) (nat.NAT, error) { return gateway, nil }
|
||||||
|
t.Cleanup(func() { discoverGateway = orig })
|
||||||
|
|
||||||
|
hook := test.NewGlobal()
|
||||||
|
origLevel := log.GetLevel()
|
||||||
|
log.SetLevel(log.DebugLevel)
|
||||||
|
t.Cleanup(func() {
|
||||||
|
hook.Reset()
|
||||||
|
log.SetLevel(origLevel)
|
||||||
|
})
|
||||||
|
|
||||||
|
return hook
|
||||||
|
}
|
||||||
|
|
||||||
|
func findEntry(hook *test.Hook, substr string) *log.Entry {
|
||||||
|
for _, entry := range hook.AllEntries() {
|
||||||
|
if strings.Contains(entry.Message, substr) {
|
||||||
|
return entry
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -4,27 +4,94 @@ package portforward
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/libp2p/go-nat"
|
"github.com/netbirdio/go-nat"
|
||||||
|
"github.com/netbirdio/go-nat/pcp"
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/internal/portforward/pcp"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// discoverGateway is the function used for NAT gateway discovery.
|
// discoverGateway is the function used for NAT gateway discovery.
|
||||||
// It can be replaced in tests to avoid real network operations.
|
// It can be replaced in tests to avoid real network operations.
|
||||||
// Tries PCP first, then falls back to NAT-PMP/UPnP.
|
|
||||||
var discoverGateway = defaultDiscoverGateway
|
var discoverGateway = defaultDiscoverGateway
|
||||||
|
|
||||||
func defaultDiscoverGateway(ctx context.Context) (nat.NAT, error) {
|
// pinholeDiscoveryTimeout is the slice of the discovery budget held back for
|
||||||
pcpGateway, err := pcp.DiscoverPCP(ctx)
|
// the IPv6 pinhole probe.
|
||||||
if err == nil {
|
//
|
||||||
return pcpGateway, nil
|
// Sizing it is coarser than it looks: PCP retransmits on a 3s socket timeout
|
||||||
}
|
// and a 3s first backoff, so a second attempt needs about 9s. Anything from
|
||||||
log.Debugf("PCP discovery failed: %v, trying NAT-PMP/UPnP", err)
|
// roughly 1s to 8s therefore buys exactly one attempt, and this only sets how
|
||||||
|
// long that attempt waits. A PCP server sits on the local link and answers in
|
||||||
|
// milliseconds, so 3s is margin rather than need, and the rest is left to
|
||||||
|
// gateway discovery, whose multicast SSDP search alone takes 5s. A probe lost
|
||||||
|
// to a dropped packet is retried by the next discovery round.
|
||||||
|
//
|
||||||
|
// It is a variable so tests can shorten it.
|
||||||
|
var pinholeDiscoveryTimeout = 3 * time.Second
|
||||||
|
|
||||||
return nat.DiscoverGateway(ctx)
|
// Discovery entry points, as variables so tests can drive the fallback without
|
||||||
|
// touching the network.
|
||||||
|
var (
|
||||||
|
discoverNATGateway = nat.DiscoverGateway
|
||||||
|
|
||||||
|
discoverPCPPinhole = func(ctx context.Context) (nat.NAT, error) {
|
||||||
|
pinhole, err := pcp.DiscoverPCP(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return pinhole, nil
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
// defaultDiscoverGateway finds a gateway that can make the WireGuard port
|
||||||
|
// reachable. DiscoverGateway prefers PCP for IPv4, races UPnP and NAT-PMP
|
||||||
|
// behind it, and attaches an IPv6 pinhole independently of which IPv4 protocol
|
||||||
|
// wins.
|
||||||
|
//
|
||||||
|
// It reports no gateway on a network offering only IPv6, having no IPv4 mapping
|
||||||
|
// to attach a pinhole to. Such a network still needs one: there is no
|
||||||
|
// translation to traverse, but the router drops inbound IPv6 until something
|
||||||
|
// opens it. Fall back to PCP alone, which yields a gateway holding just the
|
||||||
|
// pinhole.
|
||||||
|
func defaultDiscoverGateway(ctx context.Context) (nat.NAT, error) {
|
||||||
|
gatewayCtx, cancel := reserveForPinhole(ctx)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
gateway, err := discoverNATGateway(gatewayCtx)
|
||||||
|
if err == nil {
|
||||||
|
return gateway, nil
|
||||||
|
}
|
||||||
|
if !errors.Is(err, nat.ErrNoNATFound) {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
pinhole, pinholeErr := discoverPCPPinhole(ctx)
|
||||||
|
if pinholeErr != nil {
|
||||||
|
log.Debugf("no IPv6 pinhole after %v: %v", err, pinholeErr)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infof("no IPv4 gateway, continuing with an IPv6 pinhole only")
|
||||||
|
return pinhole, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// reserveForPinhole shortens ctx so that a pinhole probe still has time to run
|
||||||
|
// afterwards. Finding nothing takes gateway discovery everything it is given,
|
||||||
|
// so on the unshortened context the probe would start already expired. A budget
|
||||||
|
// too small to divide is left to gateway discovery, which is the likelier win.
|
||||||
|
func reserveForPinhole(ctx context.Context) (context.Context, context.CancelFunc) {
|
||||||
|
deadline, ok := ctx.Deadline()
|
||||||
|
if !ok {
|
||||||
|
return context.WithCancel(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
remaining := time.Until(deadline)
|
||||||
|
if remaining <= pinholeDiscoveryTimeout {
|
||||||
|
return context.WithCancel(ctx)
|
||||||
|
}
|
||||||
|
return context.WithTimeout(ctx, remaining-pinholeDiscoveryTimeout)
|
||||||
}
|
}
|
||||||
|
|
||||||
// State is persisted only for crash recovery cleanup
|
// State is persisted only for crash recovery cleanup
|
||||||
|
|||||||
@@ -0,0 +1,140 @@
|
|||||||
|
//go:build !js
|
||||||
|
|
||||||
|
package portforward
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/netbirdio/go-nat"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// stubDiscovery replaces both discovery entry points for the duration of a
|
||||||
|
// test. gatewayDelay simulates gateway discovery spending everything it is
|
||||||
|
// given before reporting that it found nothing.
|
||||||
|
func stubDiscovery(t *testing.T, gateway nat.NAT, gatewayErr error, gatewayDelay time.Duration, pinhole nat.NAT, pinholeErr error) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
origGateway, origPinhole := discoverNATGateway, discoverPCPPinhole
|
||||||
|
discoverNATGateway = func(ctx context.Context) (nat.NAT, error) {
|
||||||
|
if gatewayDelay > 0 {
|
||||||
|
select {
|
||||||
|
case <-time.After(gatewayDelay):
|
||||||
|
case <-ctx.Done():
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return gateway, gatewayErr
|
||||||
|
}
|
||||||
|
discoverPCPPinhole = func(ctx context.Context) (nat.NAT, error) {
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return pinhole, pinholeErr
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Cleanup(func() { discoverNATGateway, discoverPCPPinhole = origGateway, origPinhole })
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDefaultDiscoverGateway(t *testing.T) {
|
||||||
|
ipv4Gateway := &mockNAT{natType: "PCP+PCPv6"}
|
||||||
|
ipv6Pinhole := &mockNAT{natType: "PCP"}
|
||||||
|
otherErr := errors.New("routing table unavailable")
|
||||||
|
|
||||||
|
t.Run("an IPv4 gateway is used as is", func(t *testing.T) {
|
||||||
|
stubDiscovery(t, ipv4Gateway, nil, 0, ipv6Pinhole, nil)
|
||||||
|
|
||||||
|
got, err := defaultDiscoverGateway(context.Background())
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Same(t, ipv4Gateway, got)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("no IPv4 gateway still opens an IPv6 pinhole", func(t *testing.T) {
|
||||||
|
stubDiscovery(t, nil, nat.ErrNoNATFound, 0, ipv6Pinhole, nil)
|
||||||
|
|
||||||
|
got, err := defaultDiscoverGateway(context.Background())
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Same(t, ipv6Pinhole, got)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("no gateway and no pinhole reports the original failure", func(t *testing.T) {
|
||||||
|
stubDiscovery(t, nil, nat.ErrNoNATFound, 0, nil, errors.New("no IPv6 route"))
|
||||||
|
|
||||||
|
got, err := defaultDiscoverGateway(context.Background())
|
||||||
|
|
||||||
|
assert.Nil(t, got)
|
||||||
|
assert.ErrorIs(t, err, nat.ErrNoNATFound, "the pinhole failure must not mask why no gateway was found")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a failure other than no-gateway is reported as is", func(t *testing.T) {
|
||||||
|
stubDiscovery(t, nil, otherErr, 0, ipv6Pinhole, nil)
|
||||||
|
|
||||||
|
got, err := defaultDiscoverGateway(context.Background())
|
||||||
|
|
||||||
|
assert.Nil(t, got)
|
||||||
|
assert.ErrorIs(t, err, otherErr)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("the pinhole survives gateway discovery using its whole budget", func(t *testing.T) {
|
||||||
|
// On one shared context the probe would start already expired, which is
|
||||||
|
// how this failed against a real gateway.
|
||||||
|
reserve := 50 * time.Millisecond
|
||||||
|
origReserve := pinholeDiscoveryTimeout
|
||||||
|
pinholeDiscoveryTimeout = reserve
|
||||||
|
t.Cleanup(func() { pinholeDiscoveryTimeout = origReserve })
|
||||||
|
|
||||||
|
budget := 4 * reserve
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), budget)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
stubDiscovery(t, nil, nat.ErrNoNATFound, budget, ipv6Pinhole, nil)
|
||||||
|
|
||||||
|
got, err := defaultDiscoverGateway(ctx)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Same(t, ipv6Pinhole, got)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReserveForPinhole(t *testing.T) {
|
||||||
|
origReserve := pinholeDiscoveryTimeout
|
||||||
|
pinholeDiscoveryTimeout = time.Second
|
||||||
|
t.Cleanup(func() { pinholeDiscoveryTimeout = origReserve })
|
||||||
|
|
||||||
|
t.Run("a budget is divided", func(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
gatewayCtx, cancelGateway := reserveForPinhole(ctx)
|
||||||
|
defer cancelGateway()
|
||||||
|
|
||||||
|
deadline, ok := gatewayCtx.Deadline()
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.InDelta(t, 9*time.Second, time.Until(deadline), float64(500*time.Millisecond))
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a budget too small to divide is left whole", func(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
gatewayCtx, cancelGateway := reserveForPinhole(ctx)
|
||||||
|
defer cancelGateway()
|
||||||
|
|
||||||
|
deadline, ok := gatewayCtx.Deadline()
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.InDelta(t, 500*time.Millisecond, time.Until(deadline), float64(100*time.Millisecond))
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("no deadline stays unbounded", func(t *testing.T) {
|
||||||
|
gatewayCtx, cancelGateway := reserveForPinhole(context.Background())
|
||||||
|
defer cancelGateway()
|
||||||
|
|
||||||
|
_, ok := gatewayCtx.Deadline()
|
||||||
|
assert.False(t, ok)
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,130 @@
|
|||||||
|
package profilemanager
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
const prefsFileSuffix = ".prefs.json"
|
||||||
|
|
||||||
|
var prefsMu sync.Mutex
|
||||||
|
|
||||||
|
// Prefs is a namespaced per-profile preference store backed by a single JSON
|
||||||
|
// file next to the profile config; it is deleted together with the profile.
|
||||||
|
type Prefs struct {
|
||||||
|
path string
|
||||||
|
}
|
||||||
|
|
||||||
|
// ProfilePrefs returns the preference store of the profile identified by id.
|
||||||
|
func (s *ServiceManager) ProfilePrefs(id ID, username string) (*Prefs, error) {
|
||||||
|
if !IsValidProfileFilenameStem(id) {
|
||||||
|
return nil, fmt.Errorf("invalid profile ID: %q", id)
|
||||||
|
}
|
||||||
|
if id == defaultProfileName {
|
||||||
|
return &Prefs{path: filepath.Join(filepath.Dir(DefaultConfigPath), id.String()+prefsFileSuffix)}, nil
|
||||||
|
}
|
||||||
|
configDir, err := s.getConfigDir(username)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("get config directory for user %s: %w", username, err)
|
||||||
|
}
|
||||||
|
return &Prefs{path: filepath.Join(configDir, id.String()+prefsFileSuffix)}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get unmarshals the namespace section into v and reports whether it exists.
|
||||||
|
func (p *Prefs) Get(namespace string, v any) (bool, error) {
|
||||||
|
if namespace == "" {
|
||||||
|
return false, fmt.Errorf("empty prefs namespace")
|
||||||
|
}
|
||||||
|
|
||||||
|
prefsMu.Lock()
|
||||||
|
defer prefsMu.Unlock()
|
||||||
|
|
||||||
|
sections, err := readPrefsFile(p.path)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
raw, ok := sections[namespace]
|
||||||
|
if !ok {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(raw, v); err != nil {
|
||||||
|
return false, fmt.Errorf("decode prefs namespace %q: %w", namespace, err)
|
||||||
|
}
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Put stores v as the namespace section, replacing any previous value.
|
||||||
|
func (p *Prefs) Put(namespace string, v any) error {
|
||||||
|
if namespace == "" {
|
||||||
|
return fmt.Errorf("empty prefs namespace")
|
||||||
|
}
|
||||||
|
raw, err := json.Marshal(v)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("encode prefs namespace %q: %w", namespace, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
prefsMu.Lock()
|
||||||
|
defer prefsMu.Unlock()
|
||||||
|
|
||||||
|
sections, err := readPrefsFile(p.path)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
sections[namespace] = raw
|
||||||
|
return writePrefsFile(p.path, sections)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove deletes the namespace section; a missing one is not an error.
|
||||||
|
func (p *Prefs) Remove(namespace string) error {
|
||||||
|
if namespace == "" {
|
||||||
|
return fmt.Errorf("empty prefs namespace")
|
||||||
|
}
|
||||||
|
|
||||||
|
prefsMu.Lock()
|
||||||
|
defer prefsMu.Unlock()
|
||||||
|
|
||||||
|
sections, err := readPrefsFile(p.path)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if _, ok := sections[namespace]; !ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
delete(sections, namespace)
|
||||||
|
return writePrefsFile(p.path, sections)
|
||||||
|
}
|
||||||
|
|
||||||
|
func removePrefsFile(path string) error {
|
||||||
|
prefsMu.Lock()
|
||||||
|
defer prefsMu.Unlock()
|
||||||
|
return os.Remove(path)
|
||||||
|
}
|
||||||
|
|
||||||
|
func readPrefsFile(path string) (map[string]json.RawMessage, error) {
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return map[string]json.RawMessage{}, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("read prefs: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
sections := map[string]json.RawMessage{}
|
||||||
|
if err := json.Unmarshal(data, §ions); err != nil {
|
||||||
|
return nil, fmt.Errorf("decode prefs: %w", err)
|
||||||
|
}
|
||||||
|
return sections, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func writePrefsFile(path string, sections map[string]json.RawMessage) error {
|
||||||
|
if err := util.WriteJsonWithRestrictedPermission(context.Background(), path, sections); err != nil {
|
||||||
|
return fmt.Errorf("write prefs: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,138 @@
|
|||||||
|
package profilemanager
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
type testPrefsSection struct {
|
||||||
|
Mode uint8 `json:"mode"`
|
||||||
|
Dest string `json:"dest"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProfilePrefs_RoundTrip(t *testing.T) {
|
||||||
|
withTestSM(t, func(sm *ServiceManager, username string) {
|
||||||
|
created, err := sm.AddProfile("work", username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
prefs, err := sm.ProfilePrefs(created.ID, username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 2, Dest: "/tmp/x"}))
|
||||||
|
require.NoError(t, prefs.Put("other", map[string]int{"n": 1}))
|
||||||
|
|
||||||
|
var got testPrefsSection
|
||||||
|
found, err := prefs.Get("filedrop", &got)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.True(t, found)
|
||||||
|
assert.Equal(t, testPrefsSection{Mode: 2, Dest: "/tmp/x"}, got)
|
||||||
|
|
||||||
|
var other map[string]int
|
||||||
|
found, err = prefs.Get("other", &other)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.True(t, found)
|
||||||
|
assert.Equal(t, map[string]int{"n": 1}, other)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProfilePrefs_GetMissingNamespace(t *testing.T) {
|
||||||
|
withTestSM(t, func(sm *ServiceManager, username string) {
|
||||||
|
created, err := sm.AddProfile("work", username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
prefs, err := sm.ProfilePrefs(created.ID, username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var got testPrefsSection
|
||||||
|
found, err := prefs.Get("filedrop", &got)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.False(t, found)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProfilePrefs_RemoveNamespace(t *testing.T) {
|
||||||
|
withTestSM(t, func(sm *ServiceManager, username string) {
|
||||||
|
created, err := sm.AddProfile("work", username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
prefs, err := sm.ProfilePrefs(created.ID, username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 1}))
|
||||||
|
require.NoError(t, prefs.Put("other", map[string]int{"n": 1}))
|
||||||
|
require.NoError(t, prefs.Remove("filedrop"))
|
||||||
|
require.NoError(t, prefs.Remove("missing"))
|
||||||
|
|
||||||
|
var got testPrefsSection
|
||||||
|
found, err := prefs.Get("filedrop", &got)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.False(t, found)
|
||||||
|
|
||||||
|
var other map[string]int
|
||||||
|
found, err = prefs.Get("other", &other)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.True(t, found)
|
||||||
|
assert.Equal(t, map[string]int{"n": 1}, other)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProfilePrefs_RejectsInvalidID(t *testing.T) {
|
||||||
|
withTestSM(t, func(sm *ServiceManager, username string) {
|
||||||
|
_, err := sm.ProfilePrefs("../escape", username)
|
||||||
|
assert.Error(t, err)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProfilePrefs_RejectsEmptyNamespace(t *testing.T) {
|
||||||
|
withTestSM(t, func(sm *ServiceManager, username string) {
|
||||||
|
created, err := sm.AddProfile("work", username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
prefs, err := sm.ProfilePrefs(created.ID, username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
_, err = prefs.Get("", &testPrefsSection{})
|
||||||
|
assert.Error(t, err)
|
||||||
|
assert.Error(t, prefs.Put("", testPrefsSection{}))
|
||||||
|
assert.Error(t, prefs.Remove(""))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProfilePrefs_DefaultProfile(t *testing.T) {
|
||||||
|
withTestSM(t, func(sm *ServiceManager, username string) {
|
||||||
|
prefs, err := sm.ProfilePrefs(defaultProfileName, username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 1}))
|
||||||
|
|
||||||
|
expected := filepath.Join(filepath.Dir(DefaultConfigPath), "default"+prefsFileSuffix)
|
||||||
|
_, err = os.Stat(expected)
|
||||||
|
require.NoError(t, err)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoveProfile_DeletesPrefsFile(t *testing.T) {
|
||||||
|
withTestSM(t, func(sm *ServiceManager, username string) {
|
||||||
|
created, err := sm.AddProfile("work", username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
prefs, err := sm.ProfilePrefs(created.ID, username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 2}))
|
||||||
|
|
||||||
|
configDir, err := sm.getConfigDir(username)
|
||||||
|
require.NoError(t, err)
|
||||||
|
prefsPath := filepath.Join(configDir, created.ID.String()+prefsFileSuffix)
|
||||||
|
_, err = os.Stat(prefsPath)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.NoError(t, sm.RemoveProfile(created.ID, username))
|
||||||
|
_, err = os.Stat(prefsPath)
|
||||||
|
assert.True(t, errors.Is(err, os.ErrNotExist), "prefs file should be removed")
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -420,6 +420,11 @@ func (s *ServiceManager) RemoveProfile(id ID, username string) error {
|
|||||||
log.Warnf("failed to remove profile state file %s: %v", stateFile, err)
|
log.Warnf("failed to remove profile state file %s: %v", stateFile, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
prefsFile := filepath.Join(filepath.Dir(target.Path), id.String()+prefsFileSuffix)
|
||||||
|
if err := removePrefsFile(prefsFile); err != nil && !os.IsNotExist(err) {
|
||||||
|
log.Warnf("failed to remove profile prefs file %s: %v", prefsFile, err)
|
||||||
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -87,9 +87,10 @@ func (pm *ProfileManager) SetActiveProfileState(state *ProfileState) error {
|
|||||||
|
|
||||||
// RemoveProfileState deletes the per-profile state file (which holds the
|
// RemoveProfileState deletes the per-profile state file (which holds the
|
||||||
// account email used for the SSO login hint and the UI display). Called after
|
// account email used for the SSO login hint and the UI display). Called after
|
||||||
// a successful logout so a logged-out profile no longer shows a stale account
|
// profile removal; logout keeps the file so the next login can pass the email
|
||||||
// email. The state file only stores the email, so deleting it is equivalent to
|
// as the login_hint. The state file only stores the email, so deleting it is
|
||||||
// clearing it; the next SSO login recreates it. A missing file is not an error.
|
// equivalent to clearing it; the next SSO login recreates it. A missing file
|
||||||
|
// is not an error.
|
||||||
func (pm *ProfileManager) RemoveProfileState(profileName string) error {
|
func (pm *ProfileManager) RemoveProfileState(profileName string) error {
|
||||||
configDir, err := getConfigDir()
|
configDir, err := getConfigDir()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,82 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package systemops
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSortRouteCandidates(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
candidates []candidateRoute
|
||||||
|
wantOrder []uint32
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "longest prefix wins over metrics",
|
||||||
|
candidates: []candidateRoute{
|
||||||
|
{interfaceIndex: 1, prefixLength: 0, routeMetric: 0, interfaceMetric: 5},
|
||||||
|
{interfaceIndex: 2, prefixLength: 24, routeMetric: 100, interfaceMetric: 50},
|
||||||
|
},
|
||||||
|
wantOrder: []uint32{2, 1},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Windows ranks equal-length prefixes by route metric + interface metric,
|
||||||
|
// so a higher route metric on a low metric interface can still win.
|
||||||
|
name: "combined metric beats route metric alone",
|
||||||
|
candidates: []candidateRoute{
|
||||||
|
{interfaceIndex: 8, prefixLength: 0, routeMetric: 0, interfaceMetric: 100},
|
||||||
|
{interfaceIndex: 5, prefixLength: 0, routeMetric: 10, interfaceMetric: 5},
|
||||||
|
},
|
||||||
|
wantOrder: []uint32{5, 8},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "lower combined metric wins",
|
||||||
|
candidates: []candidateRoute{
|
||||||
|
{interfaceIndex: 5, prefixLength: 0, routeMetric: 300, interfaceMetric: 5},
|
||||||
|
{interfaceIndex: 8, prefixLength: 0, routeMetric: 0, interfaceMetric: 100},
|
||||||
|
},
|
||||||
|
wantOrder: []uint32{8, 5},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "equal combined metric falls back to route metric",
|
||||||
|
candidates: []candidateRoute{
|
||||||
|
{interfaceIndex: 1, prefixLength: 0, routeMetric: 20, interfaceMetric: 10},
|
||||||
|
{interfaceIndex: 2, prefixLength: 0, routeMetric: 5, interfaceMetric: 25},
|
||||||
|
},
|
||||||
|
wantOrder: []uint32{2, 1},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// The metrics are uint32 on the Windows side, so the sum must not wrap.
|
||||||
|
name: "combined metric beyond the uint32 range",
|
||||||
|
candidates: []candidateRoute{
|
||||||
|
{interfaceIndex: 1, prefixLength: 0, routeMetric: math.MaxUint32, interfaceMetric: 5},
|
||||||
|
{interfaceIndex: 2, prefixLength: 0, routeMetric: math.MaxUint32 - 10, interfaceMetric: 5},
|
||||||
|
},
|
||||||
|
wantOrder: []uint32{2, 1},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unknown interface metric ranks on route metric only",
|
||||||
|
candidates: []candidateRoute{
|
||||||
|
{interfaceIndex: 1, prefixLength: 0, routeMetric: 30, interfaceMetric: -1},
|
||||||
|
{interfaceIndex: 2, prefixLength: 0, routeMetric: 5, interfaceMetric: 10},
|
||||||
|
},
|
||||||
|
wantOrder: []uint32{2, 1},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
sortRouteCandidates(tt.candidates)
|
||||||
|
|
||||||
|
got := make([]uint32, 0, len(tt.candidates))
|
||||||
|
for _, c := range tt.candidates {
|
||||||
|
got = append(got, c.interfaceIndex)
|
||||||
|
}
|
||||||
|
assert.Equal(t, tt.wantOrder, got)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -882,26 +882,40 @@ func getInterfaceMetric(interfaceIndex uint32, family int16) int {
|
|||||||
return int(ipInterfaceRow.Metric)
|
return int(ipInterfaceRow.Metric)
|
||||||
}
|
}
|
||||||
|
|
||||||
// sortRouteCandidates sorts route candidates by priority: prefix length -> route metric -> interface metric
|
// sortRouteCandidates sorts route candidates by priority: prefix length -> combined metric -> route metric.
|
||||||
|
// Windows prefers the longest matching prefix and, among prefixes of the same length, the lowest metric, see
|
||||||
|
// https://learn.microsoft.com/en-us/windows-hardware/customize/desktop/unattend/microsoft-windows-tcpip-interfaces-interface-routes-route-metric
|
||||||
func sortRouteCandidates(candidates []candidateRoute) {
|
func sortRouteCandidates(candidates []candidateRoute) {
|
||||||
sort.Slice(candidates, func(i, j int) bool {
|
sort.Slice(candidates, func(i, j int) bool {
|
||||||
if candidates[i].prefixLength != candidates[j].prefixLength {
|
if candidates[i].prefixLength != candidates[j].prefixLength {
|
||||||
return candidates[i].prefixLength > candidates[j].prefixLength
|
return candidates[i].prefixLength > candidates[j].prefixLength
|
||||||
}
|
}
|
||||||
if candidates[i].routeMetric != candidates[j].routeMetric {
|
mi, mj := combinedMetric(candidates[i]), combinedMetric(candidates[j])
|
||||||
return candidates[i].routeMetric < candidates[j].routeMetric
|
if mi != mj {
|
||||||
|
return mi < mj
|
||||||
}
|
}
|
||||||
return candidates[i].interfaceMetric < candidates[j].interfaceMetric
|
return candidates[i].routeMetric < candidates[j].routeMetric
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// combinedMetric returns the effective metric Windows uses to rank routes with an equal prefix length:
|
||||||
|
// the sum of the route metric and the metric of the interface the route is on, see
|
||||||
|
// https://learn.microsoft.com/en-us/windows-server/networking/technologies/network-subsystem/net-sub-interface-metric
|
||||||
|
// An unknown interface metric contributes nothing.
|
||||||
|
func combinedMetric(candidate candidateRoute) uint64 {
|
||||||
|
if candidate.interfaceMetric < 0 {
|
||||||
|
return uint64(candidate.routeMetric)
|
||||||
|
}
|
||||||
|
return uint64(candidate.routeMetric) + uint64(candidate.interfaceMetric)
|
||||||
|
}
|
||||||
|
|
||||||
// GetBestInterface finds the best interface for reaching a destination,
|
// GetBestInterface finds the best interface for reaching a destination,
|
||||||
// excluding the VPN interface to avoid routing loops.
|
// excluding the VPN interface to avoid routing loops.
|
||||||
//
|
//
|
||||||
// Route selection priority:
|
// Route selection priority:
|
||||||
// 1. Longest prefix match (most specific route)
|
// 1. Longest prefix match (most specific route)
|
||||||
// 2. Lowest route metric
|
// 2. Lowest combined metric (route metric + interface metric)
|
||||||
// 3. Lowest interface metric
|
// 3. Lowest route metric.
|
||||||
func GetBestInterface(dest netip.Addr, vpnIntf string) (*net.Interface, error) {
|
func GetBestInterface(dest netip.Addr, vpnIntf string) (*net.Interface, error) {
|
||||||
var skipInterfaceIndex int
|
var skipInterfaceIndex int
|
||||||
if vpnIntf != "" {
|
if vpnIntf != "" {
|
||||||
@@ -925,7 +939,6 @@ func GetBestInterface(dest netip.Addr, vpnIntf string) (*net.Interface, error) {
|
|||||||
return nil, fmt.Errorf("no route to %s", dest)
|
return nil, fmt.Errorf("no route to %s", dest)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Sort routes: prefix length -> route metric -> interface metric
|
|
||||||
sortRouteCandidates(candidates)
|
sortRouteCandidates(candidates)
|
||||||
|
|
||||||
for _, candidate := range candidates {
|
for _, candidate := range candidates {
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package systemops
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"net"
|
"net"
|
||||||
|
"net/netip"
|
||||||
"syscall"
|
"syscall"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -29,6 +30,7 @@ func ensureIPv6DefaultRoute(t *testing.T) {
|
|||||||
}
|
}
|
||||||
if err := netlink.RouteAdd(route); err != nil {
|
if err := netlink.RouteAdd(route); err != nil {
|
||||||
if errors.Is(err, syscall.EEXIST) {
|
if errors.Is(err, syscall.EEXIST) {
|
||||||
|
requireUsableIPv6Nexthop(t)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
t.Skipf("install IPv6 fallback default route: %v", err)
|
t.Skipf("install IPv6 fallback default route: %v", err)
|
||||||
@@ -38,4 +40,36 @@ func ensureIPv6DefaultRoute(t *testing.T) {
|
|||||||
t.Logf("delete IPv6 fallback default route: %v", err)
|
t.Logf("delete IPv6 fallback default route: %v", err)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
requireUsableIPv6Nexthop(t)
|
||||||
|
}
|
||||||
|
|
||||||
|
// requireUsableIPv6Nexthop skips the test unless the resolved IPv6 default
|
||||||
|
// nexthop can actually carry a route. Installing the default route succeeding
|
||||||
|
// does not imply the kernel accepts it as a nexthop for a concrete prefix.
|
||||||
|
func requireUsableIPv6Nexthop(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
nexthop, err := GetNextHop(netip.IPv6Unspecified())
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("resolve IPv6 default nexthop: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
probe := &netlink.Route{
|
||||||
|
Scope: netlink.SCOPE_UNIVERSE,
|
||||||
|
Table: syscall.RT_TABLE_MAIN,
|
||||||
|
Family: netlink.FAMILY_V6,
|
||||||
|
Dst: &net.IPNet{IP: net.ParseIP("100::64"), Mask: net.CIDRMask(128, 128)},
|
||||||
|
}
|
||||||
|
require.NoError(t, addNextHop(nexthop, probe), "build IPv6 probe route")
|
||||||
|
|
||||||
|
switch err := netlink.RouteAdd(probe); {
|
||||||
|
case err == nil:
|
||||||
|
if err := netlink.RouteDel(probe); err != nil && !errors.Is(err, syscall.ESRCH) {
|
||||||
|
t.Logf("delete IPv6 probe route: %v", err)
|
||||||
|
}
|
||||||
|
case errors.Is(err, syscall.EEXIST):
|
||||||
|
default:
|
||||||
|
t.Skipf("IPv6 nexthop %s unusable for route installation: %v", nexthop, err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -18,8 +18,8 @@ type Service struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func New() (*Service, error) {
|
func New() (*Service, error) {
|
||||||
d, err := NewDetector()
|
d, err := NewDetector() //nolint:staticcheck
|
||||||
if err != nil {
|
if err != nil { //nolint:staticcheck // always errors on platforms without a sleep detector
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -37,23 +37,32 @@
|
|||||||
// Updater Process (Setup):
|
// Updater Process (Setup):
|
||||||
//
|
//
|
||||||
// 1. Receives parameters from service via command-line arguments
|
// 1. Receives parameters from service via command-line arguments
|
||||||
// 2. Runs installer with appropriate silent/quiet flags:
|
// 2. Terminates the UI so the installer does not have to replace a locked image
|
||||||
|
// file, which would otherwise leave the install needing a reboot
|
||||||
|
// 3. Runs installer with appropriate silent/quiet flags:
|
||||||
// - Windows EXE: installer.exe /S
|
// - Windows EXE: installer.exe /S
|
||||||
// - Windows MSI: msiexec.exe /i installer.msi /quiet /qn /l*v msi.log
|
// - Windows MSI: msiexec.exe /i installer.msi /qn /norestart REBOOT=ReallySuppress /l*v msi.log
|
||||||
// - macOS PKG: installer -pkg installer.pkg -target /
|
// - macOS PKG: installer -pkg installer.pkg -target /
|
||||||
// - macOS Homebrew: brew upgrade netbirdio/tap/netbird
|
// - macOS Homebrew: brew upgrade netbirdio/tap/netbird
|
||||||
// 3. Installer terminates daemon and UI processes
|
// 4. Installer terminates the daemon
|
||||||
// 4. Installer replaces binaries with new version
|
// 5. Installer replaces binaries with new version
|
||||||
// 5. Updater waits for installer to complete
|
// 6. Updater waits for installer to complete. On Windows, MSI exit codes 3010
|
||||||
// 6. Updater restarts daemon:
|
// (ERROR_SUCCESS_REBOOT_REQUIRED) and 1641 (ERROR_SUCCESS_REBOOT_INITIATED)
|
||||||
|
// are a pending-reboot outcome, not a failure: the install succeeded, but
|
||||||
|
// some files are only replaced on the next restart (the reboot itself is
|
||||||
|
// suppressed via /norestart and REBOOT=ReallySuppress), and the flow
|
||||||
|
// continues as on success
|
||||||
|
// 7. Updater restarts daemon:
|
||||||
// - Windows: netbird.exe service start
|
// - Windows: netbird.exe service start
|
||||||
// - macOS/Linux: netbird service start
|
// - macOS/Linux: netbird service start
|
||||||
// 7. Updater restarts UI:
|
// 8. Updater restarts UI:
|
||||||
// - Windows: Launches netbird-ui.exe as active console user using CreateProcessAsUser
|
// - Windows: Launches netbird-ui.exe using CreateProcessAsUser in every
|
||||||
|
// session it was terminated in, falling back to the active console session
|
||||||
// - macOS: Uses launchctl asuser to launch NetBird.app for console user
|
// - macOS: Uses launchctl asuser to launch NetBird.app for console user
|
||||||
// - Linux: Not implemented (UI typically auto-starts)
|
// - Linux: Not implemented (UI typically auto-starts)
|
||||||
// 8. Updater writes result.json with success/error status
|
// 9. Updater writes result.json with success/error status (a pending reboot is
|
||||||
// 9. Updater process exits
|
// recorded as success)
|
||||||
|
// 10. Updater process exits
|
||||||
//
|
//
|
||||||
// # Result Communication
|
// # Result Communication
|
||||||
//
|
//
|
||||||
|
|||||||
@@ -42,6 +42,9 @@ func NewWithDir(tempDir string) *Installer {
|
|||||||
// This will run by the original service process
|
// This will run by the original service process
|
||||||
func (u *Installer) RunInstallation(ctx context.Context, targetVersion string) (err error) {
|
func (u *Installer) RunInstallation(ctx context.Context, targetVersion string) (err error) {
|
||||||
resultHandler := NewResultHandler(u.tempDir)
|
resultHandler := NewResultHandler(u.tempDir)
|
||||||
|
if err := resultHandler.ClearStaleResult(); err != nil {
|
||||||
|
log.Warnf("clear stale installer result: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package installer
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
@@ -22,6 +23,12 @@ const (
|
|||||||
|
|
||||||
msiLogFile = "msi.log"
|
msiLogFile = "msi.log"
|
||||||
|
|
||||||
|
// ERROR_SUCCESS_REBOOT_REQUIRED and ERROR_SUCCESS_REBOOT_INITIATED
|
||||||
|
msiRebootRequired = 3010
|
||||||
|
msiRebootInitiated = 1641
|
||||||
|
|
||||||
|
processExitWait = 10 * time.Second
|
||||||
|
|
||||||
msiDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.msi"
|
msiDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.msi"
|
||||||
exeDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.exe"
|
exeDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.exe"
|
||||||
)
|
)
|
||||||
@@ -38,6 +45,8 @@ var (
|
|||||||
func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string, daemonFolder string) (resultErr error) {
|
func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string, daemonFolder string) (resultErr error) {
|
||||||
resultHandler := NewResultHandler(u.tempDir)
|
resultHandler := NewResultHandler(u.tempDir)
|
||||||
|
|
||||||
|
var uiSessions []uint32
|
||||||
|
|
||||||
// Always ensure daemon and UI are restarted after setup
|
// Always ensure daemon and UI are restarted after setup
|
||||||
defer func() {
|
defer func() {
|
||||||
log.Infof("starting daemon back")
|
log.Infof("starting daemon back")
|
||||||
@@ -46,7 +55,7 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
|
|||||||
}
|
}
|
||||||
|
|
||||||
log.Infof("starting UI back")
|
log.Infof("starting UI back")
|
||||||
if err := u.startUIAsUser(daemonFolder); err != nil {
|
if err := u.startUI(daemonFolder, uiSessions); err != nil {
|
||||||
log.Errorf("failed to start UI: %v", err)
|
log.Errorf("failed to start UI: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -75,6 +84,14 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// The UI holds an open handle on its own image. Left running, Restart Manager
|
||||||
|
// cannot shut it down (msiexec runs as LocalSystem here, the UI as the
|
||||||
|
// interactive user), so the MSI falls back to replacing the file on reboot and
|
||||||
|
// marks the install as restart-required. The deferred close-application action
|
||||||
|
// in the package runs too late to prevent that, it happens after
|
||||||
|
// InstallValidate has already registered the file as in use.
|
||||||
|
uiSessions = killUI()
|
||||||
|
|
||||||
var cmd *exec.Cmd
|
var cmd *exec.Cmd
|
||||||
switch installerType {
|
switch installerType {
|
||||||
case TypeExe:
|
case TypeExe:
|
||||||
@@ -84,7 +101,9 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
|
|||||||
installerDir := filepath.Dir(installerFile)
|
installerDir := filepath.Dir(installerFile)
|
||||||
logPath := filepath.Join(installerDir, msiLogFile)
|
logPath := filepath.Join(installerDir, msiLogFile)
|
||||||
log.Infof("run msi installer: %s", installerFile)
|
log.Infof("run msi installer: %s", installerFile)
|
||||||
cmd = exec.CommandContext(ctx, "msiexec.exe", "/i", filepath.Base(installerFile), "/quiet", "/qn", "/l*v", logPath)
|
// REBOOT=ReallySuppress: a silent install has no way to ask, so without it
|
||||||
|
// msiexec reboots the machine on its own if it decides one is needed.
|
||||||
|
cmd = exec.CommandContext(ctx, "msiexec.exe", "/i", filepath.Base(installerFile), "/qn", "/norestart", "REBOOT=ReallySuppress", "/l*v", logPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd.Dir = filepath.Dir(installerFile)
|
cmd.Dir = filepath.Dir(installerFile)
|
||||||
@@ -95,9 +114,13 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
|
|||||||
}
|
}
|
||||||
|
|
||||||
log.Infof("installer started with PID %d", cmd.Process.Pid)
|
log.Infof("installer started with PID %d", cmd.Process.Pid)
|
||||||
if resultErr = cmd.Wait(); resultErr != nil {
|
if err := cmd.Wait(); err != nil {
|
||||||
log.Errorf("installer process finished with error: %v", resultErr)
|
if !isRebootPending(err) {
|
||||||
return
|
resultErr = err
|
||||||
|
log.Errorf("installer process finished with error: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Warnf("installer completed but reported a pending reboot, some files will be replaced on the next restart")
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -117,16 +140,142 @@ func (u *Installer) startDaemon(daemonFolder string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *Installer) startUIAsUser(daemonFolder string) error {
|
func (u *Installer) startUI(daemonFolder string, sessionIDs []uint32) error {
|
||||||
uiPath := filepath.Join(daemonFolder, uiName)
|
uiPath := filepath.Join(daemonFolder, uiName)
|
||||||
log.Infof("starting netbird-ui: %s", uiPath)
|
log.Infof("starting netbird-ui: %s", uiPath)
|
||||||
|
|
||||||
// Get the active console session ID
|
if len(sessionIDs) == 0 {
|
||||||
sessionID := windows.WTSGetActiveConsoleSessionId()
|
sessionID := windows.WTSGetActiveConsoleSessionId()
|
||||||
if sessionID == 0xFFFFFFFF {
|
if sessionID == 0xFFFFFFFF {
|
||||||
return fmt.Errorf("no active user session found")
|
return fmt.Errorf("no active user session found")
|
||||||
|
}
|
||||||
|
sessionIDs = []uint32{sessionID}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var errs []error
|
||||||
|
for _, sessionID := range sessionIDs {
|
||||||
|
if err := startUIInSession(uiPath, sessionID); err != nil {
|
||||||
|
errs = append(errs, fmt.Errorf("session %d: %w", sessionID, err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
log.Infof("netbird-ui started successfully in session %d", sessionID)
|
||||||
|
}
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// isRebootPending reports whether the installer exit code means it succeeded but
|
||||||
|
// left work for the next restart. The reboot itself is suppressed, so this is not
|
||||||
|
// a failure.
|
||||||
|
func isRebootPending(err error) bool {
|
||||||
|
var exitErr *exec.ExitError
|
||||||
|
if !errors.As(err, &exitErr) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
switch exitErr.ExitCode() {
|
||||||
|
case msiRebootRequired, msiRebootInitiated:
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// killUI terminates any running netbird-ui process and returns the IDs of the
|
||||||
|
// interactive sessions the terminated processes belonged to. Setup starts the
|
||||||
|
// UI again in those sessions once the installer is done.
|
||||||
|
func killUI() []uint32 {
|
||||||
|
pids, err := processIDsByName(uiName)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("failed to look up %s processes: %v", uiName, err)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
sessions := make(map[uint32]struct{})
|
||||||
|
for _, pid := range pids {
|
||||||
|
var sessionID uint32
|
||||||
|
if err := windows.ProcessIdToSessionId(pid, &sessionID); err != nil {
|
||||||
|
log.Warnf("failed to look up session of %s (PID %d): %v", uiName, pid, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := terminateProcess(pid); err != nil {
|
||||||
|
log.Warnf("failed to terminate %s (PID %d): %v", uiName, pid, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
log.Infof("terminated %s (PID %d) in session %d", uiName, pid, sessionID)
|
||||||
|
|
||||||
|
if sessionID != 0 {
|
||||||
|
sessions[sessionID] = struct{}{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionIDs := make([]uint32, 0, len(sessions))
|
||||||
|
for sessionID := range sessions {
|
||||||
|
sessionIDs = append(sessionIDs, sessionID)
|
||||||
|
}
|
||||||
|
return sessionIDs
|
||||||
|
}
|
||||||
|
|
||||||
|
func processIDsByName(name string) ([]uint32, error) {
|
||||||
|
snapshot, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPPROCESS, 0)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create process snapshot: %w", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := windows.CloseHandle(snapshot); err != nil {
|
||||||
|
log.Warnf("failed to close process snapshot: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
var entry windows.ProcessEntry32
|
||||||
|
entry.Size = uint32(unsafe.Sizeof(entry))
|
||||||
|
|
||||||
|
var pids []uint32
|
||||||
|
for err = windows.Process32First(snapshot, &entry); err == nil; err = windows.Process32Next(snapshot, &entry) {
|
||||||
|
if strings.EqualFold(windows.UTF16ToString(entry.ExeFile[:]), name) {
|
||||||
|
pids = append(pids, entry.ProcessID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !errors.Is(err, windows.ERROR_NO_MORE_FILES) {
|
||||||
|
return nil, fmt.Errorf("enumerate processes: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return pids, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func terminateProcess(pid uint32) error {
|
||||||
|
handle, err := windows.OpenProcess(windows.PROCESS_TERMINATE|windows.SYNCHRONIZE, false, pid)
|
||||||
|
if err != nil {
|
||||||
|
// The process may have exited between enumeration and now.
|
||||||
|
if errors.Is(err, windows.ERROR_INVALID_PARAMETER) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return fmt.Errorf("open process: %w", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := windows.CloseHandle(handle); err != nil {
|
||||||
|
log.Warnf("failed to close process handle: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
if err := windows.TerminateProcess(handle, 0); err != nil {
|
||||||
|
return fmt.Errorf("terminate process: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Wait for the handle to signal so the image file is released before the
|
||||||
|
// installer tries to overwrite it. A timeout is reported through the returned
|
||||||
|
// event, not through err, which stays nil unless the wait itself failed.
|
||||||
|
event, err := windows.WaitForSingleObject(handle, uint32(processExitWait.Milliseconds()))
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("wait for process exit: %w", err)
|
||||||
|
}
|
||||||
|
if event != windows.WAIT_OBJECT_0 {
|
||||||
|
return fmt.Errorf("wait for process exit: unexpected wait result %#x", event)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func startUIInSession(uiPath string, sessionID uint32) error {
|
||||||
// Get the user token for that session
|
// Get the user token for that session
|
||||||
var userToken windows.Token
|
var userToken windows.Token
|
||||||
err := windows.WTSQueryUserToken(sessionID, &userToken)
|
err := windows.WTSQueryUserToken(sessionID, &userToken)
|
||||||
@@ -158,6 +307,16 @@ func (u *Installer) startUIAsUser(daemonFolder string) error {
|
|||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
|
var env *uint16
|
||||||
|
if err := windows.CreateEnvironmentBlock(&env, primaryToken, false); err != nil {
|
||||||
|
return fmt.Errorf("create environment block: %w", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if err := windows.DestroyEnvironmentBlock(env); err != nil {
|
||||||
|
log.Warnf("failed to destroy environment block: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
// Prepare startup info
|
// Prepare startup info
|
||||||
var si windows.StartupInfo
|
var si windows.StartupInfo
|
||||||
si.Cb = uint32(unsafe.Sizeof(si))
|
si.Cb = uint32(unsafe.Sizeof(si))
|
||||||
@@ -180,7 +339,7 @@ func (u *Installer) startUIAsUser(daemonFolder string) error {
|
|||||||
nil,
|
nil,
|
||||||
false,
|
false,
|
||||||
creationFlags,
|
creationFlags,
|
||||||
nil,
|
env,
|
||||||
nil,
|
nil,
|
||||||
&si,
|
&si,
|
||||||
&pi,
|
&pi,
|
||||||
@@ -197,7 +356,6 @@ func (u *Installer) startUIAsUser(daemonFolder string) error {
|
|||||||
log.Warnf("failed to close thread handle: %v", err)
|
log.Warnf("failed to close thread handle: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Infof("netbird-ui started successfully in session %d", sessionID)
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
package installer
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"os/exec"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// exitErrorWithCode returns a real *exec.ExitError carrying the given exit code.
|
||||||
|
func exitErrorWithCode(t *testing.T, code int) error {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
err := exec.Command("cmd.exe", "/c", "exit "+strconv.Itoa(code)).Run()
|
||||||
|
if err == nil {
|
||||||
|
t.Fatalf("expected a non-zero exit for code %d", code)
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsRebootPending(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
code int
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{name: "reboot required", code: msiRebootRequired, want: true},
|
||||||
|
{name: "reboot initiated", code: msiRebootInitiated, want: true},
|
||||||
|
{name: "generic failure", code: 1603, want: false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if got := isRebootPending(exitErrorWithCode(t, tt.code)); got != tt.want {
|
||||||
|
t.Errorf("isRebootPending(exit %d) = %v, want %v", tt.code, got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestProcessIDsByNameAndTerminate spawns a long-running system process, finds it
|
||||||
|
// by name and terminates it, covering the path the updater uses to release the UI
|
||||||
|
// image file before the installer replaces it.
|
||||||
|
func TestProcessIDsByNameAndTerminate(t *testing.T) {
|
||||||
|
cmd := exec.Command("ping.exe", "-n", "60", "127.0.0.1")
|
||||||
|
if err := cmd.Start(); err != nil {
|
||||||
|
t.Fatalf("start ping: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
pid := uint32(cmd.Process.Pid)
|
||||||
|
killed := false
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if !killed {
|
||||||
|
_ = cmd.Process.Kill()
|
||||||
|
}
|
||||||
|
_ = cmd.Wait()
|
||||||
|
})
|
||||||
|
|
||||||
|
// Name matching must be case-insensitive: the snapshot reports PING.EXE.
|
||||||
|
pids, err := processIDsByName("ping.exe")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("processIDsByName: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !slices.Contains(pids, pid) {
|
||||||
|
t.Fatalf("PID %d not among the ping.exe processes found: %v", pid, pids)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := terminateProcess(pid); err != nil {
|
||||||
|
t.Fatalf("terminateProcess: %v", err)
|
||||||
|
}
|
||||||
|
killed = true
|
||||||
|
|
||||||
|
// terminateProcess only returns once the handle has signalled, so the process
|
||||||
|
// is already gone and Wait must not block. It exits with the code passed to
|
||||||
|
// TerminateProcess, which is 0, so Wait reports no error.
|
||||||
|
if err := cmd.Wait(); err != nil {
|
||||||
|
t.Fatalf("wait for terminated ping: %v", err)
|
||||||
|
}
|
||||||
|
if !cmd.ProcessState.Exited() {
|
||||||
|
t.Error("process did not exit after terminateProcess")
|
||||||
|
}
|
||||||
|
|
||||||
|
remaining, err := processIDsByName("ping.exe")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("processIDsByName after terminate: %v", err)
|
||||||
|
}
|
||||||
|
if slices.Contains(remaining, pid) {
|
||||||
|
t.Errorf("PID %d still listed after terminateProcess", pid)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcessIDsByNameNoMatch(t *testing.T) {
|
||||||
|
pids, err := processIDsByName("netbird-nonexistent-process.exe")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("processIDsByName: %v", err)
|
||||||
|
}
|
||||||
|
if len(pids) != 0 {
|
||||||
|
t.Errorf("expected no matches, got %v", pids)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsRebootPendingNonExitError(t *testing.T) {
|
||||||
|
if isRebootPending(errors.New("start installer: file not found")) {
|
||||||
|
t.Error("a non-exit error must not be treated as a pending reboot")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -54,6 +54,12 @@ func (rh *ResultHandler) GetErrorResultReason() string {
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ClearStaleResult removes a result file left over from a previous installation
|
||||||
|
// attempt so result watchers cannot read an outdated outcome for the current attempt.
|
||||||
|
func (rh *ResultHandler) ClearStaleResult() error {
|
||||||
|
return rh.cleanup()
|
||||||
|
}
|
||||||
|
|
||||||
func (rh *ResultHandler) WriteSuccess() error {
|
func (rh *ResultHandler) WriteSuccess() error {
|
||||||
result := Result{
|
result := Result{
|
||||||
Success: true,
|
Success: true,
|
||||||
|
|||||||
@@ -435,7 +435,7 @@ func (m *Manager) install(ctx context.Context, pendingVersion *v.Version) error
|
|||||||
}
|
}
|
||||||
|
|
||||||
inst := installer.New()
|
inst := installer.New()
|
||||||
if err := inst.RunInstallation(ctx, pendingVersion.String()); err != nil {
|
if err := inst.RunInstallation(ctx, pendingVersion.String()); err != nil { //nolint:staticcheck // always errors on platforms without an installer
|
||||||
log.Errorf("error triggering update: %v", err)
|
log.Errorf("error triggering update: %v", err)
|
||||||
m.statusRecorder.PublishEvent(
|
m.statusRecorder.PublishEvent(
|
||||||
cProto.SystemEvent_ERROR,
|
cProto.SystemEvent_ERROR,
|
||||||
|
|||||||
@@ -22,6 +22,8 @@ import (
|
|||||||
"github.com/netbirdio/netbird/client/internal/listener"
|
"github.com/netbirdio/netbird/client/internal/listener"
|
||||||
"github.com/netbirdio/netbird/client/internal/peer"
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/client/netstate"
|
||||||
|
"github.com/netbirdio/netbird/client/netsweep"
|
||||||
"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"
|
||||||
@@ -36,11 +38,6 @@ const (
|
|||||||
AnonymizeLevelStrict = nbAnonymize.LevelStrictString
|
AnonymizeLevelStrict = nbAnonymize.LevelStrictString
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConnectionListener export internal Listener for mobile
|
|
||||||
type ConnectionListener interface {
|
|
||||||
peer.Listener
|
|
||||||
}
|
|
||||||
|
|
||||||
// RouteListener export internal RouteListener for mobile
|
// RouteListener export internal RouteListener for mobile
|
||||||
type NetworkChangeListener interface {
|
type NetworkChangeListener interface {
|
||||||
listener.NetworkChangeListener
|
listener.NetworkChangeListener
|
||||||
@@ -87,6 +84,12 @@ type Client struct {
|
|||||||
onHostDnsFn func([]string)
|
onHostDnsFn func([]string)
|
||||||
dnsManager dns.IosDnsManager
|
dnsManager dns.IosDnsManager
|
||||||
loginComplete bool
|
loginComplete bool
|
||||||
|
// netState outlives engine restarts: it mirrors the OS connectivity, not
|
||||||
|
// the engine lifecycle. Run injects it into each new ConnectClient, which
|
||||||
|
// distributes it to every reconnection loop.
|
||||||
|
netState *netstate.State
|
||||||
|
// sweeper also outlives engine restarts; NotifyNetworkChange sweeps it.
|
||||||
|
sweeper *netsweep.Sweeper
|
||||||
// preloadedConfig holds config loaded from JSON (used on tvOS where file writes are blocked)
|
// preloadedConfig holds config loaded from JSON (used on tvOS where file writes are blocked)
|
||||||
preloadedConfig *profilemanager.Config
|
preloadedConfig *profilemanager.Config
|
||||||
|
|
||||||
@@ -109,6 +112,8 @@ func NewClient(cfgFile, stateFile, cacheDir, logFilePath, deviceName string, osV
|
|||||||
ctxCancelLock: &sync.Mutex{},
|
ctxCancelLock: &sync.Mutex{},
|
||||||
networkChangeListener: networkChangeListener,
|
networkChangeListener: networkChangeListener,
|
||||||
dnsManager: dnsManager,
|
dnsManager: dnsManager,
|
||||||
|
netState: netstate.New(),
|
||||||
|
sweeper: netsweep.New(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -184,7 +189,8 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error {
|
|||||||
c.onHostDnsFn = func([]string) {}
|
c.onHostDnsFn = func([]string) {}
|
||||||
cfg.WgIface = interfaceName
|
cfg.WgIface = interfaceName
|
||||||
|
|
||||||
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder)
|
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
|
||||||
|
internal.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper))
|
||||||
c.setState(cfg, connectClient)
|
c.setState(cfg, connectClient)
|
||||||
// Persist the latest sync response so DebugBundle can include the network
|
// Persist the latest sync response so DebugBundle can include the network
|
||||||
// map. On iOS this is backed by disk to keep it out of the constrained
|
// map. On iOS this is backed by disk to keep it out of the constrained
|
||||||
@@ -193,6 +199,25 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error {
|
|||||||
return connectClient.RunOniOS(fd, c.networkChangeListener, c.dnsManager, c.stateFile, c.cacheDir, c.logFilePath)
|
return connectClient.RunOniOS(fd, c.networkChangeListener, c.dnsManager, c.stateFile, c.cacheDir, c.logFilePath)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetNetworkAvailable feeds OS-reported network availability into the client
|
||||||
|
// (e.g. from NWPathMonitor). 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.
|
||||||
|
func (c *Client) SetNetworkAvailable(available bool) {
|
||||||
|
c.netState.Set(available)
|
||||||
|
c.recorder.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.sweeper.MarkNetworkChange()
|
||||||
|
log.Infof("network change: connections marked stale")
|
||||||
|
}
|
||||||
|
|
||||||
// Stop the internal client and free the resources
|
// Stop the internal client and free the resources
|
||||||
func (c *Client) Stop() {
|
func (c *Client) Stop() {
|
||||||
c.ctxCancelLock.Lock()
|
c.ctxCancelLock.Lock()
|
||||||
@@ -331,7 +356,11 @@ func (c *Client) GetStatusDetails() *StatusDetails {
|
|||||||
|
|
||||||
// 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
|
||||||
|
|||||||
@@ -0,0 +1,43 @@
|
|||||||
|
//go:build ios
|
||||||
|
|
||||||
|
package NetBirdSDK
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Client state values, re-exported as basic constants so gomobile emits them
|
||||||
|
// into the generated 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 intentionally lacks OnStateChanged for now: adding a method to a gomobile
|
||||||
|
// interface breaks every Swift implementation, so the iOS app keeps building
|
||||||
|
// against the legacy per-state callbacks. A follow-up will extend it together
|
||||||
|
// with the app.
|
||||||
|
type ConnectionListener interface {
|
||||||
|
OnConnected()
|
||||||
|
OnDisconnected()
|
||||||
|
OnConnecting()
|
||||||
|
OnDisconnecting()
|
||||||
|
OnAddressChanged(string, string)
|
||||||
|
OnPeersListChanged(int)
|
||||||
|
}
|
||||||
|
|
||||||
|
// connectionListenerAdapter adapts the gomobile-facing ConnectionListener to
|
||||||
|
// peer.Listener.
|
||||||
|
type connectionListenerAdapter struct {
|
||||||
|
ConnectionListener
|
||||||
|
}
|
||||||
|
|
||||||
|
// OnStateChanged is dropped on iOS until the app adopts the state callback;
|
||||||
|
// the legacy per-state callbacks continue to fire.
|
||||||
|
func (a connectionListenerAdapter) OnStateChanged(peer.ClientState) {}
|
||||||
@@ -323,7 +323,7 @@ func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName strin
|
|||||||
const authInfoRequestTimeout = 30 * time.Second
|
const authInfoRequestTimeout = 30 * time.Second
|
||||||
|
|
||||||
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, forceDeviceAuth bool) (*auth.TokenInfo, error) {
|
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, forceDeviceAuth bool) (*auth.TokenInfo, error) {
|
||||||
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth)
|
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth, "")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
|
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,110 @@
|
|||||||
|
// Package netstate tracks OS-reported network availability for the client.
|
||||||
|
//
|
||||||
|
// A State instance is owned by the platform integration (e.g. the Android or
|
||||||
|
// iOS bindings, fed from ConnectivityManager callbacks or NWPathMonitor) and
|
||||||
|
// is injected into the connection retry loops (management, signal, relay,
|
||||||
|
// peer guards and the top-level connect loop), which consult it to avoid
|
||||||
|
// burning CPU and battery on reconnect attempts while the device has no
|
||||||
|
// network at all (e.g. airplane mode), and to reset their backoff as soon as
|
||||||
|
// the network returns.
|
||||||
|
//
|
||||||
|
// Consumers hold a *State that may be nil — every non-mobile platform leaves
|
||||||
|
// it unset. The read methods are safe on a nil receiver: they report online
|
||||||
|
// and never block, so consumers behave as if this package did not exist.
|
||||||
|
package netstate
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
)
|
||||||
|
|
||||||
|
// State holds the OS-reported network availability. The zero value is not
|
||||||
|
// usable; create instances with New.
|
||||||
|
type State struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
online bool
|
||||||
|
changed chan struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// New creates a State that starts online. Platforms without network tracking
|
||||||
|
// pass a nil *State instead: the read methods treat nil as always online and
|
||||||
|
// never block, so consumers need no nil guards.
|
||||||
|
func New() *State {
|
||||||
|
return &State{
|
||||||
|
online: true,
|
||||||
|
changed: make(chan struct{}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set records whether the OS reports any usable network. Transitions wake up
|
||||||
|
// all Wait callers immediately. Unlike the read methods, Set is not nil-safe:
|
||||||
|
// it is only for the platform owner that created the State with New.
|
||||||
|
func (s *State) Set(online bool) {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
if s.online == online {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.online = online
|
||||||
|
close(s.changed)
|
||||||
|
s.changed = make(chan struct{})
|
||||||
|
log.Infof("OS network availability changed: online=%t", online)
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsOnline reports whether the OS reports at least one usable network. On a
|
||||||
|
// nil receiver — no State injected — it reports online.
|
||||||
|
func (s *State) IsOnline() bool {
|
||||||
|
if s == nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
return s.online
|
||||||
|
}
|
||||||
|
|
||||||
|
// Changed returns a channel closed on the next availability transition, for
|
||||||
|
// callers that already own a select loop and cannot block in Wait. Re-read it
|
||||||
|
// after every fire: each transition installs a fresh channel. On a nil
|
||||||
|
// receiver — no State injected — it returns nil, which blocks forever in a
|
||||||
|
// select, so the caller simply never observes a transition.
|
||||||
|
func (s *State) Changed() <-chan struct{} {
|
||||||
|
if s == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
return s.changed
|
||||||
|
}
|
||||||
|
|
||||||
|
// Wait blocks while the network is offline. It reports whether it had to
|
||||||
|
// wait, so callers can reset their backoff after an outage. It returns early
|
||||||
|
// with the context error when ctx is done. On a nil receiver — no State
|
||||||
|
// injected — it returns immediately.
|
||||||
|
func (s *State) Wait(ctx context.Context) (bool, error) {
|
||||||
|
if s == nil {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
waited := false
|
||||||
|
for {
|
||||||
|
s.mu.Lock()
|
||||||
|
if s.online {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return waited, nil
|
||||||
|
}
|
||||||
|
ch := s.changed
|
||||||
|
s.mu.Unlock()
|
||||||
|
|
||||||
|
if !waited {
|
||||||
|
waited = true
|
||||||
|
log.Debugf("network is offline, pausing connection attempts")
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return waited, ctx.Err()
|
||||||
|
case <-ch:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user