mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-01 05:11:29 +02:00
Compare commits
119 Commits
netmap_pro
...
android/gu
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4263315527 | ||
|
|
56ff5237dd | ||
|
|
3a17d0381c | ||
|
|
6155c94b05 | ||
|
|
09f7fb6510 | ||
|
|
4475819f38 | ||
|
|
c8adaa45da | ||
|
|
e970daaf5f | ||
|
|
5ae323a555 | ||
|
|
19337dc056 | ||
|
|
fd06d9a3d5 | ||
|
|
1bf54ddd8f | ||
|
|
0f5d2d91fb | ||
|
|
0b7e6a9f46 | ||
|
|
f2c1070f95 | ||
|
|
df39c2b254 | ||
|
|
dd2bdc0de3 | ||
|
|
44fef45c2f | ||
|
|
3d1f209ea3 | ||
|
|
2ef457be95 | ||
|
|
63c320b6a9 | ||
|
|
4acbe2670a | ||
|
|
0fb4c8c423 | ||
|
|
42e45ff9f9 | ||
|
|
9269b56386 | ||
|
|
b3f9b82442 | ||
|
|
8a43f4f943 | ||
|
|
2f268c8141 | ||
|
|
e3c4128164 | ||
|
|
bab5572a74 | ||
|
|
9b4a5df925 | ||
|
|
1816a020c4 | ||
|
|
aa13928b76 | ||
|
|
d681670a9d | ||
|
|
4f6247b5c3 | ||
|
|
1e5b0a5c89 | ||
|
|
b65ec8b68a | ||
|
|
e13bcdbd44 | ||
|
|
46568f7af8 | ||
|
|
3358138ccc | ||
|
|
6d15d0729a | ||
|
|
0936918d24 | ||
|
|
d4a4418969 | ||
|
|
31ed241a1a | ||
|
|
178e6a8530 | ||
|
|
96963b6751 | ||
|
|
9770814f39 | ||
|
|
b0c1ed31b8 | ||
|
|
8435682ac8 | ||
|
|
ed682fad87 | ||
|
|
a12a3e4603 | ||
|
|
dc89b471fa | ||
|
|
0e520ee9f5 | ||
|
|
9620890b65 | ||
|
|
69c35e31b4 | ||
|
|
b6cd8944b1 | ||
|
|
6fc05efa6c | ||
|
|
3cda14d7f2 | ||
|
|
d9392fdbb8 | ||
|
|
82fdfa84b8 | ||
|
|
ca80e49aa0 | ||
|
|
51f17bf919 | ||
|
|
724c6a06e6 | ||
|
|
d64e9542eb | ||
|
|
3fb26d458e | ||
|
|
a411fd300c | ||
|
|
92a5ed19d3 | ||
|
|
be6777427d | ||
|
|
a1c9427d80 | ||
|
|
a59d7fba95 | ||
|
|
41d7bf4bbd | ||
|
|
b7b0d5796e | ||
|
|
21fc5b81f6 | ||
|
|
9906b9b1a1 | ||
|
|
3f8c447378 | ||
|
|
6e3f4d8722 | ||
|
|
877e889250 | ||
|
|
099ae4bc6c | ||
|
|
63d60ba490 | ||
|
|
d15830a2d0 | ||
|
|
141f3d0390 | ||
|
|
e1a24376ab | ||
|
|
62fc8d254e | ||
|
|
3a2f773d65 | ||
|
|
8f901f8899 | ||
|
|
c6bf5fbbfb | ||
|
|
f0eed7564f | ||
|
|
277d8e4c53 | ||
|
|
e70a69bbcf | ||
|
|
a48618c074 | ||
|
|
39193396f5 | ||
|
|
5343402385 | ||
|
|
62703ca23e | ||
|
|
cc64a93953 | ||
|
|
831325d6e2 | ||
|
|
8f64173574 | ||
|
|
76877e83c4 | ||
|
|
ecd398d895 | ||
|
|
aa92ad3fb1 | ||
|
|
fd94fdb42b | ||
|
|
30d15ecc3d | ||
|
|
3d87547d95 | ||
|
|
4d4cc551fd | ||
|
|
8e02154bf5 | ||
|
|
08e46aa62f | ||
|
|
e0c25ba4ba | ||
|
|
2560c6bd6c | ||
|
|
96ac15d292 | ||
|
|
488bbcb22b | ||
|
|
b7bbb44286 | ||
|
|
58318481e6 | ||
|
|
7cd5c1732b | ||
|
|
816d80602f | ||
|
|
d0d6dd4b0c | ||
|
|
47352e6e45 | ||
|
|
91acb8147c | ||
|
|
c9d387bd0d | ||
|
|
3aa6c02b93 | ||
|
|
f6900fb07c |
18
.coderabbit.yaml
Normal file
18
.coderabbit.yaml
Normal file
@@ -0,0 +1,18 @@
|
|||||||
|
# yaml-language-server: $schema=https://coderabbit.ai/integrations/schema.v2.json
|
||||||
|
language: en-US
|
||||||
|
reviews:
|
||||||
|
profile: chill
|
||||||
|
request_changes_workflow: false
|
||||||
|
high_level_summary: true
|
||||||
|
poem: false
|
||||||
|
review_status: true
|
||||||
|
auto_review:
|
||||||
|
enabled: true
|
||||||
|
drafts: false
|
||||||
|
path_filters:
|
||||||
|
- "!**/*.tsx"
|
||||||
|
- "!**/*.ts"
|
||||||
|
- "!**/*.js"
|
||||||
|
- "!**/*.svg"
|
||||||
|
chat:
|
||||||
|
auto_reply: true
|
||||||
@@ -6,7 +6,6 @@ RUN apt-get update && export DEBIAN_FRONTEND=noninteractive \
|
|||||||
iptables=1.8.9-2 \
|
iptables=1.8.9-2 \
|
||||||
libgl1-mesa-dev=22.3.6-1+deb12u1 \
|
libgl1-mesa-dev=22.3.6-1+deb12u1 \
|
||||||
xorg-dev=1:7.7+23 \
|
xorg-dev=1:7.7+23 \
|
||||||
libayatana-appindicator3-dev=0.5.92-1 \
|
|
||||||
&& apt-get clean \
|
&& apt-get clean \
|
||||||
&& rm -rf /var/lib/apt/lists/* \
|
&& rm -rf /var/lib/apt/lists/* \
|
||||||
&& go install -v golang.org/x/tools/gopls@latest
|
&& go install -v golang.org/x/tools/gopls@latest
|
||||||
|
|||||||
12
.github/workflows/agent-network-e2e.yml
vendored
12
.github/workflows/agent-network-e2e.yml
vendored
@@ -5,6 +5,13 @@ on:
|
|||||||
schedule:
|
schedule:
|
||||||
- cron: "0 3 * * *"
|
- cron: "0 3 * * *"
|
||||||
workflow_dispatch:
|
workflow_dispatch:
|
||||||
|
inputs:
|
||||||
|
bedrock_model:
|
||||||
|
description: >-
|
||||||
|
Bedrock inference-profile id to drive the matrix with, exactly as
|
||||||
|
AWS issues it. Leave empty for the Sonnet 4.6 default.
|
||||||
|
required: false
|
||||||
|
default: ""
|
||||||
|
|
||||||
concurrency:
|
concurrency:
|
||||||
group: ${{ github.workflow }}-${{ github.ref }}
|
group: ${{ github.workflow }}-${{ github.ref }}
|
||||||
@@ -51,6 +58,9 @@ jobs:
|
|||||||
# token (and URL, for gateways) is unset, so partial coverage is fine.
|
# token (and URL, for gateways) is unset, so partial coverage is fine.
|
||||||
OPENAI_TOKEN: ${{ secrets.E2E_OPENAI_TOKEN }}
|
OPENAI_TOKEN: ${{ secrets.E2E_OPENAI_TOKEN }}
|
||||||
ANTHROPIC_TOKEN: ${{ secrets.E2E_ANTHROPIC_TOKEN }}
|
ANTHROPIC_TOKEN: ${{ secrets.E2E_ANTHROPIC_TOKEN }}
|
||||||
|
# Moonshot AI platform key (platform.kimi.ai); drives both Kimi wire
|
||||||
|
# shapes (OpenAI /v1 and Anthropic /anthropic) through kimi_api.
|
||||||
|
KIMI_TOKEN: ${{ secrets.E2E_KIMI_TOKEN }}
|
||||||
VERCEL_URL: ${{ secrets.E2E_VERCEL_URL }}
|
VERCEL_URL: ${{ secrets.E2E_VERCEL_URL }}
|
||||||
VERCEL_TOKEN: ${{ secrets.E2E_VERCEL_TOKEN }}
|
VERCEL_TOKEN: ${{ secrets.E2E_VERCEL_TOKEN }}
|
||||||
OPENROUTER_URL: ${{ secrets.E2E_OPENROUTER_URL }}
|
OPENROUTER_URL: ${{ secrets.E2E_OPENROUTER_URL }}
|
||||||
@@ -59,6 +69,8 @@ jobs:
|
|||||||
CLOUDFLARE_TOKEN: ${{ secrets.E2E_CLOUDFLARE_TOKEN }}
|
CLOUDFLARE_TOKEN: ${{ secrets.E2E_CLOUDFLARE_TOKEN }}
|
||||||
AWS_BEARER_TOKEN_BEDROCK: ${{ secrets.E2E_AWS_BEARER_TOKEN_BEDROCK }}
|
AWS_BEARER_TOKEN_BEDROCK: ${{ secrets.E2E_AWS_BEARER_TOKEN_BEDROCK }}
|
||||||
AWS_REGION: ${{ secrets.E2E_AWS_REGION }}
|
AWS_REGION: ${{ secrets.E2E_AWS_REGION }}
|
||||||
|
# Bedrock model override: dispatch input wins, then the repo variable, else the test default.
|
||||||
|
AWS_BEDROCK_MODEL: ${{ inputs.bedrock_model || vars.E2E_AWS_BEDROCK_MODEL }}
|
||||||
# Vertex (Anthropic-on-Vertex): SA + project required; region defaults
|
# Vertex (Anthropic-on-Vertex): SA + project required; region defaults
|
||||||
# to "global", model to a pinned claude snapshot.
|
# to "global", model to a pinned claude snapshot.
|
||||||
GOOGLE_VERTEX_SA_BASE64: ${{ secrets.E2E_GOOGLE_VERTEX_SA_BASE64 }}
|
GOOGLE_VERTEX_SA_BASE64: ${{ secrets.E2E_GOOGLE_VERTEX_SA_BASE64 }}
|
||||||
|
|||||||
98
.github/workflows/frontend-ui.yml
vendored
Normal file
98
.github/workflows/frontend-ui.yml
vendored
Normal file
@@ -0,0 +1,98 @@
|
|||||||
|
name: UI Frontend
|
||||||
|
|
||||||
|
on:
|
||||||
|
pull_request:
|
||||||
|
paths:
|
||||||
|
- "client/ui/frontend/**"
|
||||||
|
- "client/ui/i18n/**"
|
||||||
|
- "client/ui/**/*.go"
|
||||||
|
- ".github/workflows/frontend-ui.yml"
|
||||||
|
push:
|
||||||
|
branches:
|
||||||
|
- main
|
||||||
|
paths:
|
||||||
|
- "client/ui/frontend/**"
|
||||||
|
- "client/ui/i18n/**"
|
||||||
|
- "client/ui/**/*.go"
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
concurrency:
|
||||||
|
group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }}
|
||||||
|
cancel-in-progress: true
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
lint-and-build:
|
||||||
|
name: Lint & Build
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
timeout-minutes: 15
|
||||||
|
defaults:
|
||||||
|
run:
|
||||||
|
working-directory: client/ui/frontend
|
||||||
|
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"
|
||||||
|
|
||||||
|
- name: Set up pnpm
|
||||||
|
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
|
||||||
|
with:
|
||||||
|
version: 11
|
||||||
|
|
||||||
|
# Bindings are generated by wails3 from the Go service definitions and
|
||||||
|
# are not checked in (see client/ui/frontend/bindings/). Without them,
|
||||||
|
# typecheck/build fail on missing module imports.
|
||||||
|
- name: Set up Go
|
||||||
|
uses: actions/setup-go@4b73464bb391d4059bd26b0524d20df3927bd417 # v6.3.0
|
||||||
|
with:
|
||||||
|
go-version-file: "go.mod"
|
||||||
|
cache: false
|
||||||
|
|
||||||
|
# wails3 CLI links against GTK4 / WebKitGTK 6.0 via its internal/operatingsystem
|
||||||
|
# package, so the dev libraries must be present before `go install`.
|
||||||
|
- name: Install Wails Linux system dependencies
|
||||||
|
run: |
|
||||||
|
sudo apt-get update
|
||||||
|
sudo apt-get install -y --no-install-recommends \
|
||||||
|
pkg-config \
|
||||||
|
libgtk-4-dev \
|
||||||
|
libwebkitgtk-6.0-dev
|
||||||
|
|
||||||
|
- name: Install wails3 CLI
|
||||||
|
# Version derived from go.mod so the binding generator always matches
|
||||||
|
# the wails runtime the daemon links against.
|
||||||
|
working-directory: ${{ github.workspace }}
|
||||||
|
run: |
|
||||||
|
WAILS_VERSION=$(go list -m -f '{{.Version}}' github.com/wailsapp/wails/v3)
|
||||||
|
go install github.com/wailsapp/wails/v3/cmd/wails3@$WAILS_VERSION
|
||||||
|
|
||||||
|
- name: Get pnpm store directory
|
||||||
|
id: pnpm-store
|
||||||
|
run: echo "path=$(pnpm store path --silent)" >> "$GITHUB_OUTPUT"
|
||||||
|
|
||||||
|
- name: Cache pnpm store
|
||||||
|
uses: actions/cache@v4
|
||||||
|
with:
|
||||||
|
path: ${{ steps.pnpm-store.outputs.path }}
|
||||||
|
key: ${{ runner.os }}-pnpm-${{ hashFiles('client/ui/frontend/pnpm-lock.yaml') }}
|
||||||
|
restore-keys: |
|
||||||
|
${{ runner.os }}-pnpm-
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: pnpm install --frozen-lockfile --ignore-scripts
|
||||||
|
|
||||||
|
- name: Generate Wails bindings
|
||||||
|
run: pnpm run bindings
|
||||||
|
|
||||||
|
- name: Lint, typecheck, format
|
||||||
|
run: pnpm check
|
||||||
|
|
||||||
|
- name: Build
|
||||||
|
run: pnpm build
|
||||||
10
.github/workflows/golang-test-darwin.yml
vendored
10
.github/workflows/golang-test-darwin.yml
vendored
@@ -45,7 +45,15 @@ jobs:
|
|||||||
run: git --no-pager diff --exit-code
|
run: git --no-pager diff --exit-code
|
||||||
|
|
||||||
- name: Test
|
- name: Test
|
||||||
run: NETBIRD_STORE_ENGINE=${{ matrix.store }} CI=true go test -coverprofile=coverage.txt -tags 'devcert privileged' -exec 'sudo --preserve-env=CI,NETBIRD_STORE_ENGINE' -timeout 5m -p 1 $(go list ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/testutil/privileged)
|
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
|
||||||
|
# which fails to compile until the frontend has been built. The Wails UI
|
||||||
|
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
|
||||||
|
# before goreleaser.
|
||||||
|
# `go list -e` lets the listing succeed even though the embed fails to
|
||||||
|
# resolve; the grep then drops the broken package by path. Without -e,
|
||||||
|
# go list aborts with empty stdout and `go test` falls back to the repo
|
||||||
|
# root, which has no Go files.
|
||||||
|
run: NETBIRD_STORE_ENGINE=${{ matrix.store }} CI=true go test -coverprofile=coverage.txt -tags 'devcert privileged' -exec 'sudo --preserve-env=CI,NETBIRD_STORE_ENGINE' -timeout 5m -p 1 $(go list -e ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/ui -e /client/testutil/privileged)
|
||||||
|
|
||||||
- name: Upload coverage reports to Codecov
|
- name: Upload coverage reports to Codecov
|
||||||
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f #v7.0.0
|
uses: codecov/codecov-action@fb8b3582c8e4def4969c97caa2f19720cb33a72f #v7.0.0
|
||||||
|
|||||||
17
.github/workflows/golang-test-linux.yml
vendored
17
.github/workflows/golang-test-linux.yml
vendored
@@ -53,7 +53,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
if: steps.cache.outputs.cache-hit != 'true'
|
if: steps.cache.outputs.cache-hit != 'true'
|
||||||
run: sudo apt update && sudo apt install -y -q libgtk-3-dev libayatana-appindicator3-dev libgl1-mesa-dev xorg-dev gcc-multilib libpcap-dev
|
run: sudo apt update && sudo apt install -y -q libgtk-4-dev libwebkitgtk-6.0-dev libsoup-3.0-dev libgl1-mesa-dev xorg-dev gcc-multilib libpcap-dev
|
||||||
|
|
||||||
- name: Install 32-bit libpcap
|
- name: Install 32-bit libpcap
|
||||||
if: steps.cache.outputs.cache-hit != 'true'
|
if: steps.cache.outputs.cache-hit != 'true'
|
||||||
@@ -145,7 +145,7 @@ jobs:
|
|||||||
${{ runner.os }}-gotest-cache-
|
${{ runner.os }}-gotest-cache-
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: sudo apt update && sudo apt install -y -q libgtk-3-dev libayatana-appindicator3-dev libgl1-mesa-dev xorg-dev gcc-multilib libpcap-dev
|
run: sudo apt update && sudo apt install -y -q libgtk-4-dev libwebkitgtk-6.0-dev libsoup-3.0-dev libgl1-mesa-dev xorg-dev gcc-multilib libpcap-dev
|
||||||
|
|
||||||
- name: Install 32-bit libpcap
|
- name: Install 32-bit libpcap
|
||||||
if: matrix.arch == '386'
|
if: matrix.arch == '386'
|
||||||
@@ -158,7 +158,15 @@ jobs:
|
|||||||
run: git --no-pager diff --exit-code
|
run: git --no-pager diff --exit-code
|
||||||
|
|
||||||
- name: Test
|
- name: Test
|
||||||
run: CGO_ENABLED=1 GOARCH=${{ matrix.arch }} CI=true go test -coverprofile=coverage.txt -tags 'devcert privileged' -exec 'sudo --preserve-env=CI,CGO_ENABLED' -timeout 10m -p 1 $(go list ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/testutil/privileged)
|
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
|
||||||
|
# which fails to compile until the frontend has been built. The Wails UI
|
||||||
|
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
|
||||||
|
# before goreleaser.
|
||||||
|
# `go list -e` lets the listing succeed even though the embed fails to
|
||||||
|
# resolve; the grep then drops the broken package by path. Without -e,
|
||||||
|
# go list aborts with empty stdout and `go test` falls back to the repo
|
||||||
|
# root, which has no Go files.
|
||||||
|
run: CGO_ENABLED=1 GOARCH=${{ matrix.arch }} CI=true go test -coverprofile=coverage.txt -tags 'devcert privileged' -exec 'sudo --preserve-env=CI,CGO_ENABLED' -timeout 10m -p 1 $(go list -e ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/ui -e /client/testutil/privileged)
|
||||||
|
|
||||||
- name: Upload coverage reports to Codecov
|
- name: Upload coverage reports to Codecov
|
||||||
if: matrix.arch == 'amd64'
|
if: matrix.arch == 'amd64'
|
||||||
@@ -168,7 +176,6 @@ jobs:
|
|||||||
slug: netbirdio/netbird
|
slug: netbirdio/netbird
|
||||||
flags: unit,client
|
flags: unit,client
|
||||||
|
|
||||||
|
|
||||||
test_client_on_docker:
|
test_client_on_docker:
|
||||||
name: "Client (Docker) / Unit"
|
name: "Client (Docker) / Unit"
|
||||||
needs: [build-cache]
|
needs: [build-cache]
|
||||||
@@ -229,7 +236,7 @@ jobs:
|
|||||||
sh -c ' \
|
sh -c ' \
|
||||||
apk update; apk add --no-cache \
|
apk update; apk add --no-cache \
|
||||||
ca-certificates iptables ip6tables dbus dbus-dev libpcap-dev build-base; \
|
ca-certificates iptables ip6tables dbus dbus-dev libpcap-dev build-base; \
|
||||||
go test -buildvcs=false -tags "devcert privileged" -v -timeout 10m -p 1 $(go list -buildvcs=false ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/ui -e /upload-server -e /client/testutil/privileged)
|
go test -buildvcs=false -tags "devcert privileged" -v -timeout 10m -p 1 $(go list -e -buildvcs=false ./... | grep -v -e /management -e /signal -e /relay -e /proxy -e /combined -e /client/ui -e /upload-server -e /client/testutil/privileged)
|
||||||
'
|
'
|
||||||
|
|
||||||
test_relay:
|
test_relay:
|
||||||
|
|||||||
9
.github/workflows/golang-test-windows.yml
vendored
9
.github/workflows/golang-test-windows.yml
vendored
@@ -65,8 +65,15 @@ jobs:
|
|||||||
- run: PsExec64 -s -w ${{ github.workspace }} C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe env -w GOCACHE=${{ env.modcache }}
|
- run: PsExec64 -s -w ${{ github.workspace }} C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe env -w GOCACHE=${{ env.modcache }}
|
||||||
- run: PsExec64 -s -w ${{ github.workspace }} C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe mod tidy
|
- run: PsExec64 -s -w ${{ github.workspace }} C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe mod tidy
|
||||||
- name: Generate test script
|
- name: Generate test script
|
||||||
|
# Exclude client/ui: its main.go uses //go:embed all:frontend/dist,
|
||||||
|
# which fails to compile until the frontend has been built. The Wails UI
|
||||||
|
# has no Go-side unit tests, and its release pipeline runs `pnpm build`
|
||||||
|
# before goreleaser.
|
||||||
|
# `go list -e` lets the listing succeed even though the embed fails to
|
||||||
|
# resolve; the Where-Object pipeline then drops the broken package by
|
||||||
|
# path. Without -e, go list aborts with empty stdout.
|
||||||
run: |
|
run: |
|
||||||
$packages = go list ./... | Where-Object { $_ -notmatch '/management' } | Where-Object { $_ -notmatch '/relay' } | Where-Object { $_ -notmatch '/signal' } | Where-Object { $_ -notmatch '/proxy' } | Where-Object { $_ -notmatch '/combined' }
|
$packages = go list -e ./... | Where-Object { $_ -notmatch '/management' } | Where-Object { $_ -notmatch '/relay' } | Where-Object { $_ -notmatch '/signal' } | Where-Object { $_ -notmatch '/proxy' } | Where-Object { $_ -notmatch '/combined' } | Where-Object { $_ -notmatch '/client/ui' }
|
||||||
$goExe = "C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe"
|
$goExe = "C:\hostedtoolcache\windows\go\${{ steps.go.outputs.go-version }}\x64\bin\go.exe"
|
||||||
$cmd = "$goExe test -tags `"devcert privileged`" -timeout 10m -p 1 $($packages -join ' ') > test-out.txt 2>&1"
|
$cmd = "$goExe test -tags `"devcert privileged`" -timeout 10m -p 1 $($packages -join ' ') > test-out.txt 2>&1"
|
||||||
Set-Content -Path "${{ github.workspace }}\run-tests.cmd" -Value $cmd
|
Set-Content -Path "${{ github.workspace }}\run-tests.cmd" -Value $cmd
|
||||||
|
|||||||
25
.github/workflows/golangci-lint.yml
vendored
25
.github/workflows/golangci-lint.yml
vendored
@@ -22,7 +22,15 @@ jobs:
|
|||||||
uses: codespell-project/actions-codespell@8f01853be192eb0f849a5c7d721450e7a467c579 # v2.2
|
uses: codespell-project/actions-codespell@8f01853be192eb0f849a5c7d721450e7a467c579 # v2.2
|
||||||
with:
|
with:
|
||||||
ignore_words_list: erro,clienta,hastable,iif,groupd,testin,groupe,cros,ans,deriver,te,userA,ede,additionals,flate,recordin,unparseable
|
ignore_words_list: erro,clienta,hastable,iif,groupd,testin,groupe,cros,ans,deriver,te,userA,ede,additionals,flate,recordin,unparseable
|
||||||
skip: go.mod,go.sum,**/proxy/web/**
|
# Non-English UI translations trip codespell on real foreign words
|
||||||
|
# (de: "Sie", "oder", "ist"). Only en/common.json is the source of
|
||||||
|
# truth that should be spell-checked. List each translated locale
|
||||||
|
# dir below and add new ones as languages are added under
|
||||||
|
# client/ui/i18n/locales/. Single-star globs are matched per path
|
||||||
|
# segment by codespell and behave the same across versions; the
|
||||||
|
# recursive "**" form did not take effect with the codespell shipped
|
||||||
|
# by this action.
|
||||||
|
skip: go.mod,go.sum,*/proxy/web/*,*pnpm-lock.yaml,*package-lock.json,*/locales/de/*,*/locales/es/*,*/locales/fr/*,*/locales/hu/*,*/locales/it/*,*/locales/pt/*,*/locales/ru/*,*/locales/zh-CN/*,*/i18n/TRANSLATING.md
|
||||||
golangci:
|
golangci:
|
||||||
strategy:
|
strategy:
|
||||||
fail-fast: false
|
fail-fast: false
|
||||||
@@ -37,7 +45,7 @@ jobs:
|
|||||||
display_name: Linux
|
display_name: Linux
|
||||||
name: ${{ matrix.display_name }}
|
name: ${{ matrix.display_name }}
|
||||||
runs-on: ${{ matrix.os }}
|
runs-on: ${{ matrix.os }}
|
||||||
timeout-minutes: 15
|
timeout-minutes: 25
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||||
@@ -54,7 +62,16 @@ jobs:
|
|||||||
cache: false
|
cache: false
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
if: matrix.os == 'ubuntu-latest'
|
if: matrix.os == 'ubuntu-latest'
|
||||||
run: sudo apt update && sudo apt install -y -q libgtk-3-dev libayatana-appindicator3-dev libgl1-mesa-dev xorg-dev libpcap-dev
|
run: sudo apt update && sudo apt install -y -q libgtk-4-dev libwebkitgtk-6.0-dev libsoup-3.0-dev libgl1-mesa-dev xorg-dev libpcap-dev
|
||||||
|
- name: Stub Wails frontend bundle
|
||||||
|
# client/ui/main.go has //go:embed all:frontend/dist. The
|
||||||
|
# directory is produced by `pnpm run build` and is gitignored, so
|
||||||
|
# lint-only runs (no frontend toolchain) need a placeholder file
|
||||||
|
# for the embed pattern to match.
|
||||||
|
shell: bash
|
||||||
|
run: |
|
||||||
|
mkdir -p client/ui/frontend/dist
|
||||||
|
touch client/ui/frontend/dist/.embed-placeholder
|
||||||
- name: golangci-lint
|
- name: golangci-lint
|
||||||
uses: golangci/golangci-lint-action@82606bf257cbaff209d206a39f5134f0cfbfd2ee #v9.2.1
|
uses: golangci/golangci-lint-action@82606bf257cbaff209d206a39f5134f0cfbfd2ee #v9.2.1
|
||||||
with:
|
with:
|
||||||
@@ -62,4 +79,4 @@ jobs:
|
|||||||
skip-cache: true
|
skip-cache: true
|
||||||
skip-save-cache: true
|
skip-save-cache: true
|
||||||
cache-invalidation-interval: 0
|
cache-invalidation-interval: 0
|
||||||
args: --timeout=12m
|
args: --timeout=20m
|
||||||
|
|||||||
81
.github/workflows/release.yml
vendored
81
.github/workflows/release.yml
vendored
@@ -9,7 +9,7 @@ on:
|
|||||||
pull_request:
|
pull_request:
|
||||||
|
|
||||||
env:
|
env:
|
||||||
SIGN_PIPE_VER: "v0.1.6"
|
SIGN_PIPE_VER: "v0.1.8"
|
||||||
GORELEASER_VER: "v2.16.0"
|
GORELEASER_VER: "v2.16.0"
|
||||||
PRODUCT_NAME: "NetBird"
|
PRODUCT_NAME: "NetBird"
|
||||||
COPYRIGHT: "NetBird GmbH"
|
COPYRIGHT: "NetBird GmbH"
|
||||||
@@ -216,9 +216,9 @@ jobs:
|
|||||||
- name: Install goversioninfo
|
- name: Install goversioninfo
|
||||||
run: go install github.com/josephspurrier/goversioninfo/cmd/goversioninfo@233067e
|
run: go install github.com/josephspurrier/goversioninfo/cmd/goversioninfo@233067e
|
||||||
- name: Generate windows syso amd64
|
- name: Generate windows syso amd64
|
||||||
run: goversioninfo -icon client/ui/assets/netbird.ico -manifest client/manifest.xml -product-name ${{ env.PRODUCT_NAME }} -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/resources_windows_amd64.syso
|
run: goversioninfo -icon client/ui/build/windows/icon.ico -manifest client/manifest.xml -product-name ${{ env.PRODUCT_NAME }} -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/resources_windows_amd64.syso
|
||||||
- name: Generate windows syso arm64
|
- name: Generate windows syso arm64
|
||||||
run: goversioninfo -arm -64 -icon client/ui/assets/netbird.ico -manifest client/manifest.xml -product-name ${{ env.PRODUCT_NAME }} -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/resources_windows_arm64.syso
|
run: goversioninfo -arm -64 -icon client/ui/build/windows/icon.ico -manifest client/manifest.xml -product-name ${{ env.PRODUCT_NAME }} -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/resources_windows_arm64.syso
|
||||||
- name: Run GoReleaser
|
- name: Run GoReleaser
|
||||||
id: goreleaser
|
id: goreleaser
|
||||||
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
|
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
|
||||||
@@ -397,8 +397,18 @@ jobs:
|
|||||||
- name: check git status
|
- name: check git status
|
||||||
run: git --no-pager diff --exit-code
|
run: git --no-pager diff --exit-code
|
||||||
|
|
||||||
|
- name: Set up Node.js
|
||||||
|
uses: actions/setup-node@v4
|
||||||
|
with:
|
||||||
|
node-version: '22'
|
||||||
|
|
||||||
|
- name: Set up pnpm
|
||||||
|
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
|
||||||
|
with:
|
||||||
|
version: 11
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: sudo apt update && sudo apt install -y -q libappindicator3-dev gir1.2-appindicator3-0.1 libxxf86vm-dev gcc-mingw-w64-x86-64
|
run: sudo apt update && sudo apt install -y -q libgtk-4-dev libwebkitgtk-6.0-dev libsoup-3.0-dev gcc-mingw-w64-x86-64
|
||||||
|
|
||||||
- name: Decode GPG signing key
|
- name: Decode GPG signing key
|
||||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name == github.repository
|
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name == github.repository
|
||||||
@@ -417,10 +427,16 @@ jobs:
|
|||||||
echo "/tmp/llvm-mingw-20250709-ucrt-ubuntu-22.04-x86_64/bin" >> $GITHUB_PATH
|
echo "/tmp/llvm-mingw-20250709-ucrt-ubuntu-22.04-x86_64/bin" >> $GITHUB_PATH
|
||||||
- name: Install goversioninfo
|
- name: Install goversioninfo
|
||||||
run: go install github.com/josephspurrier/goversioninfo/cmd/goversioninfo@233067e
|
run: go install github.com/josephspurrier/goversioninfo/cmd/goversioninfo@233067e
|
||||||
|
- name: Install wails3 CLI
|
||||||
|
# Version derived from go.mod so the binding generator always matches
|
||||||
|
# the wails runtime the binary links against.
|
||||||
|
run: |
|
||||||
|
WAILS_VERSION=$(go list -m -f '{{.Version}}' github.com/wailsapp/wails/v3)
|
||||||
|
go install github.com/wailsapp/wails/v3/cmd/wails3@$WAILS_VERSION
|
||||||
- name: Generate windows syso amd64
|
- name: Generate windows syso amd64
|
||||||
run: goversioninfo -64 -icon client/ui/assets/netbird.ico -manifest client/ui/manifest.xml -product-name ${{ env.PRODUCT_NAME }}-"UI" -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/ui/resources_windows_amd64.syso
|
run: goversioninfo -64 -icon client/ui/build/windows/icon.ico -manifest client/ui/build/windows/wails.exe.manifest -product-name ${{ env.PRODUCT_NAME }}-"UI" -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/ui/resources_windows_amd64.syso
|
||||||
- name: Generate windows syso arm64
|
- name: Generate windows syso arm64
|
||||||
run: goversioninfo -arm -64 -icon client/ui/assets/netbird.ico -manifest client/ui/manifest.xml -product-name ${{ env.PRODUCT_NAME }}-"UI" -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/ui/resources_windows_arm64.syso
|
run: goversioninfo -arm -64 -icon client/ui/build/windows/icon.ico -manifest client/ui/build/windows/wails.exe.manifest -product-name ${{ env.PRODUCT_NAME }}-"UI" -copyright "${{ env.COPYRIGHT }}" -ver-major ${{ steps.semver_parser.outputs.major }} -ver-minor ${{ steps.semver_parser.outputs.minor }} -ver-patch ${{ steps.semver_parser.outputs.patch }} -ver-build 0 -file-version ${{ steps.semver_parser.outputs.fullversion }}.0 -product-version ${{ steps.semver_parser.outputs.fullversion }}.0 -o client/ui/resources_windows_arm64.syso
|
||||||
|
|
||||||
- name: Run GoReleaser
|
- name: Run GoReleaser
|
||||||
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
|
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
|
||||||
@@ -489,6 +505,20 @@ jobs:
|
|||||||
run: go mod tidy
|
run: go mod tidy
|
||||||
- name: check git status
|
- name: check git status
|
||||||
run: git --no-pager diff --exit-code
|
run: git --no-pager diff --exit-code
|
||||||
|
- name: Set up Node.js
|
||||||
|
uses: actions/setup-node@v4
|
||||||
|
with:
|
||||||
|
node-version: '22'
|
||||||
|
- name: Set up pnpm
|
||||||
|
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
|
||||||
|
with:
|
||||||
|
version: 11
|
||||||
|
- name: Install wails3 CLI
|
||||||
|
# Version derived from go.mod so the binding generator always matches
|
||||||
|
# the wails runtime the binary links against.
|
||||||
|
run: |
|
||||||
|
WAILS_VERSION=$(go list -m -f '{{.Version}}' github.com/wailsapp/wails/v3)
|
||||||
|
go install github.com/wailsapp/wails/v3/cmd/wails3@$WAILS_VERSION
|
||||||
- name: Run GoReleaser
|
- name: Run GoReleaser
|
||||||
id: goreleaser
|
id: goreleaser
|
||||||
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
|
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
|
||||||
@@ -576,23 +606,6 @@ jobs:
|
|||||||
- name: Move wintun.dll into dist
|
- name: Move wintun.dll into dist
|
||||||
run: mv ${{ env.downloadPath }}\wintun\bin\${{ matrix.wintun_arch }}\wintun.dll ${{ github.workspace }}\dist\${{ env.PackageWorkdir }}\
|
run: mv ${{ env.downloadPath }}\wintun\bin\${{ matrix.wintun_arch }}\wintun.dll ${{ github.workspace }}\dist\${{ env.PackageWorkdir }}\
|
||||||
|
|
||||||
- name: Download Mesa3D (amd64 only)
|
|
||||||
id: download-mesa3d
|
|
||||||
if: matrix.arch == 'amd64'
|
|
||||||
uses: netbirdio/shared-actions/actions/win-download-and-verify@be5df6047383da2236e02243cceb857d8567c27e # v0.0.2
|
|
||||||
with:
|
|
||||||
url: https://pkgs.netbird.io/mesa3d/MesaForWindows-x64-20.1.8.7z
|
|
||||||
destination: ${{ env.downloadPath }}\mesa3d.7z
|
|
||||||
sha256: 71c7cb64ec229a1d6b8d62fa08e1889ed2bd17c0eeede8689daf0f25cb31d6b9
|
|
||||||
|
|
||||||
- name: Extract Mesa3D driver (amd64 only)
|
|
||||||
if: matrix.arch == 'amd64'
|
|
||||||
run: 7z x -o"${{ env.downloadPath }}" "${{ env.downloadPath }}/mesa3d.7z"
|
|
||||||
|
|
||||||
- name: Move opengl32.dll into dist (amd64 only)
|
|
||||||
if: matrix.arch == 'amd64'
|
|
||||||
run: mv ${{ env.downloadPath }}\opengl32.dll ${{ github.workspace }}\dist\${{ env.PackageWorkdir }}\
|
|
||||||
|
|
||||||
- name: Download EnVar plugin for NSIS
|
- name: Download EnVar plugin for NSIS
|
||||||
uses: netbirdio/shared-actions/actions/win-download-and-verify@be5df6047383da2236e02243cceb857d8567c27e # v0.0.2
|
uses: netbirdio/shared-actions/actions/win-download-and-verify@be5df6047383da2236e02243cceb857d8567c27e # v0.0.2
|
||||||
with:
|
with:
|
||||||
@@ -615,6 +628,28 @@ jobs:
|
|||||||
if: matrix.arch == 'amd64'
|
if: matrix.arch == 'amd64'
|
||||||
run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z"
|
run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z"
|
||||||
|
|
||||||
|
- name: Set up Go for wails3 CLI
|
||||||
|
uses: actions/setup-go@v5
|
||||||
|
with:
|
||||||
|
go-version-file: "go.mod"
|
||||||
|
cache: false
|
||||||
|
|
||||||
|
- name: Install wails3 CLI
|
||||||
|
# Version derived from go.mod so the bootstrapper payload always
|
||||||
|
# matches the wails runtime the binary links against.
|
||||||
|
shell: bash
|
||||||
|
run: |
|
||||||
|
WAILS_VERSION=$(go list -m -f '{{.Version}}' github.com/wailsapp/wails/v3)
|
||||||
|
go install github.com/wailsapp/wails/v3/cmd/wails3@$WAILS_VERSION
|
||||||
|
|
||||||
|
- name: Stage WebView2 bootstrapper for installers
|
||||||
|
# Both client/installer.nsis and client/netbird.wxs reference
|
||||||
|
# client/MicrosoftEdgeWebview2Setup.exe. wails3 writes it there.
|
||||||
|
# The signing pipeline (netbirdio/sign-pipelines) does the same
|
||||||
|
# step for release builds; this mirrors it for PR sanity testing.
|
||||||
|
shell: bash
|
||||||
|
run: wails3 generate webview2bootstrapper -dir client
|
||||||
|
|
||||||
- name: Build NSIS installer
|
- name: Build NSIS installer
|
||||||
shell: pwsh
|
shell: pwsh
|
||||||
env:
|
env:
|
||||||
|
|||||||
87
.github/workflows/test-infrastructure-files.yml
vendored
87
.github/workflows/test-infrastructure-files.yml
vendored
@@ -249,78 +249,35 @@ jobs:
|
|||||||
docker compose exec management ls -l /var/lib/netbird/ | grep -i GeoLite2-City_[0-9]*.mmdb
|
docker compose exec management ls -l /var/lib/netbird/ | grep -i GeoLite2-City_[0-9]*.mmdb
|
||||||
docker compose exec management ls -l /var/lib/netbird/ | grep -i geonames_[0-9]*.db
|
docker compose exec management ls -l /var/lib/netbird/ | grep -i geonames_[0-9]*.db
|
||||||
|
|
||||||
test-getting-started-script:
|
test-legacy-getting-started-scripts:
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- name: Install jq
|
|
||||||
run: sudo apt-get install -y jq
|
|
||||||
|
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: run script with Zitadel PostgreSQL
|
- name: Verify Dex retirement notice
|
||||||
run: NETBIRD_DOMAIN=use-ip bash -x infrastructure_files/getting-started-with-zitadel.sh
|
|
||||||
|
|
||||||
- name: test Caddy file gen postgres
|
|
||||||
run: test -f Caddyfile
|
|
||||||
|
|
||||||
- name: test docker-compose file gen postgres
|
|
||||||
run: test -f docker-compose.yml
|
|
||||||
|
|
||||||
- name: test management.json file gen postgres
|
|
||||||
run: test -f management.json
|
|
||||||
|
|
||||||
- name: test turnserver.conf file gen postgres
|
|
||||||
run: |
|
run: |
|
||||||
set -x
|
if infrastructure_files/getting-started-with-dex.sh >stdout.txt 2>stderr.txt; then
|
||||||
test -f turnserver.conf
|
echo "Expected the retired Dex installer to fail"
|
||||||
grep external-ip turnserver.conf
|
exit 1
|
||||||
|
fi
|
||||||
|
test ! -s stdout.txt
|
||||||
|
grep -Fq "Dex support is not deprecated." stderr.txt
|
||||||
|
grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-quickstart" stderr.txt
|
||||||
|
grep -Fq "https://docs.netbird.io/selfhosted/identity-providers/local" stderr.txt
|
||||||
|
grep -Fq "removed in NetBird v0.80" stderr.txt
|
||||||
|
|
||||||
- name: test zitadel.env file gen postgres
|
- name: Verify Zitadel retirement notice
|
||||||
run: test -f zitadel.env
|
|
||||||
|
|
||||||
- name: test dashboard.env file gen postgres
|
|
||||||
run: test -f dashboard.env
|
|
||||||
|
|
||||||
- name: test relay.env file gen postgres
|
|
||||||
run: test -f relay.env
|
|
||||||
|
|
||||||
- name: test zdb.env file gen postgres
|
|
||||||
run: test -f zdb.env
|
|
||||||
|
|
||||||
- name: Postgres run cleanup
|
|
||||||
run: |
|
run: |
|
||||||
docker compose down --volumes --rmi all
|
if bash infrastructure_files/getting-started-with-zitadel.sh >stdout.txt 2>stderr.txt; then
|
||||||
rm -rf docker-compose.yml Caddyfile zitadel.env dashboard.env machinekey/zitadel-admin-sa.token turnserver.conf management.json zdb.env
|
echo "Expected the retired Zitadel installer to fail"
|
||||||
|
exit 1
|
||||||
- name: run script with Zitadel CockroachDB
|
fi
|
||||||
run: bash -x infrastructure_files/getting-started-with-zitadel.sh
|
test ! -s stdout.txt
|
||||||
env:
|
grep -Fq "Zitadel support and existing Zitadel deployments are not deprecated." stderr.txt
|
||||||
NETBIRD_DOMAIN: use-ip
|
grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-quickstart" stderr.txt
|
||||||
ZITADEL_DATABASE: cockroach
|
grep -Fq "https://docs.netbird.io/selfhosted/identity-providers/zitadel" stderr.txt
|
||||||
|
grep -Fq "https://docs.netbird.io/selfhosted/selfhosted-guide" stderr.txt
|
||||||
- name: test Caddy file gen CockroachDB
|
grep -Fq "removed in NetBird v0.80" stderr.txt
|
||||||
run: test -f Caddyfile
|
|
||||||
|
|
||||||
- name: test docker-compose file gen CockroachDB
|
|
||||||
run: test -f docker-compose.yml
|
|
||||||
|
|
||||||
- name: test management.json file gen CockroachDB
|
|
||||||
run: test -f management.json
|
|
||||||
|
|
||||||
- name: test turnserver.conf file gen CockroachDB
|
|
||||||
run: |
|
|
||||||
set -x
|
|
||||||
test -f turnserver.conf
|
|
||||||
grep external-ip turnserver.conf
|
|
||||||
|
|
||||||
- name: test zitadel.env file gen CockroachDB
|
|
||||||
run: test -f zitadel.env
|
|
||||||
|
|
||||||
- name: test dashboard.env file gen CockroachDB
|
|
||||||
run: test -f dashboard.env
|
|
||||||
|
|
||||||
- name: test relay.env file gen CockroachDB
|
|
||||||
run: test -f relay.env
|
|
||||||
|
|||||||
2
.github/workflows/wasm-build-validation.yml
vendored
2
.github/workflows/wasm-build-validation.yml
vendored
@@ -27,7 +27,7 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
go-version-file: "go.mod"
|
go-version-file: "go.mod"
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: sudo apt update && sudo apt install -y -q libgtk-3-dev libayatana-appindicator3-dev libgl1-mesa-dev xorg-dev libpcap-dev
|
run: sudo apt update && sudo apt install -y -q libgtk-4-dev libwebkitgtk-6.0-dev libsoup-3.0-dev libgl1-mesa-dev xorg-dev libpcap-dev
|
||||||
- name: Install golangci-lint
|
- name: Install golangci-lint
|
||||||
uses: golangci/golangci-lint-action@82606bf257cbaff209d206a39f5134f0cfbfd2ee #v9.2.1
|
uses: golangci/golangci-lint-action@82606bf257cbaff209d206a39f5134f0cfbfd2ee #v9.2.1
|
||||||
with:
|
with:
|
||||||
|
|||||||
@@ -114,6 +114,16 @@ linters:
|
|||||||
- linters:
|
- linters:
|
||||||
- staticcheck
|
- staticcheck
|
||||||
text: "QF1012"
|
text: "QF1012"
|
||||||
|
# client/ui/main.go uses //go:embed all:frontend/dist; the
|
||||||
|
# directory is populated by `pnpm build` in the release pipeline
|
||||||
|
# and missing at lint time, so the embed parses to "no matching
|
||||||
|
# files found" — surfaced by golangci-lint's typecheck pre-pass.
|
||||||
|
# Suppress just that one diagnostic; the rest of the package
|
||||||
|
# (services/, tray.go, grpc.go, ...) still gets linted normally.
|
||||||
|
- linters:
|
||||||
|
- typecheck
|
||||||
|
path: client/ui/main\.go
|
||||||
|
text: "pattern all:frontend/dist"
|
||||||
paths:
|
paths:
|
||||||
- third_party$
|
- third_party$
|
||||||
- builtin$
|
- builtin$
|
||||||
|
|||||||
@@ -212,6 +212,7 @@ nfpms:
|
|||||||
description: Netbird client.
|
description: Netbird client.
|
||||||
homepage: https://netbird.io/
|
homepage: https://netbird.io/
|
||||||
license: BSD-3-Clause
|
license: BSD-3-Clause
|
||||||
|
vendor: NetBird
|
||||||
id: netbird_deb
|
id: netbird_deb
|
||||||
bindir: /usr/bin
|
bindir: /usr/bin
|
||||||
builds:
|
builds:
|
||||||
@@ -226,6 +227,7 @@ nfpms:
|
|||||||
description: Netbird client.
|
description: Netbird client.
|
||||||
homepage: https://netbird.io/
|
homepage: https://netbird.io/
|
||||||
license: BSD-3-Clause
|
license: BSD-3-Clause
|
||||||
|
vendor: NetBird
|
||||||
id: netbird_rpm
|
id: netbird_rpm
|
||||||
bindir: /usr/bin
|
bindir: /usr/bin
|
||||||
builds:
|
builds:
|
||||||
@@ -271,8 +273,8 @@ dockers_v2:
|
|||||||
- netbirdio/netbird
|
- netbirdio/netbird
|
||||||
- ghcr.io/netbirdio/netbird
|
- ghcr.io/netbirdio/netbird
|
||||||
tags:
|
tags:
|
||||||
- "v{{ .Version }}-rootless"
|
- "{{ .Version }}-rootless"
|
||||||
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}"
|
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}rootless-latest{{ end }}"
|
||||||
dockerfile: client/Dockerfile-rootless
|
dockerfile: client/Dockerfile-rootless
|
||||||
extra_files:
|
extra_files:
|
||||||
- client/netbird-entrypoint.sh
|
- client/netbird-entrypoint.sh
|
||||||
|
|||||||
@@ -2,6 +2,15 @@ version: 2
|
|||||||
env:
|
env:
|
||||||
- SKIP_PUBLISH={{ if index .Env "SKIP_PUBLISH" }}{{ .Env.SKIP_PUBLISH }}{{ else }}true{{ end }}
|
- SKIP_PUBLISH={{ if index .Env "SKIP_PUBLISH" }}{{ .Env.SKIP_PUBLISH }}{{ else }}true{{ end }}
|
||||||
project_name: netbird-ui
|
project_name: netbird-ui
|
||||||
|
|
||||||
|
before:
|
||||||
|
hooks:
|
||||||
|
# Bindings are gitignored; regenerate before the frontend build so
|
||||||
|
# the @wailsio/runtime Vite plugin can resolve them (vite refuses to
|
||||||
|
# build without them).
|
||||||
|
- sh -c 'cd client/ui && wails3 generate bindings -clean=true -ts'
|
||||||
|
- sh -c 'cd client/ui/frontend && pnpm install --frozen-lockfile && pnpm build'
|
||||||
|
|
||||||
builds:
|
builds:
|
||||||
- id: netbird-ui
|
- id: netbird-ui
|
||||||
dir: client/ui
|
dir: client/ui
|
||||||
@@ -15,6 +24,8 @@ builds:
|
|||||||
ldflags:
|
ldflags:
|
||||||
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
||||||
mod_timestamp: "{{ .CommitTimestamp }}"
|
mod_timestamp: "{{ .CommitTimestamp }}"
|
||||||
|
tags:
|
||||||
|
- production
|
||||||
|
|
||||||
- id: netbird-ui-windows-amd64
|
- id: netbird-ui-windows-amd64
|
||||||
dir: client/ui
|
dir: client/ui
|
||||||
@@ -30,6 +41,8 @@ builds:
|
|||||||
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
||||||
- -H windowsgui
|
- -H windowsgui
|
||||||
mod_timestamp: "{{ .CommitTimestamp }}"
|
mod_timestamp: "{{ .CommitTimestamp }}"
|
||||||
|
tags:
|
||||||
|
- production
|
||||||
|
|
||||||
- id: netbird-ui-windows-arm64
|
- id: netbird-ui-windows-arm64
|
||||||
dir: client/ui
|
dir: client/ui
|
||||||
@@ -46,6 +59,8 @@ builds:
|
|||||||
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
||||||
- -H windowsgui
|
- -H windowsgui
|
||||||
mod_timestamp: "{{ .CommitTimestamp }}"
|
mod_timestamp: "{{ .CommitTimestamp }}"
|
||||||
|
tags:
|
||||||
|
- production
|
||||||
|
|
||||||
archives:
|
archives:
|
||||||
- id: linux-arch
|
- id: linux-arch
|
||||||
@@ -62,6 +77,8 @@ nfpms:
|
|||||||
- maintainer: Netbird <dev@netbird.io>
|
- maintainer: Netbird <dev@netbird.io>
|
||||||
description: Netbird client UI.
|
description: Netbird client UI.
|
||||||
homepage: https://netbird.io/
|
homepage: https://netbird.io/
|
||||||
|
license: BSD-3-Clause
|
||||||
|
vendor: NetBird
|
||||||
id: netbird_ui_deb
|
id: netbird_ui_deb
|
||||||
package_name: netbird-ui
|
package_name: netbird-ui
|
||||||
builds:
|
builds:
|
||||||
@@ -71,9 +88,9 @@ nfpms:
|
|||||||
scripts:
|
scripts:
|
||||||
postinstall: "release_files/ui-post-install.sh"
|
postinstall: "release_files/ui-post-install.sh"
|
||||||
contents:
|
contents:
|
||||||
- src: client/ui/build/netbird.desktop
|
- src: client/ui/build/linux/netbird.desktop
|
||||||
dst: /usr/share/applications/netbird.desktop
|
dst: /usr/share/applications/org.wails.netbird.desktop
|
||||||
- src: client/ui/assets/netbird.png
|
- src: client/ui/build/appicon.png
|
||||||
dst: /usr/share/pixmaps/netbird.png
|
dst: /usr/share/pixmaps/netbird.png
|
||||||
dependencies:
|
dependencies:
|
||||||
- netbird
|
- netbird
|
||||||
@@ -81,6 +98,8 @@ nfpms:
|
|||||||
- maintainer: Netbird <dev@netbird.io>
|
- maintainer: Netbird <dev@netbird.io>
|
||||||
description: Netbird client UI.
|
description: Netbird client UI.
|
||||||
homepage: https://netbird.io/
|
homepage: https://netbird.io/
|
||||||
|
license: BSD-3-Clause
|
||||||
|
vendor: NetBird
|
||||||
id: netbird_ui_rpm
|
id: netbird_ui_rpm
|
||||||
package_name: netbird-ui
|
package_name: netbird-ui
|
||||||
builds:
|
builds:
|
||||||
@@ -90,12 +109,13 @@ nfpms:
|
|||||||
scripts:
|
scripts:
|
||||||
postinstall: "release_files/ui-post-install.sh"
|
postinstall: "release_files/ui-post-install.sh"
|
||||||
contents:
|
contents:
|
||||||
- src: client/ui/build/netbird.desktop
|
- src: client/ui/build/linux/netbird.desktop
|
||||||
dst: /usr/share/applications/netbird.desktop
|
dst: /usr/share/applications/org.wails.netbird.desktop
|
||||||
- src: client/ui/assets/netbird.png
|
- src: client/ui/build/appicon.png
|
||||||
dst: /usr/share/pixmaps/netbird.png
|
dst: /usr/share/pixmaps/netbird.png
|
||||||
dependencies:
|
dependencies:
|
||||||
- netbird
|
- netbird
|
||||||
|
|
||||||
rpm:
|
rpm:
|
||||||
signature:
|
signature:
|
||||||
key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}'
|
key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}'
|
||||||
|
|||||||
@@ -1,6 +1,15 @@
|
|||||||
version: 2
|
version: 2
|
||||||
|
|
||||||
project_name: netbird-ui
|
project_name: netbird-ui
|
||||||
|
|
||||||
|
before:
|
||||||
|
hooks:
|
||||||
|
# Bindings are gitignored; regenerate before the frontend build so
|
||||||
|
# the @wailsio/runtime Vite plugin can resolve them (vite refuses to
|
||||||
|
# build without them).
|
||||||
|
- sh -c 'cd client/ui && wails3 generate bindings -clean=true -ts'
|
||||||
|
- sh -c 'cd client/ui/frontend && pnpm install --frozen-lockfile && pnpm build'
|
||||||
|
|
||||||
builds:
|
builds:
|
||||||
- id: netbird-ui-darwin
|
- id: netbird-ui-darwin
|
||||||
dir: client/ui
|
dir: client/ui
|
||||||
@@ -21,7 +30,7 @@ builds:
|
|||||||
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
- -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser
|
||||||
mod_timestamp: "{{ .CommitTimestamp }}"
|
mod_timestamp: "{{ .CommitTimestamp }}"
|
||||||
tags:
|
tags:
|
||||||
- load_wgnt_from_rsrc
|
- production
|
||||||
|
|
||||||
universal_binaries:
|
universal_binaries:
|
||||||
- id: netbird-ui-darwin
|
- id: netbird-ui-darwin
|
||||||
|
|||||||
@@ -79,13 +79,21 @@ dependencies are installed. Here is a short guide on how that can be done.
|
|||||||
|
|
||||||
### Requirements
|
### Requirements
|
||||||
|
|
||||||
#### Go 1.21
|
#### Go 1.25
|
||||||
|
|
||||||
Follow the installation guide from https://go.dev/
|
Follow the installation guide from https://go.dev/
|
||||||
|
|
||||||
#### UI client - Fyne toolkit
|
#### UI client - Wails v3 + React
|
||||||
|
|
||||||
We use the fyne toolkit in our UI client. You can follow its requirement guide to have all its dependencies installed: https://developer.fyne.io/started/#prerequisites
|
The desktop UI client (`client/ui`) is built with [Wails v3](https://v3.wails.io/) and a React frontend rendered in a WebView. To build it you need:
|
||||||
|
|
||||||
|
- Go ≥ 1.25
|
||||||
|
- Node ≥ 20 and **pnpm** (`corepack enable && corepack prepare pnpm@latest --activate`)
|
||||||
|
- The `wails3` CLI: `go install github.com/wailsapp/wails/v3/cmd/wails3@latest`
|
||||||
|
- The `task` runner: `go install github.com/go-task/task/v3/cmd/task@latest`
|
||||||
|
- Linux only: `libwebkitgtk-6.0-dev`, `libgtk-4-dev`, `libsoup-3.0-dev`
|
||||||
|
|
||||||
|
All UI build, dev-loop, and cross-compile commands are described in the [UI client](#ui-client) section below.
|
||||||
|
|
||||||
#### gRPC
|
#### gRPC
|
||||||
You can follow the instructions from the quickstarter guide https://grpc.io/docs/languages/go/quickstart/#prerequisites and then run the `generate.sh` files located in each `proto` directory to generate changes.
|
You can follow the instructions from the quickstarter guide https://grpc.io/docs/languages/go/quickstart/#prerequisites and then run the `generate.sh` files located in each `proto` directory to generate changes.
|
||||||
@@ -214,6 +222,49 @@ To start NetBird the client in the foreground:
|
|||||||
sudo ./client up --log-level debug --log-file console
|
sudo ./client up --log-level debug --log-file console
|
||||||
```
|
```
|
||||||
> On Windows use a powershell with administrator privileges
|
> On Windows use a powershell with administrator privileges
|
||||||
|
|
||||||
|
#### UI client
|
||||||
|
|
||||||
|
The desktop UI lives in `client/ui` and is built with Wails v3 (see [Requirements](#ui-client---wails-v3--react)). All commands run from `client/ui`.
|
||||||
|
|
||||||
|
Live-reload development (Vite + Go binary + `*.go` watcher):
|
||||||
|
|
||||||
|
```
|
||||||
|
cd client/ui
|
||||||
|
task dev
|
||||||
|
```
|
||||||
|
|
||||||
|
Pass daemon flags after `--`, pointing the UI at the socket the daemon serves:
|
||||||
|
|
||||||
|
```
|
||||||
|
task dev -- --daemon-addr=unix:///var/run/netbird.sock # Linux, macOS
|
||||||
|
task dev -- --daemon-addr=npipe://netbird # Windows
|
||||||
|
```
|
||||||
|
|
||||||
|
On Windows the daemon serves a named pipe (`npipe://netbird`). Which path that
|
||||||
|
ends up being depends on what the daemon may create: as a service or elevated it
|
||||||
|
serves `\\.\pipe\ProtectedPrefix\Administrators\netbird`, which no unprivileged
|
||||||
|
process can take from it, and otherwise it falls back to `\\.\pipe\netbird`.
|
||||||
|
Clients try both and check who owns the pipe before using the plain one. Avoid
|
||||||
|
`tcp://127.0.0.1:41731`: loopback TCP carries no caller identity, so the daemon
|
||||||
|
refuses the operations that require an administrator and you will not exercise
|
||||||
|
those paths.
|
||||||
|
|
||||||
|
Production build (frontend assets embedded into the binary, output in `client/ui/bin/`):
|
||||||
|
|
||||||
|
```
|
||||||
|
cd client/ui
|
||||||
|
task build
|
||||||
|
```
|
||||||
|
|
||||||
|
Cross-compile the Windows binary from Linux (requires the mingw-w64 toolchain, e.g. `sudo apt install gcc-mingw-w64-x86-64`):
|
||||||
|
|
||||||
|
```
|
||||||
|
CGO_ENABLED=1 task windows:build
|
||||||
|
```
|
||||||
|
|
||||||
|
> macOS cross-compile from Linux is not supported (signing and notarization need a real Mac).
|
||||||
|
|
||||||
#### Signal service
|
#### Signal service
|
||||||
|
|
||||||
To start NetBird's signal, execute:
|
To start NetBird's signal, execute:
|
||||||
@@ -251,10 +302,10 @@ Create dist directory
|
|||||||
mkdir -p dist/netbird_windows_amd64
|
mkdir -p dist/netbird_windows_amd64
|
||||||
```
|
```
|
||||||
|
|
||||||
UI client
|
UI client (built with Wails v3 — see the [UI client](#ui-client) section above)
|
||||||
```shell
|
```shell
|
||||||
CC=x86_64-w64-mingw32-gcc CGO_ENABLED=1 GOOS=windows GOARCH=amd64 go build -o netbird-ui.exe -ldflags "-s -w -H windowsgui" ./client/ui
|
(cd client/ui && CGO_ENABLED=1 task windows:build)
|
||||||
mv netbird-ui.exe ./dist/netbird_windows_amd64/
|
mv client/ui/bin/netbird-ui.exe ./dist/netbird_windows_amd64/
|
||||||
```
|
```
|
||||||
|
|
||||||
Client
|
Client
|
||||||
@@ -291,8 +342,6 @@ go test -exec sudo ./...
|
|||||||
```
|
```
|
||||||
> On Windows use a powershell with administrator privileges
|
> On Windows use a powershell with administrator privileges
|
||||||
|
|
||||||
> Non-GTK environments will need the `libayatana-appindicator3-dev` (debian/ubuntu) package installed
|
|
||||||
|
|
||||||
## Checklist before submitting a PR
|
## Checklist before submitting a PR
|
||||||
As a critical network service and open-source project, we must enforce a few things before submitting the pull-requests:
|
As a critical network service and open-source project, we must enforce a few things before submitting the pull-requests:
|
||||||
- Keep functions as simple as possible, with a single purpose
|
- Keep functions as simple as possible, with a single purpose
|
||||||
|
|||||||
@@ -1,16 +1,47 @@
|
|||||||
# NetBird Agent Network
|
# NetBird Agent Network
|
||||||
|
|
||||||
Agent Network is NetBird's access control layer for AI agents and the people who run
|
Agent Network is NetBird's access control layer for AI agents and the people who run them.
|
||||||
them. It gives every agent a real identity, tied to your identity provider (IdP), and
|
It gives every agent a real identity, tied to an identity provider (IdP), and governs what it can reach: LLM APIs and
|
||||||
governs what it can reach — the LLM APIs and AI gateways it can call, and the internal
|
AI gateways it can call, and the internal resources it can access. Traffic flows only over the encrypted NetBird tunnel,
|
||||||
resources it can access. Traffic flows only over the encrypted NetBird tunnel, scoped by
|
scoped by policy, with no API keys or other credentials to leak. It also gives you control over cost and token usage.
|
||||||
policy, with no API keys to leak.
|
|
||||||
|
|
||||||
> **Beta.** Agent Network is open source and can be self-hosted on your own
|
Because every LLM request passes through an
|
||||||
> infrastructure.
|
identity-aware proxy, you can:
|
||||||
|
|
||||||
|
- **Set spending and rate limits** per agent, per user, or per team — with hard caps
|
||||||
|
that stop requests once a budget is reached.
|
||||||
|
- **Restrict models and providers** so agents can only call approved (and cost-appropriate)
|
||||||
|
endpoints, keeping expensive models off-limits unless explicitly allowed.
|
||||||
|
- **Attribute usage** by tracking token consumption and cost per identity, group, or cost center so every
|
||||||
|
request is tied back to the agent and person responsible.
|
||||||
|
- **Reuse your existing AI gateway** — point the proxy at a gateway you already run,
|
||||||
|
keeping its routing and config in place while it adds identity on top, so you skip
|
||||||
|
API key distribution.
|
||||||
|
|
||||||
|
https://github.com/user-attachments/assets/44d18286-d8ab-49f8-a457-98ccd66f3268
|
||||||
|
|
||||||
|
> **Beta.** Agent Network is in beta, but it's stable and already running in
|
||||||
|
> production environments. It's fully open source and can be self-hosted on your own
|
||||||
|
> infrastructure, with no vendor lock-in and no data leaving your environment.
|
||||||
|
|
||||||
## How it works
|
## How it works
|
||||||
|
|
||||||
|
Say you have a simple use case: your Engineering or IT team needs access to Claude Code or Codex, and you want visibility into usage plus the ability to enforce budgets.
|
||||||
|
How can you do that without creating a dedicated API key for every team?
|
||||||
|
|
||||||
|
With Agent Network you get a private endpoint inside your network, for example: https://mirror.netbird.ai
|
||||||
|
Teams configure their agents to point to that endpoint instead of using individual API keys directly.
|
||||||
|
|
||||||
|
This endpoint is only reachable when users are connected to your NetBird network and authenticated through your IdP. Otherwise, it is not accessible from the public internet.
|
||||||
|
You can then use this private endpoint to configure your AI agents, whether that is Claude Code, Codex, or another tool.
|
||||||
|
|
||||||
|
## Quickstart
|
||||||
|
|
||||||
|
Full step-by-step setup:
|
||||||
|
**https://docs.netbird.io/agent-network/quickstart**
|
||||||
|
|
||||||
|
## Architecture
|
||||||
|
|
||||||
Agent Network is built on two existing NetBird capabilities:
|
Agent Network is built on two existing NetBird capabilities:
|
||||||
|
|
||||||
- **Overlay network** — the encrypted WireGuard mesh between peers.
|
- **Overlay network** — the encrypted WireGuard mesh between peers.
|
||||||
@@ -22,6 +53,9 @@ LLM traffic is routed through the proxy's identity-aware pipeline, while interna
|
|||||||
resources (databases, internal APIs, self-hosted models) are reached directly over
|
resources (databases, internal APIs, self-hosted models) are reached directly over
|
||||||
peer-to-peer WireGuard tunnels, governed by the same identities and access policies.
|
peer-to-peer WireGuard tunnels, governed by the same identities and access policies.
|
||||||
|
|
||||||
|
<img width="4720" height="2218" alt="image" src="https://github.com/user-attachments/assets/1afa5da1-4b82-4f8a-a7a8-f417efadf1eb" />
|
||||||
|
|
||||||
|
|
||||||
## Where the code lives
|
## Where the code lives
|
||||||
|
|
||||||
There is no separate "agent-network" service — it reuses the reverse-proxy and management
|
There is no separate "agent-network" service — it reuses the reverse-proxy and management
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"slices"
|
"slices"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -56,6 +57,12 @@ type DnsReadyListener interface {
|
|||||||
dns.ReadyListener
|
dns.ReadyListener
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TunSettings is a snapshot of the settings the TUN device is rebuilt with
|
||||||
|
type TunSettings struct {
|
||||||
|
Routes string
|
||||||
|
SearchDomains string
|
||||||
|
}
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
formatter.SetLogcatFormatter(log.StandardLogger())
|
formatter.SetLogcatFormatter(log.StandardLogger())
|
||||||
}
|
}
|
||||||
@@ -75,13 +82,34 @@ type Client struct {
|
|||||||
connectClient *internal.ConnectClient
|
connectClient *internal.ConnectClient
|
||||||
config *profilemanager.Config
|
config *profilemanager.Config
|
||||||
cacheDir string
|
cacheDir string
|
||||||
|
// Identifies the running profile for the SSO login hint; see profile_state.go.
|
||||||
|
cfgPath string
|
||||||
|
|
||||||
|
stateChangeMu sync.Mutex
|
||||||
|
stateChangeSubID string
|
||||||
|
eventSub *peer.EventSubscription
|
||||||
|
// Closed to stop the watch goroutines from delivering buffered items to a
|
||||||
|
// listener that has been removed or replaced. See stopStateChangeWatchLocked.
|
||||||
|
stateChangeDone chan struct{}
|
||||||
|
|
||||||
|
// Latched "the server wants an interactive login": survives the engine
|
||||||
|
// restarts that replace the run loop's context state. See Client.Status.
|
||||||
|
// Guarded by loginRequiredMu together with loginCleared, which counts
|
||||||
|
// clears so a stale observation cannot re-latch over one.
|
||||||
|
loginRequiredMu sync.Mutex
|
||||||
|
loginRequired bool
|
||||||
|
loginCleared uint64
|
||||||
|
|
||||||
|
extendMu sync.Mutex
|
||||||
|
extendCancel context.CancelFunc
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cc *internal.ConnectClient) {
|
func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cfgPath string, cc *internal.ConnectClient) {
|
||||||
c.stateMu.Lock()
|
c.stateMu.Lock()
|
||||||
defer c.stateMu.Unlock()
|
defer c.stateMu.Unlock()
|
||||||
c.config = cfg
|
c.config = cfg
|
||||||
c.cacheDir = cacheDir
|
c.cacheDir = cacheDir
|
||||||
|
c.cfgPath = cfgPath
|
||||||
c.connectClient = cc
|
c.connectClient = cc
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -91,6 +119,16 @@ func (c *Client) stateSnapshot() (*profilemanager.Config, string, *internal.Conn
|
|||||||
return c.config, c.cacheDir, c.connectClient
|
return c.config, c.cacheDir, c.connectClient
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// authSnapshot returns the config together with the path it was loaded from, in
|
||||||
|
// one lock: the path identifies the profile whose account email backs the login
|
||||||
|
// hint, so reading it separately could pair one profile's config with another's
|
||||||
|
// hint when a profile switch lands in between.
|
||||||
|
func (c *Client) authSnapshot() (*profilemanager.Config, string, *internal.ConnectClient) {
|
||||||
|
c.stateMu.RLock()
|
||||||
|
defer c.stateMu.RUnlock()
|
||||||
|
return c.config, c.cfgPath, c.connectClient
|
||||||
|
}
|
||||||
|
|
||||||
func (c *Client) getConnectClient() *internal.ConnectClient {
|
func (c *Client) getConnectClient() *internal.ConnectClient {
|
||||||
c.stateMu.RLock()
|
c.stateMu.RLock()
|
||||||
defer c.stateMu.RUnlock()
|
defer c.stateMu.RUnlock()
|
||||||
@@ -143,16 +181,21 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
|
|||||||
defer c.ctxCancel()
|
defer c.ctxCancel()
|
||||||
c.ctxCancelLock.Unlock()
|
c.ctxCancelLock.Unlock()
|
||||||
|
|
||||||
auth := NewAuthWithConfig(ctx, cfg)
|
auth := NewAuthWithConfig(ctx, cfg, cfgFile)
|
||||||
err = auth.login(urlOpener, isAndroidTV)
|
err = auth.login(urlOpener, isAndroidTV)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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)
|
||||||
c.setState(cfg, cacheDir, connectClient)
|
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
||||||
|
// This path runs the interactive SSO flow, so reaching here means the peer
|
||||||
|
// is authenticated again — release the latch Status() reports from. Clear
|
||||||
|
// only once the fresh connect client is installed: until then Status()
|
||||||
|
// still reads the previous run's context state, which holds the NeedsLogin
|
||||||
|
// that prompted this login, and would re-latch what was just cleared.
|
||||||
|
c.clearLoginRequired()
|
||||||
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)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -187,7 +230,7 @@ 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)
|
||||||
c.setState(cfg, cacheDir, 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)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -216,6 +259,24 @@ func (c *Client) RenewTun(fd int) error {
|
|||||||
return e.RenewTun(fd)
|
return e.RenewTun(fd)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *Client) GetTunSettings() (*TunSettings, error) {
|
||||||
|
cc := c.getConnectClient()
|
||||||
|
if cc == nil {
|
||||||
|
return nil, fmt.Errorf("engine not running")
|
||||||
|
}
|
||||||
|
|
||||||
|
e := cc.Engine()
|
||||||
|
if e == nil {
|
||||||
|
return nil, fmt.Errorf("engine not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
routes, searchDomains := e.TunSettings()
|
||||||
|
return &TunSettings{
|
||||||
|
Routes: strings.Join(routes, ";"),
|
||||||
|
SearchDomains: strings.Join(searchDomains, ";"),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
// 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.
|
// It works both with and without a running engine.
|
||||||
func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (string, error) {
|
func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (string, error) {
|
||||||
@@ -247,6 +308,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (strin
|
|||||||
deps.SyncResponse = resp
|
deps.SyncResponse = resp
|
||||||
|
|
||||||
if e := cc.Engine(); e != nil {
|
if e := cc.Engine(); e != nil {
|
||||||
|
deps.RefreshStatus = func() {
|
||||||
|
e.RunHealthProbes(context.Background(), true)
|
||||||
|
}
|
||||||
if cm := e.GetClientMetrics(); cm != nil {
|
if cm := e.GetClientMetrics(); cm != nil {
|
||||||
deps.ClientMetrics = cm
|
deps.ClientMetrics = cm
|
||||||
}
|
}
|
||||||
@@ -296,6 +360,13 @@ func (c *Client) SetInfoLogLevel() {
|
|||||||
// PeersList return with the list of the PeerInfos
|
// PeersList return with the list of the PeerInfos
|
||||||
func (c *Client) PeersList() *PeerInfoArray {
|
func (c *Client) PeersList() *PeerInfoArray {
|
||||||
|
|
||||||
|
// The recorder only caches transfer counters and handshake times; nothing
|
||||||
|
// refreshes them on its own, so without this they read as zero. The desktop
|
||||||
|
// daemon does the same before serving a full peer status.
|
||||||
|
if err := c.recorder.RefreshWireGuardStats(); err != nil {
|
||||||
|
log.Debugf("failed to refresh WireGuard stats: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
fullStatus := c.recorder.GetFullStatus()
|
fullStatus := c.recorder.GetFullStatus()
|
||||||
|
|
||||||
peerInfos := make([]PeerInfo, len(fullStatus.Peers))
|
peerInfos := make([]PeerInfo, len(fullStatus.Peers))
|
||||||
@@ -306,6 +377,20 @@ func (c *Client) PeersList() *PeerInfoArray {
|
|||||||
FQDN: p.FQDN,
|
FQDN: p.FQDN,
|
||||||
ConnStatus: int(p.ConnStatus),
|
ConnStatus: int(p.ConnStatus),
|
||||||
Routes: PeerRoutes{routes: maps.Keys(p.GetRoutes())},
|
Routes: PeerRoutes{routes: maps.Keys(p.GetRoutes())},
|
||||||
|
|
||||||
|
PubKey: p.PubKey,
|
||||||
|
Latency: formatDuration(p.Latency),
|
||||||
|
LatencyMs: p.Latency.Milliseconds(),
|
||||||
|
BytesRx: p.BytesRx,
|
||||||
|
BytesTx: p.BytesTx,
|
||||||
|
ConnStatusUpdate: formatTime(p.ConnStatusUpdate),
|
||||||
|
Relayed: p.Relayed,
|
||||||
|
RosenpassEnabled: p.RosenpassEnabled,
|
||||||
|
LastWireguardHandshake: formatTime(p.LastWireguardHandshake),
|
||||||
|
LocalIceCandidateType: p.LocalIceCandidateType,
|
||||||
|
RemoteIceCandidateType: p.RemoteIceCandidateType,
|
||||||
|
LocalIceCandidateEndpoint: p.LocalIceCandidateEndpoint,
|
||||||
|
RemoteIceCandidateEndpoint: p.RemoteIceCandidateEndpoint,
|
||||||
}
|
}
|
||||||
peerInfos[n] = pi
|
peerInfos[n] = pi
|
||||||
}
|
}
|
||||||
@@ -436,10 +521,6 @@ func (c *Client) RemoveConnectionListener() {
|
|||||||
c.recorder.RemoveConnectionListener()
|
c.recorder.RemoveConnectionListener()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) toggleRoute(command routeCommand) error {
|
|
||||||
return command.toggleRoute()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Client) getRouteManager() (routemanager.Manager, error) {
|
func (c *Client) getRouteManager() (routemanager.Manager, error) {
|
||||||
client := c.getConnectClient()
|
client := c.getConnectClient()
|
||||||
if client == nil {
|
if client == nil {
|
||||||
@@ -459,22 +540,22 @@ func (c *Client) getRouteManager() (routemanager.Manager, error) {
|
|||||||
return manager, nil
|
return manager, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) SelectRoute(route string) error {
|
func (c *Client) SelectRoute(id string) error {
|
||||||
manager, err := c.getRouteManager()
|
manager, err := c.getRouteManager()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return c.toggleRoute(selectRouteCommand{route: route, manager: manager})
|
return manager.SelectRoutes([]route.NetID{route.NetID(id)}, true)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) DeselectRoute(route string) error {
|
func (c *Client) DeselectRoute(id string) error {
|
||||||
manager, err := c.getRouteManager()
|
manager, err := c.getRouteManager()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return c.toggleRoute(deselectRouteCommand{route: route, manager: manager})
|
return manager.DeselectRoutes([]route.NetID{route.NetID(id)})
|
||||||
}
|
}
|
||||||
|
|
||||||
// getNetworkDomainsFromRoute extracts domains from a route and enriches each domain
|
// getNetworkDomainsFromRoute extracts domains from a route and enriches each domain
|
||||||
@@ -509,3 +590,28 @@ func exportEnvList(list *EnvList) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// formatDuration renders a duration for display, trimming the fractional part
|
||||||
|
// to two digits so latencies read as "12.34ms" rather than "12.345678ms".
|
||||||
|
func formatDuration(d time.Duration) string {
|
||||||
|
ds := d.String()
|
||||||
|
dotIndex := strings.Index(ds, ".")
|
||||||
|
if dotIndex == -1 {
|
||||||
|
return ds
|
||||||
|
}
|
||||||
|
|
||||||
|
endIndex := min(dotIndex+3, len(ds))
|
||||||
|
|
||||||
|
// Skip the remaining digits so only the unit suffix is appended back.
|
||||||
|
unitStart := endIndex
|
||||||
|
for unitStart < len(ds) && ds[unitStart] >= '0' && ds[unitStart] <= '9' {
|
||||||
|
unitStart++
|
||||||
|
}
|
||||||
|
return ds[:endIndex] + ds[unitStart:]
|
||||||
|
}
|
||||||
|
|
||||||
|
// formatTime renders a timestamp in UTC using a fixed layout. The zero time is
|
||||||
|
// passed through as-is so the UI can recognise it and show "never" instead.
|
||||||
|
func formatTime(t time.Time) string {
|
||||||
|
return t.UTC().Format("2006-01-02 15:04:05")
|
||||||
|
}
|
||||||
|
|||||||
@@ -4,6 +4,8 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/internal/auth"
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
"github.com/netbirdio/netbird/client/system"
|
"github.com/netbirdio/netbird/client/system"
|
||||||
@@ -53,11 +55,14 @@ func NewAuth(cfgPath string, mgmURL string) (*Auth, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewAuthWithConfig instantiate Auth based on existing config
|
// NewAuthWithConfig instantiate Auth based on existing config. cfgPath is the
|
||||||
func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config) *Auth {
|
// file the config was loaded from; it identifies the profile whose account email
|
||||||
|
// backs the login_hint.
|
||||||
|
func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config, cfgPath string) *Auth {
|
||||||
return &Auth{
|
return &Auth{
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
config: config,
|
config: config,
|
||||||
|
cfgPath: cfgPath,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -150,12 +155,14 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
jwtToken := ""
|
jwtToken := ""
|
||||||
|
email := ""
|
||||||
if needsLogin {
|
if needsLogin {
|
||||||
tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV)
|
tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("interactive sso login failed: %v", err)
|
return fmt.Errorf("interactive sso login failed: %v", err)
|
||||||
}
|
}
|
||||||
jwtToken = tokenInfo.GetTokenToUse()
|
jwtToken = tokenInfo.GetTokenToUse()
|
||||||
|
email = tokenInfo.Email
|
||||||
}
|
}
|
||||||
|
|
||||||
err, _ = authClient.Login(a.ctx, "", jwtToken)
|
err, _ = authClient.Login(a.ctx, "", jwtToken)
|
||||||
@@ -163,17 +170,42 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
|
|||||||
return fmt.Errorf("login failed: %v", err)
|
return fmt.Errorf("login failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Stored after Login, not before: a rejected token must not leave a hint
|
||||||
|
// pointing at an account that cannot be used.
|
||||||
|
if email != "" && a.cfgPath != "" {
|
||||||
|
if err := writeProfileEmail(a.cfgPath, email); err != nil {
|
||||||
|
log.Warnf("failed to store profile account email: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
go urlOpener.OnLoginSuccess()
|
go urlOpener.OnLoginSuccess()
|
||||||
|
|
||||||
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)
|
||||||
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
|
||||||
|
// leaves the choice to the IdP, which is how accounts get switched.
|
||||||
|
if a.cfgPath != "" {
|
||||||
|
if hint := readProfileEmail(a.cfgPath); hint != "" {
|
||||||
|
if setter, ok := oAuthFlow.(loginHintSetter); ok {
|
||||||
|
setter.SetLoginHint(hint)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO())
|
flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err)
|
return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err)
|
||||||
|
|||||||
@@ -12,12 +12,30 @@ const (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// PeerInfo describe information about the peers. It designed for the UI usage
|
// PeerInfo describe information about the peers. It designed for the UI usage
|
||||||
|
//
|
||||||
|
// The fields below ConnStatus back the peer detail screen. Durations and times
|
||||||
|
// are pre-formatted into strings so the UI does not have to know Go's layouts;
|
||||||
|
// Latency is additionally exposed as LatencyMs for colour coding.
|
||||||
type PeerInfo struct {
|
type PeerInfo struct {
|
||||||
IP string
|
IP string
|
||||||
IPv6 string
|
IPv6 string
|
||||||
FQDN string
|
FQDN string
|
||||||
ConnStatus int
|
ConnStatus int
|
||||||
Routes PeerRoutes
|
Routes PeerRoutes
|
||||||
|
|
||||||
|
PubKey string
|
||||||
|
Latency string
|
||||||
|
LatencyMs int64
|
||||||
|
BytesRx int64
|
||||||
|
BytesTx int64
|
||||||
|
ConnStatusUpdate string
|
||||||
|
Relayed bool
|
||||||
|
RosenpassEnabled bool
|
||||||
|
LastWireguardHandshake string
|
||||||
|
LocalIceCandidateType string
|
||||||
|
RemoteIceCandidateType string
|
||||||
|
LocalIceCandidateEndpoint string
|
||||||
|
RemoteIceCandidateEndpoint string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PeerInfo) GetPeerRoutes() *PeerRoutes {
|
func (p *PeerInfo) GetPeerRoutes() *PeerRoutes {
|
||||||
|
|||||||
@@ -13,18 +13,17 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// Android-specific config filename (different from desktop default.json)
|
|
||||||
defaultConfigFilename = "netbird.cfg"
|
|
||||||
// Subdirectory for non-default profiles (must match Java Preferences.java)
|
|
||||||
profilesSubdir = "profiles"
|
|
||||||
// Android uses a single user context per app (non-empty username required by ServiceManager)
|
// Android uses a single user context per app (non-empty username required by ServiceManager)
|
||||||
androidUsername = "android"
|
androidUsername = "android"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Profile represents a profile for gomobile
|
// Profile represents a profile for gomobile
|
||||||
type Profile struct {
|
type Profile struct {
|
||||||
ID string
|
ID string
|
||||||
Name string
|
Name string
|
||||||
|
// Email is the account this profile last logged in with, "" if it never
|
||||||
|
// completed an SSO login or was logged out. See profile_state.go.
|
||||||
|
Email string
|
||||||
IsActive bool
|
IsActive bool
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -101,6 +100,7 @@ func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) {
|
|||||||
profiles = append(profiles, &Profile{
|
profiles = append(profiles, &Profile{
|
||||||
ID: p.ID.String(),
|
ID: p.ID.String(),
|
||||||
Name: p.Name,
|
Name: p.Name,
|
||||||
|
Email: pm.profileEmail(p.ID.String()),
|
||||||
IsActive: p.IsActive,
|
IsActive: p.IsActive,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -123,7 +123,22 @@ func (pm *ProfileManager) GetActiveProfile() (*Profile, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to resolve active profile %q: %w", activeState.ID, err)
|
return nil, fmt.Errorf("failed to resolve active profile %q: %w", activeState.ID, err)
|
||||||
}
|
}
|
||||||
return &Profile{ID: prof.ID.String(), Name: prof.Name, IsActive: true}, nil
|
return &Profile{
|
||||||
|
ID: prof.ID.String(),
|
||||||
|
Name: prof.Name,
|
||||||
|
Email: pm.profileEmail(prof.ID.String()),
|
||||||
|
IsActive: true,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// profileEmail returns the account email recorded for a profile. Display-only, so
|
||||||
|
// an unresolvable path degrades to "" rather than an error.
|
||||||
|
func (pm *ProfileManager) profileEmail(id string) string {
|
||||||
|
configPath, err := pm.getProfileConfigPath(id)
|
||||||
|
if err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return readProfileEmail(configPath)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SwitchProfile switches to a different profile
|
// SwitchProfile switches to a different profile
|
||||||
@@ -185,10 +200,28 @@ 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.
|
||||||
|
if err := removeProfileEmail(configPath); err != nil {
|
||||||
|
log.Warnf("failed to clear stored account email for profile %s: %v", id, err)
|
||||||
|
}
|
||||||
|
|
||||||
log.Infof("logged out from profile: %s", id)
|
log.Infof("logged out from profile: %s", id)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RenameProfile changes a profile's display name. The profile ID, and therefore
|
||||||
|
// its on-disk filename, is left untouched: only the "name" field of the config
|
||||||
|
// is rewritten. This works for the default profile too, whose config lives in
|
||||||
|
// netbird.cfg rather than under profiles/.
|
||||||
|
func (pm *ProfileManager) RenameProfile(id string, newName string) error {
|
||||||
|
if err := pm.serviceMgr.RenameProfile(profilemanager.ID(id), androidUsername, newName); err != nil {
|
||||||
|
return fmt.Errorf("failed to rename profile: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infof("renamed profile %s to: %s", id, newName)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// RemoveProfile deletes a profile
|
// RemoveProfile deletes a profile
|
||||||
func (pm *ProfileManager) RemoveProfile(id string) error {
|
func (pm *ProfileManager) RemoveProfile(id string) error {
|
||||||
// Use ServiceManager (removes profile from profiles/ directory)
|
// Use ServiceManager (removes profile from profiles/ directory)
|
||||||
|
|||||||
108
client/android/profile_state.go
Normal file
108
client/android/profile_state.go
Normal file
@@ -0,0 +1,108 @@
|
|||||||
|
package android
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
"github.com/netbirdio/netbird/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// Android-specific config filename (different from desktop default.json)
|
||||||
|
defaultConfigFilename = "netbird.cfg"
|
||||||
|
// Subdirectory for non-default profiles (must match Java Preferences.java)
|
||||||
|
profilesSubdir = "profiles"
|
||||||
|
// profileAccountSuffix names the file holding the profile's account email.
|
||||||
|
// Deliberately not ".state.json", which desktop uses for the same data:
|
||||||
|
// there the email and the engine's state manager live in different
|
||||||
|
// directories, but on Android both resolve under files/, so sharing the name
|
||||||
|
// would have the two overwrite each other — the state manager rewrites the
|
||||||
|
// whole file from its own keys (see statemanager.Manager.PersistState), and
|
||||||
|
// this package's writer does the same in reverse.
|
||||||
|
profileAccountSuffix = ".account.json"
|
||||||
|
)
|
||||||
|
|
||||||
|
// profileAccountPathFor derives the account file path from a profile's config
|
||||||
|
// path: netbird.cfg -> netbird.account.json, <id>.json -> <id>.account.json.
|
||||||
|
//
|
||||||
|
// Deriving from the config path rather than resolving the active profile keeps
|
||||||
|
// the write on the profile the login actually ran for: Auth.login runs in a
|
||||||
|
// goroutine, so the active profile can change under a flow already in flight.
|
||||||
|
func profileAccountPathFor(configPath string) (string, error) {
|
||||||
|
if configPath == "" {
|
||||||
|
return "", fmt.Errorf("empty config path")
|
||||||
|
}
|
||||||
|
|
||||||
|
base := filepath.Base(configPath)
|
||||||
|
stem := strings.TrimSuffix(base, filepath.Ext(base))
|
||||||
|
if stem == "" || stem == "." {
|
||||||
|
return "", fmt.Errorf("config path %q has no filename stem", configPath)
|
||||||
|
}
|
||||||
|
|
||||||
|
return filepath.Join(filepath.Dir(configPath), stem+profileAccountSuffix), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// readProfileEmail returns the account email stored for the profile whose config
|
||||||
|
// lives at configPath. A missing or unreadable file yields "", which leaves the
|
||||||
|
// account choice to the IdP.
|
||||||
|
func readProfileEmail(configPath string) string {
|
||||||
|
accountPath, err := profileAccountPathFor(configPath)
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("no profile account path for login hint: %v", err)
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var state profilemanager.ProfileState
|
||||||
|
if _, err := util.ReadJson(accountPath, &state); err != nil {
|
||||||
|
if !os.IsNotExist(err) {
|
||||||
|
log.Debugf("failed to read profile account for login hint: %v", err)
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
return state.Email
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeProfileEmail records the account email for the profile whose config lives
|
||||||
|
// at configPath, so later logins can pass it as an OIDC login_hint. An empty
|
||||||
|
// email is ignored rather than blanking what is already stored.
|
||||||
|
func writeProfileEmail(configPath string, email string) error {
|
||||||
|
if email == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
accountPath, err := profileAccountPathFor(configPath)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("resolve profile account path: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
state := profilemanager.ProfileState{Email: email}
|
||||||
|
if err := util.WriteJsonWithRestrictedPermission(context.Background(), accountPath, state); err != nil {
|
||||||
|
return fmt.Errorf("write profile account: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// removeProfileEmail drops the stored account email. Called on logout: while the
|
||||||
|
// email is on disk it goes out as a login_hint, which would steer the next login
|
||||||
|
// straight back into the account just logged out of. Mirrors the desktop UI's
|
||||||
|
// RemoveProfileState call.
|
||||||
|
func removeProfileEmail(configPath string) error {
|
||||||
|
accountPath, err := profileAccountPathFor(configPath)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("resolve profile account path: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.Remove(accountPath); err != nil && !os.IsNotExist(err) {
|
||||||
|
return fmt.Errorf("remove profile account: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
161
client/android/profile_state_test.go
Normal file
161
client/android/profile_state_test.go
Normal file
@@ -0,0 +1,161 @@
|
|||||||
|
package android
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestProfileAccountPathFor(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
configPath string
|
||||||
|
want string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "default profile",
|
||||||
|
configPath: "/data/data/io.netbird.client/files/netbird.cfg",
|
||||||
|
want: "/data/data/io.netbird.client/files/netbird.account.json",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "id profile",
|
||||||
|
configPath: "/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.json",
|
||||||
|
want: "/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.account.json",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "legacy name-keyed profile is handled the same way",
|
||||||
|
configPath: "/data/data/io.netbird.client/files/profiles/work.json",
|
||||||
|
want: "/data/data/io.netbird.client/files/profiles/work.account.json",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty path is rejected",
|
||||||
|
configPath: "",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, err := profileAccountPathFor(tt.configPath)
|
||||||
|
if tt.wantErr {
|
||||||
|
if err == nil {
|
||||||
|
t.Fatalf("expected an error, got path %q", got)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if got != tt.want {
|
||||||
|
t.Errorf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProfileAccountPathForDefaultDoesNotCollide(t *testing.T) {
|
||||||
|
root := "/data/data/io.netbird.client/files"
|
||||||
|
|
||||||
|
defaultAccount, err := profileAccountPathFor(filepath.Join(root, defaultConfigFilename))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("default profile: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
idAccount, err := profileAccountPathFor(filepath.Join(root, profilesSubdir, "abc123.json"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("id profile: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if defaultAccount == idAccount {
|
||||||
|
t.Fatalf("default and id profile share an account file: %q", defaultAccount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The account file must never land on the engine state file: on Android both
|
||||||
|
// resolve under files/, and the state manager rewrites the whole file from its
|
||||||
|
// own keys, so sharing a path would have the two overwrite each other. The
|
||||||
|
// expected names here mirror ProfileManager.GetStateFilePath.
|
||||||
|
func TestProfileAccountPathAvoidsEngineStateFile(t *testing.T) {
|
||||||
|
root := "/data/data/io.netbird.client/files"
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
configPath string
|
||||||
|
engineState string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
configPath: filepath.Join(root, defaultConfigFilename),
|
||||||
|
engineState: filepath.Join(root, "state.json"),
|
||||||
|
},
|
||||||
|
{
|
||||||
|
configPath: filepath.Join(root, profilesSubdir, "abc123.json"),
|
||||||
|
engineState: filepath.Join(root, profilesSubdir, "abc123.state.json"),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, c := range cases {
|
||||||
|
account, err := profileAccountPathFor(c.configPath)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("%s: %v", c.configPath, err)
|
||||||
|
}
|
||||||
|
if account == c.engineState {
|
||||||
|
t.Errorf("account file collides with the engine state file: %q", account)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteThenReadProfileEmail(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "profiles", "abc123.json")
|
||||||
|
if err := ensureDirFor(t, configPath); err != nil {
|
||||||
|
t.Fatalf("prepare dir: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := readProfileEmail(configPath); got != "" {
|
||||||
|
t.Errorf("expected no email before a login, got %q", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
const email = "user@example.com"
|
||||||
|
if err := writeProfileEmail(configPath, email); err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := readProfileEmail(configPath); got != email {
|
||||||
|
t.Errorf("got %q, want %q", got, email)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := removeProfileEmail(configPath); err != nil {
|
||||||
|
t.Fatalf("remove: %v", err)
|
||||||
|
}
|
||||||
|
if got := readProfileEmail(configPath); got != "" {
|
||||||
|
t.Errorf("expected no email after logout, got %q", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Logout may run on a never-logged-in profile, so a second remove must pass.
|
||||||
|
if err := removeProfileEmail(configPath); err != nil {
|
||||||
|
t.Fatalf("second remove should be a no-op: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteProfileEmailIgnoresEmpty(t *testing.T) {
|
||||||
|
configPath := filepath.Join(t.TempDir(), "profiles", "abc123.json")
|
||||||
|
if err := ensureDirFor(t, configPath); err != nil {
|
||||||
|
t.Fatalf("prepare dir: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
const email = "user@example.com"
|
||||||
|
if err := writeProfileEmail(configPath, email); err != nil {
|
||||||
|
t.Fatalf("write: %v", err)
|
||||||
|
}
|
||||||
|
if err := writeProfileEmail(configPath, ""); err != nil {
|
||||||
|
t.Fatalf("write empty: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := readProfileEmail(configPath); got != email {
|
||||||
|
t.Errorf("empty write clobbered the stored email: got %q, want %q", got, email)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func ensureDirFor(t *testing.T, path string) error {
|
||||||
|
t.Helper()
|
||||||
|
return os.MkdirAll(filepath.Dir(path), 0o700)
|
||||||
|
}
|
||||||
@@ -1,70 +0,0 @@
|
|||||||
//go:build android
|
|
||||||
|
|
||||||
package android
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
|
||||||
"golang.org/x/exp/maps"
|
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/internal/routemanager"
|
|
||||||
"github.com/netbirdio/netbird/route"
|
|
||||||
)
|
|
||||||
|
|
||||||
func executeRouteToggle(id string, manager routemanager.Manager,
|
|
||||||
operationName string,
|
|
||||||
routeOperation func(routes []route.NetID, allRoutes []route.NetID) error) error {
|
|
||||||
netID := route.NetID(id)
|
|
||||||
routes := []route.NetID{netID}
|
|
||||||
|
|
||||||
routesMap := manager.GetClientRoutesWithNetID()
|
|
||||||
routes = route.ExpandV6ExitPairs(routes, routesMap)
|
|
||||||
|
|
||||||
log.Debugf("%s with ids: %v", operationName, routes)
|
|
||||||
|
|
||||||
if err := routeOperation(routes, maps.Keys(routesMap)); err != nil {
|
|
||||||
log.Debugf("error when %s: %s", operationName, err)
|
|
||||||
return fmt.Errorf("error %s: %w", operationName, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
manager.TriggerSelection(manager.GetClientRoutes())
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type routeCommand interface {
|
|
||||||
toggleRoute() error
|
|
||||||
}
|
|
||||||
|
|
||||||
type selectRouteCommand struct {
|
|
||||||
route string
|
|
||||||
manager routemanager.Manager
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s selectRouteCommand) toggleRoute() error {
|
|
||||||
routeSelector := s.manager.GetRouteSelector()
|
|
||||||
if routeSelector == nil {
|
|
||||||
return fmt.Errorf("no route selector available")
|
|
||||||
}
|
|
||||||
|
|
||||||
routeOperation := func(routes []route.NetID, allRoutes []route.NetID) error {
|
|
||||||
return routeSelector.SelectRoutes(routes, true, allRoutes)
|
|
||||||
}
|
|
||||||
|
|
||||||
return executeRouteToggle(s.route, s.manager, "selecting route", routeOperation)
|
|
||||||
}
|
|
||||||
|
|
||||||
type deselectRouteCommand struct {
|
|
||||||
route string
|
|
||||||
manager routemanager.Manager
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d deselectRouteCommand) toggleRoute() error {
|
|
||||||
routeSelector := d.manager.GetRouteSelector()
|
|
||||||
if routeSelector == nil {
|
|
||||||
return fmt.Errorf("no route selector available")
|
|
||||||
}
|
|
||||||
|
|
||||||
return executeRouteToggle(d.route, d.manager, "deselecting route", routeSelector.DeselectRoutes)
|
|
||||||
}
|
|
||||||
312
client/android/session.go
Normal file
312
client/android/session.go
Normal file
@@ -0,0 +1,312 @@
|
|||||||
|
//go:build android
|
||||||
|
|
||||||
|
package android
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
cProto "github.com/netbirdio/netbird/client/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
// StateChangeListener receives client state notifications.
|
||||||
|
//
|
||||||
|
// OnStateChanged is a payload-free wake-up whenever the state snapshot
|
||||||
|
// changed: connection state, the run-loop status label (e.g. NeedsLogin) or
|
||||||
|
// the session deadline. It mirrors the daemon's SubscribeStatus stream
|
||||||
|
// trigger — on each signal the consumer pulls the fresh values via
|
||||||
|
// Status() / SessionExpiresAtUnix().
|
||||||
|
//
|
||||||
|
// OnSessionExpiring forwards the engine's session-expiry warnings, fired at
|
||||||
|
// sessionwatch.WarningLead before the deadline and again at FinalWarningLead
|
||||||
|
// (finalWarning true). The second one is suppressed when the user dismissed
|
||||||
|
// the first via DismissSessionWarning. The daemon turns the same events into
|
||||||
|
// its tray notification.
|
||||||
|
type StateChangeListener interface {
|
||||||
|
OnStateChanged()
|
||||||
|
OnSessionExpiring(expiresAtUnix int64, leadMinutes int64, finalWarning bool)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Status returns the connect run-loop's status label — the same value the
|
||||||
|
// desktop daemon serves in StatusResponse.Status. "NeedsLogin" means the
|
||||||
|
// management server rejected the peer and an interactive login is required.
|
||||||
|
//
|
||||||
|
// The label is latched: the run loop keeps its status in a per-run context
|
||||||
|
// state, which a restart replaces with a fresh Idle one, so an engine restart
|
||||||
|
// (network change, always-on) would otherwise erase the fact that the peer
|
||||||
|
// still needs to log in. Only a successful interactive login or extend clears
|
||||||
|
// it — see clearLoginRequired.
|
||||||
|
func (c *Client) Status() string {
|
||||||
|
latched, generation := c.loginRequiredState()
|
||||||
|
if latched {
|
||||||
|
return string(internal.StatusNeedsLogin)
|
||||||
|
}
|
||||||
|
cc := c.getConnectClient()
|
||||||
|
if cc == nil {
|
||||||
|
return string(internal.StatusIdle)
|
||||||
|
}
|
||||||
|
status := cc.Status()
|
||||||
|
if status == internal.StatusNeedsLogin {
|
||||||
|
c.latchLoginRequired(generation)
|
||||||
|
}
|
||||||
|
return string(status)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) loginRequiredState() (bool, uint64) {
|
||||||
|
c.loginRequiredMu.Lock()
|
||||||
|
defer c.loginRequiredMu.Unlock()
|
||||||
|
return c.loginRequired, c.loginCleared
|
||||||
|
}
|
||||||
|
|
||||||
|
// latchLoginRequired records a NeedsLogin observation, unless a clear landed
|
||||||
|
// while the caller was reading the run loop's status: cc.Status() is read
|
||||||
|
// outside the lock, so a login or extend completing in that window would
|
||||||
|
// otherwise be undone by this stale observation, stranding the UI on
|
||||||
|
// "login required" over a healthy session.
|
||||||
|
func (c *Client) latchLoginRequired(observedGeneration uint64) {
|
||||||
|
c.loginRequiredMu.Lock()
|
||||||
|
defer c.loginRequiredMu.Unlock()
|
||||||
|
if c.loginCleared != observedGeneration {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.loginRequired = true
|
||||||
|
}
|
||||||
|
|
||||||
|
// clearLoginRequired releases the latch after a successful interactive login
|
||||||
|
// or session extend, and invalidates any observation already in flight.
|
||||||
|
func (c *Client) clearLoginRequired() {
|
||||||
|
c.loginRequiredMu.Lock()
|
||||||
|
defer c.loginRequiredMu.Unlock()
|
||||||
|
c.loginRequired = false
|
||||||
|
c.loginCleared++
|
||||||
|
}
|
||||||
|
|
||||||
|
// SessionExpiresAtUnix returns the SSO session deadline as unix seconds, or 0
|
||||||
|
// when no deadline is known (not SSO-registered, expiry disabled, or the
|
||||||
|
// engine has not received one yet). A past value means the session expired.
|
||||||
|
// Mirror of StatusResponse.sessionExpiresAt on the desktop daemon.
|
||||||
|
func (c *Client) SessionExpiresAtUnix() int64 {
|
||||||
|
deadline := c.recorder.GetSessionExpiresAt()
|
||||||
|
if deadline.IsZero() {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return deadline.Unix()
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetStateChangeListener registers the state notification listener.
|
||||||
|
// Replaces any previously registered listener; remove it with
|
||||||
|
// RemoveStateChangeListener.
|
||||||
|
func (c *Client) SetStateChangeListener(listener StateChangeListener) {
|
||||||
|
c.stateChangeMu.Lock()
|
||||||
|
defer c.stateChangeMu.Unlock()
|
||||||
|
c.stopStateChangeWatchLocked()
|
||||||
|
if listener == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Both subscriptions are buffered (one pending tick, ten pending events),
|
||||||
|
// so unsubscribing is not enough to stop callbacks: the loops would drain
|
||||||
|
// what is already queued and deliver it to a listener the caller has
|
||||||
|
// already removed or replaced. Gate every callback on this registration's
|
||||||
|
// own signal, which is closed before unsubscribing.
|
||||||
|
done := make(chan struct{})
|
||||||
|
c.stateChangeDone = done
|
||||||
|
|
||||||
|
id, ch := c.recorder.SubscribeToStateChanges()
|
||||||
|
c.stateChangeSubID = id
|
||||||
|
// The channel is closed by UnsubscribeFromStateChanges, which ends the
|
||||||
|
// goroutine. Ticks are coalesced (buffer of one), so a burst of changes
|
||||||
|
// wakes the listener once.
|
||||||
|
go func() {
|
||||||
|
for range ch {
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
listener.OnStateChanged()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
c.eventSub = c.recorder.SubscribeToEvents()
|
||||||
|
go watchSessionWarnings(c.eventSub, listener, done)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RemoveStateChangeListener unregisters the state notification listener.
|
||||||
|
func (c *Client) RemoveStateChangeListener() {
|
||||||
|
c.stateChangeMu.Lock()
|
||||||
|
defer c.stateChangeMu.Unlock()
|
||||||
|
c.stopStateChangeWatchLocked()
|
||||||
|
}
|
||||||
|
|
||||||
|
// DismissSessionWarning records the user's "Dismiss" on the first expiry
|
||||||
|
// warning and suppresses the final one for the current deadline. A refreshed
|
||||||
|
// deadline re-arms both. No-op while the engine is not running.
|
||||||
|
func (c *Client) DismissSessionWarning() {
|
||||||
|
cc := c.getConnectClient()
|
||||||
|
if cc == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
engine := cc.Engine()
|
||||||
|
if engine == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
engine.DismissSessionWarning()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExtendAuthSession runs the interactive SSO flow to obtain a fresh JWT and
|
||||||
|
// asks the management server to extend the session deadline. The tunnel is
|
||||||
|
// untouched: no resync, no reconnect. Async; the result arrives on the
|
||||||
|
// listener. Mirror of the daemon's RequestExtendAuthSession /
|
||||||
|
// WaitExtendAuthSession RPC pair, with URLOpener playing the "UI opens the
|
||||||
|
// browser" role.
|
||||||
|
//
|
||||||
|
// Only one flow may be in flight: the PKCE step binds a fixed loopback port,
|
||||||
|
// so a second concurrent flow would fail on that bind. Call
|
||||||
|
// CancelExtendAuthSession when the user abandons the browser.
|
||||||
|
func (c *Client) ExtendAuthSession(urlOpener URLOpener, isAndroidTV bool, resultListener ErrListener) {
|
||||||
|
ctx, err := c.beginExtend()
|
||||||
|
if err != nil {
|
||||||
|
resultListener.OnError(err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
defer c.endExtend()
|
||||||
|
if err := c.extendAuthSession(ctx, urlOpener, isAndroidTV); err != nil {
|
||||||
|
resultListener.OnError(err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
resultListener.OnSuccess()
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// CancelExtendAuthSession aborts an in-flight ExtendAuthSession. The tunnel is
|
||||||
|
// left alone — unlike the login flow, which cancels the whole client context
|
||||||
|
// by stopping the engine. Without this the abandoned PKCE wait keeps its
|
||||||
|
// loopback port for the full flow timeout and blocks every later attempt.
|
||||||
|
// No-op when no flow is running.
|
||||||
|
func (c *Client) CancelExtendAuthSession() {
|
||||||
|
c.extendMu.Lock()
|
||||||
|
defer c.extendMu.Unlock()
|
||||||
|
if c.extendCancel != nil {
|
||||||
|
c.extendCancel()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) stopStateChangeWatchLocked() {
|
||||||
|
// Signal first, unsubscribe second: closing the channels only stops new
|
||||||
|
// items, and the loops would still hand whatever is buffered to a listener
|
||||||
|
// that is no longer registered.
|
||||||
|
if c.stateChangeDone != nil {
|
||||||
|
close(c.stateChangeDone)
|
||||||
|
c.stateChangeDone = nil
|
||||||
|
}
|
||||||
|
if c.stateChangeSubID != "" {
|
||||||
|
c.recorder.UnsubscribeFromStateChanges(c.stateChangeSubID)
|
||||||
|
c.stateChangeSubID = ""
|
||||||
|
}
|
||||||
|
if c.eventSub != nil {
|
||||||
|
// Closes the channel, which ends watchSessionWarnings.
|
||||||
|
c.recorder.UnsubscribeFromEvents(c.eventSub)
|
||||||
|
c.eventSub = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// watchSessionWarnings forwards the engine's session-expiry warnings to the
|
||||||
|
// listener. The event stream also carries unrelated traffic — network-map
|
||||||
|
// updates on every sync, DNS and route errors — so everything but an
|
||||||
|
// AUTHENTICATION event carrying the session-warning marker is dropped. Exits
|
||||||
|
// when the subscription is closed by UnsubscribeFromEvents, or earlier when
|
||||||
|
// done is closed — the stream buffers up to ten events, and a deregistered
|
||||||
|
// listener must not receive the ones already queued.
|
||||||
|
func watchSessionWarnings(sub *peer.EventSubscription, listener StateChangeListener, done <-chan struct{}) {
|
||||||
|
for ev := range sub.Events() {
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
if ev.GetCategory() != cProto.SystemEvent_AUTHENTICATION {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
meta := ev.GetMetadata()
|
||||||
|
if meta[sessionwatch.MetaSessionWarning] != "true" {
|
||||||
|
// Other AUTHENTICATION events exist (e.g. a deadline rejected as
|
||||||
|
// out of range); they carry no warning marker.
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
deadline, err := sessionwatch.ParseExpiresAt(meta[sessionwatch.MetaSessionExpiresAt])
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("session warning event with unparsable deadline: %v", err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
lead, err := sessionwatch.ParseLeadMinutes(meta[sessionwatch.MetaSessionLeadMinutes])
|
||||||
|
if err != nil {
|
||||||
|
// Informational only — the deadline above is what drives the UI.
|
||||||
|
lead = 0
|
||||||
|
}
|
||||||
|
listener.OnSessionExpiring(deadline.Unix(), int64(lead),
|
||||||
|
meta[sessionwatch.MetaSessionFinal] == "true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) beginExtend() (context.Context, error) {
|
||||||
|
c.extendMu.Lock()
|
||||||
|
defer c.extendMu.Unlock()
|
||||||
|
if c.extendCancel != nil {
|
||||||
|
return nil, fmt.Errorf("session extend already in progress")
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
c.extendCancel = cancel
|
||||||
|
return ctx, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) endExtend() {
|
||||||
|
c.extendMu.Lock()
|
||||||
|
defer c.extendMu.Unlock()
|
||||||
|
if c.extendCancel != nil {
|
||||||
|
c.extendCancel()
|
||||||
|
c.extendCancel = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isAndroidTV bool) error {
|
||||||
|
cfg, cfgPath, cc := c.authSnapshot()
|
||||||
|
if cfg == nil || cc == nil {
|
||||||
|
return fmt.Errorf("engine is not running")
|
||||||
|
}
|
||||||
|
engine := cc.Engine()
|
||||||
|
if engine == nil {
|
||||||
|
return fmt.Errorf("engine is not initialized")
|
||||||
|
}
|
||||||
|
|
||||||
|
authClient, err := auth.NewAuth(ctx, cfg.PrivateKey, cfg.ManagementURL, cfg)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to create auth client: %v", err)
|
||||||
|
}
|
||||||
|
defer authClient.Close()
|
||||||
|
|
||||||
|
// Passing the config path makes the flow pick up the login_hint: an extend
|
||||||
|
// renews the session of the account already signed in, so it must not stop to
|
||||||
|
// offer a choice.
|
||||||
|
a := NewAuthWithConfig(ctx, cfg, cfgPath)
|
||||||
|
tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("interactive sso login failed: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := engine.ExtendAuthSession(ctx, tokenInfo.GetTokenToUse()); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
c.clearLoginRequired()
|
||||||
|
|
||||||
|
go urlOpener.OnLoginSuccess()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
66
client/cmd/daemon_error.go
Normal file
66
client/cmd/daemon_error.go
Normal file
@@ -0,0 +1,66 @@
|
|||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"google.golang.org/genproto/googleapis/rpc/errdetails"
|
||||||
|
gstatus "google.golang.org/grpc/status"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||||
|
)
|
||||||
|
|
||||||
|
// daemonCallError prepares a daemon error for display. A refusal the daemon
|
||||||
|
// raised because the operation needs root/administrator is already guidance
|
||||||
|
// written for the user, so it is surfaced on its own instead of buried under the
|
||||||
|
// gRPC envelope and the name of the RPC that hit it. Anything else is wrapped
|
||||||
|
// with context as usual.
|
||||||
|
func daemonCallError(context string, err error) error {
|
||||||
|
if guidance, ok := privilegeGuidance(err); ok {
|
||||||
|
return errors.New(guidance)
|
||||||
|
}
|
||||||
|
return fmt.Errorf("%s: %w", context, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// privilegeGuidance renders the daemon's privilege refusal as a summary and the
|
||||||
|
// command that performs the operation with the privileges it needs. It reports
|
||||||
|
// false for any other error.
|
||||||
|
func privilegeGuidance(err error) (string, bool) {
|
||||||
|
info, ok := privilegeErrorInfo(err)
|
||||||
|
if !ok {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
|
||||||
|
summary := info.GetMetadata()[ipcauth.ErrorMetaSummary]
|
||||||
|
command := info.GetMetadata()[ipcauth.ErrorMetaCommand]
|
||||||
|
if summary == "" {
|
||||||
|
// Detail without a summary: fall back to the status message, which
|
||||||
|
// carries the same text.
|
||||||
|
summary = strings.TrimSpace(gstatus.Convert(err).Message())
|
||||||
|
}
|
||||||
|
if command == "" {
|
||||||
|
return summary, true
|
||||||
|
}
|
||||||
|
|
||||||
|
return fmt.Sprintf("%s\n\n %s\n", summary, command), true
|
||||||
|
}
|
||||||
|
|
||||||
|
// privilegeErrorInfo returns the daemon's privilege-refusal detail, if the error
|
||||||
|
// carries one.
|
||||||
|
func privilegeErrorInfo(err error) (*errdetails.ErrorInfo, bool) {
|
||||||
|
if err == nil {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, detail := range gstatus.Convert(err).Details() {
|
||||||
|
info, ok := detail.(*errdetails.ErrorInfo)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if info.GetReason() == ipcauth.ErrorReasonPrivilegeRequired && info.GetDomain() == ipcauth.ErrorDomain {
|
||||||
|
return info, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
@@ -17,16 +17,26 @@ import (
|
|||||||
"github.com/netbirdio/netbird/client/internal"
|
"github.com/netbirdio/netbird/client/internal"
|
||||||
"github.com/netbirdio/netbird/client/internal/auth"
|
"github.com/netbirdio/netbird/client/internal/auth"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
"github.com/netbirdio/netbird/client/proto"
|
"github.com/netbirdio/netbird/client/proto"
|
||||||
|
"github.com/netbirdio/netbird/client/server"
|
||||||
"github.com/netbirdio/netbird/client/system"
|
"github.com/netbirdio/netbird/client/system"
|
||||||
"github.com/netbirdio/netbird/util"
|
"github.com/netbirdio/netbird/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// extendSessionFlag drives the `netbird login --extend` flow: refresh the
|
||||||
|
// SSO session expiry on the management server without tearing down the
|
||||||
|
// tunnel. Mutually exclusive with setup-key login (a setup-key cannot
|
||||||
|
// refresh an SSO-tracked peer — see auth.errSetupKeyOnSSOExpiredPeer).
|
||||||
|
var extendSessionFlag bool
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
loginCmd.PersistentFlags().BoolVar(&noBrowser, noBrowserFlag, false, noBrowserDesc)
|
loginCmd.PersistentFlags().BoolVar(&noBrowser, noBrowserFlag, false, noBrowserDesc)
|
||||||
loginCmd.PersistentFlags().BoolVar(&showQR, showQRFlag, false, showQRDesc)
|
loginCmd.PersistentFlags().BoolVar(&showQR, showQRFlag, false, showQRDesc)
|
||||||
loginCmd.PersistentFlags().StringVar(&profileName, profileNameFlag, "", profileNameDesc)
|
loginCmd.PersistentFlags().StringVar(&profileName, profileNameFlag, "", profileNameDesc)
|
||||||
loginCmd.PersistentFlags().StringVarP(&configPath, "config", "c", "", "(DEPRECATED) Netbird config file location")
|
loginCmd.PersistentFlags().StringVarP(&configPath, "config", "c", "", "(DEPRECATED) Netbird config file location")
|
||||||
|
loginCmd.PersistentFlags().BoolVar(&extendSessionFlag, "extend", false,
|
||||||
|
"refresh the SSO session expiry without tearing down the tunnel (requires an active connection)")
|
||||||
}
|
}
|
||||||
|
|
||||||
var loginCmd = &cobra.Command{
|
var loginCmd = &cobra.Command{
|
||||||
@@ -61,6 +71,16 @@ var loginCmd = &cobra.Command{
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if extendSessionFlag {
|
||||||
|
if providedSetupKey != "" {
|
||||||
|
return fmt.Errorf("--extend cannot be combined with a setup key; setup keys can only enrol new peers")
|
||||||
|
}
|
||||||
|
if err := doExtendSession(ctx, cmd); err != nil {
|
||||||
|
return fmt.Errorf("extend session failed: %v", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// workaround to run without service
|
// workaround to run without service
|
||||||
if util.FindFirstLogPath(logFiles) == "" {
|
if util.FindFirstLogPath(logFiles) == "" {
|
||||||
if err := doForegroundLogin(ctx, cmd, providedSetupKey, activeProf); err != nil {
|
if err := doForegroundLogin(ctx, cmd, providedSetupKey, activeProf); err != nil {
|
||||||
@@ -152,6 +172,65 @@ func doDaemonLogin(ctx context.Context, cmd *cobra.Command, providedSetupKey str
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// doExtendSession drives the daemon's RequestExtendAuthSession /
|
||||||
|
// WaitExtendAuthSession pair. The user is sent through a regular SSO flow
|
||||||
|
// (browser + verification URL) and the resulting JWT is forwarded to the
|
||||||
|
// management server's ExtendAuthSession RPC. The tunnel stays up
|
||||||
|
// throughout — no Down/Up, no network-map resync.
|
||||||
|
func doExtendSession(ctx context.Context, cmd *cobra.Command) error {
|
||||||
|
conn, err := DialClientGRPCServer(ctx, daemonAddr)
|
||||||
|
if err != nil {
|
||||||
|
//nolint
|
||||||
|
return fmt.Errorf("failed to connect to daemon error: %v\n"+
|
||||||
|
"If the daemon is not running please run: "+
|
||||||
|
"\nnetbird service install \nnetbird service start\n", err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
client := proto.NewDaemonServiceClient(conn)
|
||||||
|
|
||||||
|
req := &proto.RequestExtendAuthSessionRequest{}
|
||||||
|
// Pre-fill the IdP login hint from the active profile so the user
|
||||||
|
// doesn't have to retype their email. Best-effort: we still proceed
|
||||||
|
// without a hint if the lookup fails.
|
||||||
|
pm := profilemanager.NewProfileManager()
|
||||||
|
if active, perr := pm.GetActiveProfile(); perr == nil {
|
||||||
|
if profState, sperr := pm.GetProfileState(active.ID); sperr == nil && profState.Email != "" {
|
||||||
|
req.Hint = &profState.Email
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
startResp, err := client.RequestExtendAuthSession(ctx, req)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("start extend session: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
uri := startResp.GetVerificationURIComplete()
|
||||||
|
if uri == "" {
|
||||||
|
uri = startResp.GetVerificationURI()
|
||||||
|
}
|
||||||
|
openURL(cmd, uri, startResp.GetUserCode(), noBrowser, showQR)
|
||||||
|
|
||||||
|
waitResp, err := client.WaitExtendAuthSession(ctx, &proto.WaitExtendAuthSessionRequest{
|
||||||
|
DeviceCode: startResp.GetDeviceCode(),
|
||||||
|
UserCode: startResp.GetUserCode(),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("wait for extend session: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if ts := waitResp.GetSessionExpiresAt(); ts.IsValid() && !ts.AsTime().IsZero() {
|
||||||
|
deadline := ts.AsTime().Local()
|
||||||
|
cmd.Printf("Session extended. New expiry: %s\n", deadline.Format("2006-01-02 15:04:05 MST"))
|
||||||
|
} else {
|
||||||
|
// Management reported the peer is not eligible (e.g. login
|
||||||
|
// expiration disabled on the account). Surface that fact
|
||||||
|
// instead of pretending the call succeeded.
|
||||||
|
cmd.Println("Session extension call completed, but the management server did not return a new deadline (peer may not be SSO-tracked or login expiration is disabled).")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func getActiveProfile(ctx context.Context, pm *profilemanager.ProfileManager, profileName string, username string) (*profilemanager.Profile, error) {
|
func getActiveProfile(ctx context.Context, pm *profilemanager.ProfileManager, profileName string, username string) (*profilemanager.Profile, error) {
|
||||||
// switch profile if provided
|
// switch profile if provided
|
||||||
|
|
||||||
@@ -254,6 +333,14 @@ func doForegroundLogin(ctx context.Context, cmd *cobra.Command, setupKey string,
|
|||||||
return fmt.Errorf("read config file %s: %v", configFilePath, err)
|
return fmt.Errorf("read config file %s: %v", configFilePath, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Mirror runInForegroundMode: recover residual state (DNS, firewall,
|
||||||
|
// ssh config, legacy routing) from a previous unclean shutdown and
|
||||||
|
// enable advanced routing before dialing management.
|
||||||
|
if err := server.RestoreResidualState(ctx, profilemanager.NewServiceManager(configFilePath).GetStatePath()); err != nil {
|
||||||
|
log.Warnf("failed to restore residual state: %v", err)
|
||||||
|
}
|
||||||
|
nbnet.Init()
|
||||||
|
|
||||||
err = foregroundLogin(ctx, cmd, config, setupKey, activeProf.ID)
|
err = foregroundLogin(ctx, cmd, config, setupKey, activeProf.ID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("foreground login failed: %v", err)
|
return fmt.Errorf("foreground login failed: %v", err)
|
||||||
|
|||||||
@@ -46,7 +46,7 @@ var logoutCmd = &cobra.Command{
|
|||||||
}
|
}
|
||||||
|
|
||||||
if _, err := daemonClient.Logout(ctx, req); err != nil {
|
if _, err := daemonClient.Logout(ctx, req); err != nil {
|
||||||
return fmt.Errorf("deregister: %v", err)
|
return daemonCallError("deregister", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
cmd.Println("Deregistered successfully")
|
cmd.Println("Deregistered successfully")
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ import (
|
|||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
"github.com/spf13/pflag"
|
"github.com/spf13/pflag"
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
"google.golang.org/grpc/credentials/insecure"
|
|
||||||
|
|
||||||
daddr "github.com/netbirdio/netbird/client/internal/daemonaddr"
|
daddr "github.com/netbirdio/netbird/client/internal/daemonaddr"
|
||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
@@ -91,6 +90,7 @@ var (
|
|||||||
// Don't resolve for service commands — they create the socket, not connect to it.
|
// Don't resolve for service commands — they create the socket, not connect to it.
|
||||||
if !isServiceCmd(cmd) {
|
if !isServiceCmd(cmd) {
|
||||||
daemonAddr = daddr.ResolveUnixDaemonAddr(daemonAddr)
|
daemonAddr = daddr.ResolveUnixDaemonAddr(daemonAddr)
|
||||||
|
daemonAddr = daddr.ResolveDaemonAddr(daemonAddr)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
},
|
},
|
||||||
@@ -143,10 +143,10 @@ func init() {
|
|||||||
|
|
||||||
defaultDaemonAddr := "unix:///var/run/netbird.sock"
|
defaultDaemonAddr := "unix:///var/run/netbird.sock"
|
||||||
if runtime.GOOS == "windows" {
|
if runtime.GOOS == "windows" {
|
||||||
defaultDaemonAddr = "tcp://127.0.0.1:41731"
|
defaultDaemonAddr = daddr.WindowsPipeAddr
|
||||||
}
|
}
|
||||||
|
|
||||||
rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp]://[path|host:port]")
|
rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp|npipe]://[path|host:port|name]")
|
||||||
rootCmd.PersistentFlags().StringVarP(&managementURL, "management-url", "m", "", fmt.Sprintf("Management Service URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultManagementURL))
|
rootCmd.PersistentFlags().StringVarP(&managementURL, "management-url", "m", "", fmt.Sprintf("Management Service URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultManagementURL))
|
||||||
rootCmd.PersistentFlags().StringVar(&adminURL, "admin-url", "", fmt.Sprintf("Admin Panel URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultAdminURL))
|
rootCmd.PersistentFlags().StringVar(&adminURL, "admin-url", "", fmt.Sprintf("Admin Panel URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultAdminURL))
|
||||||
rootCmd.PersistentFlags().StringVarP(&logLevel, "log-level", "l", "info", "sets NetBird log level")
|
rootCmd.PersistentFlags().StringVarP(&logLevel, "log-level", "l", "info", "sets NetBird log level")
|
||||||
@@ -269,12 +269,10 @@ func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, e
|
|||||||
ctx, cancel := context.WithTimeout(ctx, time.Second*10)
|
ctx, cancel := context.WithTimeout(ctx, time.Second*10)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
return grpc.DialContext(
|
target, opts := daddr.DialTarget(addr)
|
||||||
ctx,
|
opts = append(opts, grpc.WithBlock())
|
||||||
strings.TrimPrefix(addr, "tcp://"),
|
|
||||||
grpc.WithTransportCredentials(insecure.NewCredentials()),
|
return grpc.DialContext(ctx, target, opts...)
|
||||||
grpc.WithBlock(),
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// WithBackOff execute function in backoff cycle.
|
// WithBackOff execute function in backoff cycle.
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net/http"
|
||||||
"runtime"
|
"runtime"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -22,15 +23,26 @@ var serviceCmd = &cobra.Command{
|
|||||||
Short: "Manage the NetBird daemon service",
|
Short: "Manage the NetBird daemon service",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const defaultJSONSocket = "unix:///var/run/netbird-http.sock"
|
||||||
|
|
||||||
var (
|
var (
|
||||||
serviceName string
|
serviceName string
|
||||||
serviceEnvVars []string
|
serviceEnvVars []string
|
||||||
|
jsonSocket string
|
||||||
|
enableJSONSocket bool
|
||||||
)
|
)
|
||||||
|
|
||||||
type program struct {
|
type program struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
serv *grpc.Server
|
serv *grpc.Server
|
||||||
|
jsonServ *http.Server
|
||||||
|
// jsonClient is the gateway's own connection to the daemon. It is held so
|
||||||
|
// shutting the gateway down also closes it: nothing else references it once
|
||||||
|
// the handlers are registered, so its transport goroutines would otherwise
|
||||||
|
// outlive the server.
|
||||||
|
jsonClient *grpc.ClientConn
|
||||||
|
jsonServMu sync.Mutex
|
||||||
serverInstance *server.Server
|
serverInstance *server.Server
|
||||||
serverInstanceMu sync.Mutex
|
serverInstanceMu sync.Mutex
|
||||||
}
|
}
|
||||||
@@ -46,6 +58,8 @@ func init() {
|
|||||||
serviceCmd.PersistentFlags().BoolVar(&updateSettingsDisabled, "disable-update-settings", false, "Disables update settings feature. If enabled, the client will not be able to change or edit any settings. To persist this setting, use: netbird service install --disable-update-settings")
|
serviceCmd.PersistentFlags().BoolVar(&updateSettingsDisabled, "disable-update-settings", false, "Disables update settings feature. If enabled, the client will not be able to change or edit any settings. To persist this setting, use: netbird service install --disable-update-settings")
|
||||||
serviceCmd.PersistentFlags().BoolVar(&captureEnabled, "enable-capture", false, "Enables packet capture via 'netbird debug capture'. To persist, use: netbird service install --enable-capture")
|
serviceCmd.PersistentFlags().BoolVar(&captureEnabled, "enable-capture", false, "Enables packet capture via 'netbird debug capture'. To persist, use: netbird service install --enable-capture")
|
||||||
serviceCmd.PersistentFlags().BoolVar(&networksDisabled, "disable-networks", false, "Disables network selection. If enabled, the client will not allow listing, selecting, or deselecting networks. To persist, use: netbird service install --disable-networks")
|
serviceCmd.PersistentFlags().BoolVar(&networksDisabled, "disable-networks", false, "Disables network selection. If enabled, the client will not allow listing, selecting, or deselecting networks. To persist, use: netbird service install --disable-networks")
|
||||||
|
serviceCmd.PersistentFlags().BoolVar(&enableJSONSocket, "enable-json-socket", false, "Enables the HTTP/JSON API socket served by grpc-gateway. To persist, use: netbird service install --enable-json-socket")
|
||||||
|
serviceCmd.PersistentFlags().StringVar(&jsonSocket, "json-socket", defaultJSONSocket, "HTTP/JSON API socket address [unix|tcp]://[path|host:port]. Requires --enable-json-socket to serve. To persist, use: netbird service install --enable-json-socket --json-socket")
|
||||||
|
|
||||||
rootCmd.PersistentFlags().StringVarP(&serviceName, "service", "s", defaultServiceName, "Netbird system service name")
|
rootCmd.PersistentFlags().StringVarP(&serviceName, "service", "s", defaultServiceName, "Netbird system service name")
|
||||||
serviceEnvDesc := `Sets extra environment variables for the service. ` +
|
serviceEnvDesc := `Sets extra environment variables for the service. ` +
|
||||||
|
|||||||
@@ -5,9 +5,7 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"runtime"
|
||||||
"os"
|
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
@@ -16,69 +14,157 @@ import (
|
|||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/daemonaddr"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||||
"github.com/netbirdio/netbird/client/proto"
|
"github.com/netbirdio/netbird/client/proto"
|
||||||
"github.com/netbirdio/netbird/client/server"
|
"github.com/netbirdio/netbird/client/server"
|
||||||
"github.com/netbirdio/netbird/client/system"
|
"github.com/netbirdio/netbird/client/system"
|
||||||
"github.com/netbirdio/netbird/util"
|
"github.com/netbirdio/netbird/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func validateJSONSocketFlags() error {
|
||||||
|
if serviceCmd.PersistentFlags().Changed("json-socket") && !enableJSONSocket {
|
||||||
|
return fmt.Errorf("--json-socket requires --enable-json-socket to configure the daemon JSON gateway")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// daemonServerOptions installs the transport credentials that expose each
|
||||||
|
// caller's kernel-authenticated identity to the handlers, which is what lets
|
||||||
|
// the daemon require root/administrator for privileged operations.
|
||||||
|
//
|
||||||
|
// The handshake exchanges no bytes, so older CLI and UI binaries still
|
||||||
|
// interoperate. Callers on a TCP socket carry no identity at all: the daemon
|
||||||
|
// keeps serving them, and the privileged operations deny them, so a warning is
|
||||||
|
// logged to make the loss of functionality visible.
|
||||||
|
func daemonServerOptions(network string) []grpc.ServerOption {
|
||||||
|
if network == "tcp" {
|
||||||
|
log.Warnf("daemon is listening on TCP (%s): callers carry no verifiable identity over TCP, "+
|
||||||
|
"so privileged operations (SSH root login, SSH auth, enabling the SSH server, management URL changes, "+
|
||||||
|
"deregistration) will be denied. Use a unix socket, or npipe:// on Windows", daemonAddr)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
creds := ipcauth.NewTransportCredentials()
|
||||||
|
if creds == nil {
|
||||||
|
log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied", runtime.GOOS)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return []grpc.ServerOption{grpc.Creds(creds)}
|
||||||
|
}
|
||||||
|
|
||||||
func (p *program) Start(svc service.Service) error {
|
func (p *program) Start(svc service.Service) error {
|
||||||
// Start should not block. Do the actual work async.
|
// Start should not block. Do the actual work async.
|
||||||
log.Info("starting NetBird service") //nolint
|
log.Info("starting NetBird service") //nolint
|
||||||
|
|
||||||
|
if err := validateJSONSocketFlags(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
// Collect static system and platform information
|
// Collect static system and platform information
|
||||||
system.UpdateStaticInfoAsync()
|
system.UpdateStaticInfoAsync()
|
||||||
|
|
||||||
// in any case, even if configuration does not exists we run daemon to serve CLI gRPC API.
|
// A daemon installed before named-pipe support has the loopback TCP address
|
||||||
p.serv = grpc.NewServer()
|
// persisted. Move it to the named pipe so an upgraded daemon can identify
|
||||||
|
// its callers instead of silently serving an unauthenticated socket.
|
||||||
split := strings.Split(daemonAddr, "://")
|
if migrated, ok := daemonaddr.MigrateLegacy(daemonAddr); ok {
|
||||||
switch split[0] {
|
log.Infof("daemon address %q predates named-pipe support, listening on %q so callers can be identified", daemonAddr, migrated)
|
||||||
case "unix":
|
daemonAddr = migrated
|
||||||
// cleanup failed close
|
|
||||||
stat, err := os.Stat(split[1])
|
|
||||||
if err == nil && !stat.IsDir() {
|
|
||||||
if err := os.Remove(split[1]); err != nil {
|
|
||||||
log.Debugf("remove socket file: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case "tcp":
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("unsupported daemon address protocol: %v", split[0])
|
|
||||||
}
|
}
|
||||||
|
|
||||||
listen, err := net.Listen(split[0], split[1])
|
network, _, err := parseListenAddress(daemonAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("listen daemon interface: %w", err)
|
return fmt.Errorf("parse daemon address: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// in any case, even if configuration does not exists we run daemon to serve CLI gRPC API.
|
||||||
|
p.serv = grpc.NewServer(daemonServerOptions(network)...)
|
||||||
|
|
||||||
|
daemonListener, jsonListener, err := listenDaemonSockets()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
defer listen.Close()
|
// Fatal here rather than inside serve, so serve's deferred listener
|
||||||
|
// closes run before the process exits.
|
||||||
if split[0] == "unix" {
|
if err := p.serve(daemonListener, jsonListener); err != nil {
|
||||||
if err := os.Chmod(split[1], 0666); err != nil {
|
log.Fatalf("failed to %v", err)
|
||||||
log.Errorf("failed setting daemon permissions: %v", split[1])
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
serverInstance := server.New(p.ctx, util.FindFirstLogPath(logFiles), configPath, profilesDisabled, updateSettingsDisabled, captureEnabled, networksDisabled)
|
|
||||||
if err := serverInstance.Start(); err != nil {
|
|
||||||
log.Fatalf("failed to start daemon: %v", err)
|
|
||||||
}
|
|
||||||
proto.RegisterDaemonServiceServer(p.serv, serverInstance)
|
|
||||||
|
|
||||||
p.serverInstanceMu.Lock()
|
|
||||||
p.serverInstance = serverInstance
|
|
||||||
p.serverInstanceMu.Unlock()
|
|
||||||
|
|
||||||
log.Printf("started daemon server: %v", split[1])
|
|
||||||
if err := p.serv.Serve(listen); err != nil {
|
|
||||||
log.Errorf("failed to serve daemon requests: %v", err)
|
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// listenDaemonSockets opens the daemon control socket and, when it is enabled, the
|
||||||
|
// JSON gateway socket. The control socket is closed again if the second one fails,
|
||||||
|
// so a failed start leaves nothing listening. The returned JSON listener is nil
|
||||||
|
// when the socket is disabled.
|
||||||
|
func listenDaemonSockets() (*socketListener, *socketListener, error) {
|
||||||
|
daemonListener, err := listenOnAddress(daemonAddr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, fmt.Errorf("listen daemon interface: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !enableJSONSocket {
|
||||||
|
removeStaleUnixSocketForAddress(jsonSocket)
|
||||||
|
return daemonListener, nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
jsonListener, err := listenOnAddress(jsonSocket)
|
||||||
|
if err != nil {
|
||||||
|
if cerr := daemonListener.Close(); cerr != nil {
|
||||||
|
log.Debugf("close daemon listener: %v", cerr)
|
||||||
|
}
|
||||||
|
return nil, nil, fmt.Errorf("listen daemon JSON interface: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return daemonListener, jsonListener, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// serve brings up the daemon server on an already-open control socket and blocks
|
||||||
|
// until it stops. jsonListener is nil when the JSON socket is disabled. A returned
|
||||||
|
// error means the daemon cannot run at all and the caller is expected to exit; the
|
||||||
|
// failures it recovers from on its own are logged here.
|
||||||
|
func (p *program) serve(daemonListener, jsonListener *socketListener) error {
|
||||||
|
defer daemonListener.Close()
|
||||||
|
if jsonListener != nil {
|
||||||
|
defer jsonListener.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// chmodUnixSocket is a no-op for a nil listener and for a non-unix one.
|
||||||
|
if err := daemonListener.chmodUnixSocket("daemon"); err != nil {
|
||||||
|
log.Error(err)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err := jsonListener.chmodUnixSocket("daemon JSON"); err != nil {
|
||||||
|
log.Error(err)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
serverInstance := server.New(p.ctx, util.FindFirstLogPath(logFiles), configPath, profilesDisabled, updateSettingsDisabled, captureEnabled, networksDisabled)
|
||||||
|
if err := serverInstance.Start(); err != nil {
|
||||||
|
return fmt.Errorf("start daemon: %w", err)
|
||||||
|
}
|
||||||
|
proto.RegisterDaemonServiceServer(p.serv, serverInstance)
|
||||||
|
|
||||||
|
p.serverInstanceMu.Lock()
|
||||||
|
p.serverInstance = serverInstance
|
||||||
|
p.serverInstanceMu.Unlock()
|
||||||
|
|
||||||
|
if jsonListener == nil {
|
||||||
|
log.Debug("daemon JSON socket disabled")
|
||||||
|
} else if err := p.startJSONGateway(jsonListener, daemonAddr); err != nil {
|
||||||
|
return fmt.Errorf("start daemon JSON server: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Printf("started daemon server: %v", daemonListener.address)
|
||||||
|
if err := p.serv.Serve(daemonListener.Listener); err != nil {
|
||||||
|
log.Errorf("failed to serve daemon requests: %v", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (p *program) Stop(srv service.Service) error {
|
func (p *program) Stop(srv service.Service) error {
|
||||||
p.serverInstanceMu.Lock()
|
p.serverInstanceMu.Lock()
|
||||||
if p.serverInstance != nil {
|
if p.serverInstance != nil {
|
||||||
@@ -92,6 +178,25 @@ func (p *program) Stop(srv service.Service) error {
|
|||||||
|
|
||||||
p.cancel()
|
p.cancel()
|
||||||
|
|
||||||
|
p.jsonServMu.Lock()
|
||||||
|
jsonServ, jsonClient := p.jsonServ, p.jsonClient
|
||||||
|
p.jsonServMu.Unlock()
|
||||||
|
if jsonClient != nil {
|
||||||
|
if err := jsonClient.Close(); err != nil {
|
||||||
|
log.Debugf("close daemon JSON gateway client: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if jsonServ != nil {
|
||||||
|
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||||
|
if err := jsonServ.Shutdown(shutdownCtx); err != nil {
|
||||||
|
log.Errorf("failed to stop daemon JSON server gracefully: %v", err)
|
||||||
|
if err := jsonServ.Close(); err != nil {
|
||||||
|
log.Errorf("failed to close daemon JSON server: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
shutdownCancel()
|
||||||
|
}
|
||||||
|
|
||||||
if p.serv != nil {
|
if p.serv != nil {
|
||||||
p.serv.Stop()
|
p.serv.Stop()
|
||||||
}
|
}
|
||||||
@@ -148,6 +253,9 @@ var runCmd = &cobra.Command{
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
if err := validateJSONSocketFlags(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
return s.Run()
|
return s.Run()
|
||||||
},
|
},
|
||||||
@@ -162,6 +270,9 @@ var startCmd = &cobra.Command{
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
if err := validateJSONSocketFlags(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
if err := s.Start(); err != nil {
|
if err := s.Start(); err != nil {
|
||||||
return fmt.Errorf("start service: %w", err)
|
return fmt.Errorf("start service: %w", err)
|
||||||
@@ -198,6 +309,9 @@ var restartCmd = &cobra.Command{
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
if err := validateJSONSocketFlags(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
if err := s.Restart(); err != nil {
|
if err := s.Restart(); err != nil {
|
||||||
return fmt.Errorf("restart service: %w", err)
|
return fmt.Errorf("restart service: %w", err)
|
||||||
|
|||||||
@@ -67,6 +67,10 @@ func buildServiceArguments() []string {
|
|||||||
args = append(args, "--disable-networks")
|
args = append(args, "--disable-networks")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if enableJSONSocket {
|
||||||
|
args = append(args, "--enable-json-socket", "--json-socket", jsonSocket)
|
||||||
|
}
|
||||||
|
|
||||||
return args
|
return args
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -106,6 +110,10 @@ func configurePlatformSpecificSettings(svcConfig *service.Config) error {
|
|||||||
|
|
||||||
// Create fully configured service config for install/reconfigure
|
// Create fully configured service config for install/reconfigure
|
||||||
func createServiceConfigForInstall() (*service.Config, error) {
|
func createServiceConfigForInstall() (*service.Config, error) {
|
||||||
|
if err := validateJSONSocketFlags(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
svcConfig, err := newSVCConfig()
|
svcConfig, err := newSVCConfig()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("create service config: %w", err)
|
return nil, fmt.Errorf("create service config: %w", err)
|
||||||
|
|||||||
150
client/cmd/service_json_gateway.go
Normal file
150
client/cmd/service_json_gateway.go
Normal file
@@ -0,0 +1,150 @@
|
|||||||
|
//go:build !ios && !android
|
||||||
|
|
||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"google.golang.org/grpc"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/daemonaddr"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||||
|
"github.com/netbirdio/netbird/client/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
// jsonPeerIdentity is the context key under which the connecting HTTP client's
|
||||||
|
// identity is stashed for the lifetime of its connection.
|
||||||
|
type jsonPeerIdentity struct{}
|
||||||
|
|
||||||
|
// jsonPeerIdentityValue pairs the identity with whether it could be read at
|
||||||
|
// all, so an unreadable identity is forwarded as "unknown" rather than omitted.
|
||||||
|
type jsonPeerIdentityValue struct {
|
||||||
|
id ipcauth.Identity
|
||||||
|
known bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// jsonConnContext reads the identity of the client connecting to the JSON
|
||||||
|
// socket and stashes it on the connection's context. The gateway re-dials the
|
||||||
|
// daemon in-process, so the daemon would otherwise see every JSON request as
|
||||||
|
// coming from the daemon itself.
|
||||||
|
func jsonConnContext(ctx context.Context, c net.Conn) context.Context {
|
||||||
|
value := jsonPeerIdentityValue{}
|
||||||
|
id, err := ipcauth.ConnIdentity(c)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("json gateway: cannot read HTTP client identity, privileged operations will be denied for this connection: %v", err)
|
||||||
|
} else {
|
||||||
|
value.id = id
|
||||||
|
value.known = true
|
||||||
|
}
|
||||||
|
return context.WithValue(ctx, jsonPeerIdentity{}, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// forwardIdentity stamps the HTTP client's identity onto every call the gateway
|
||||||
|
// makes to the daemon.
|
||||||
|
//
|
||||||
|
// It is an interceptor on the gateway's client connection rather than a
|
||||||
|
// runtime.WithMetadata annotator because grpc-gateway skips annotators when no
|
||||||
|
// request header maps to metadata, which an HTTP/1.0 request with no Host header
|
||||||
|
// over a unix socket achieves. The daemon would then receive no marker, see its own
|
||||||
|
// identity as the transport peer, and authorize the request as the daemon itself.
|
||||||
|
// An interceptor runs for every RPC whatever the request looked like.
|
||||||
|
func forwardIdentity(ctx context.Context) context.Context {
|
||||||
|
value, ok := ctx.Value(jsonPeerIdentity{}).(jsonPeerIdentityValue)
|
||||||
|
if !ok {
|
||||||
|
// No ConnContext ran for this request, so forward an unknown identity:
|
||||||
|
// the daemon must not mistake its own identity for the client's.
|
||||||
|
return ipcauth.WithForwardedIdentity(ctx, ipcauth.Identity{}, false)
|
||||||
|
}
|
||||||
|
return ipcauth.WithForwardedIdentity(ctx, value.id, value.known)
|
||||||
|
}
|
||||||
|
|
||||||
|
func forwardIdentityUnary(ctx context.Context, method string, req, reply any, cc *grpc.ClientConn, invoker grpc.UnaryInvoker, opts ...grpc.CallOption) error {
|
||||||
|
return invoker(forwardIdentity(ctx), method, req, reply, cc, opts...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func forwardIdentityStream(ctx context.Context, desc *grpc.StreamDesc, cc *grpc.ClientConn, method string, streamer grpc.Streamer, opts ...grpc.CallOption) (grpc.ClientStream, error) {
|
||||||
|
return streamer(forwardIdentity(ctx), desc, cc, method, opts...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// reservedHeaderWarning limits the dropped-header warning to the first occurrence.
|
||||||
|
var reservedHeaderWarning sync.Once
|
||||||
|
|
||||||
|
// jsonIncomingHeaderMatcher keeps an HTTP client from supplying the metadata the
|
||||||
|
// gateway uses to forward its identity. grpc-gateway turns "Grpc-Metadata-<key>"
|
||||||
|
// headers into gRPC metadata and joins them ahead of what its annotators add, so
|
||||||
|
// without this filter a JSON client could send its own x-netbird-fwd-uid and the
|
||||||
|
// daemon would authorize that instead of the client's real identity.
|
||||||
|
func jsonIncomingHeaderMatcher(key string) (string, bool) {
|
||||||
|
mapped, ok := runtime.DefaultHeaderMatcher(key)
|
||||||
|
if !ok {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
if ipcauth.IsReservedForwardKey(mapped) {
|
||||||
|
// Warn once: any client can send these on every request, so warning each
|
||||||
|
// time hands it a way to fill the log. The rest are debug-level.
|
||||||
|
reservedHeaderWarning.Do(func() {
|
||||||
|
log.Warnf("json gateway: dropping reserved header %q from a request: only the gateway may set the caller's identity", key)
|
||||||
|
})
|
||||||
|
log.Debugf("json gateway: dropping reserved header %q", key)
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
return mapped, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint string) error {
|
||||||
|
if jsonListener.network == "tcp" {
|
||||||
|
log.Warnf("daemon JSON socket is listening on TCP (%s): callers carry no verifiable identity over TCP, "+
|
||||||
|
"so privileged operations will be denied for JSON clients", jsonListener.address)
|
||||||
|
}
|
||||||
|
|
||||||
|
mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher))
|
||||||
|
|
||||||
|
// grpc.NewClient does not connect until the first request, so registering
|
||||||
|
// the handler here cannot block daemon startup.
|
||||||
|
target, opts := daemonaddr.DialTarget(daemonEndpoint)
|
||||||
|
opts = append(opts,
|
||||||
|
grpc.WithChainUnaryInterceptor(forwardIdentityUnary),
|
||||||
|
grpc.WithChainStreamInterceptor(forwardIdentityStream),
|
||||||
|
)
|
||||||
|
conn, err := grpc.NewClient(target, opts...)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("create daemon client for JSON gateway: %w", err)
|
||||||
|
}
|
||||||
|
if err := proto.RegisterDaemonServiceHandler(p.ctx, mux, conn); err != nil {
|
||||||
|
if cerr := conn.Close(); cerr != nil {
|
||||||
|
log.Debugf("close daemon client after failed JSON gateway registration: %v", cerr)
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
jsonServer := &http.Server{
|
||||||
|
Handler: mux,
|
||||||
|
ReadHeaderTimeout: 5 * time.Second,
|
||||||
|
BaseContext: func(net.Listener) context.Context {
|
||||||
|
return p.ctx
|
||||||
|
},
|
||||||
|
ConnContext: jsonConnContext,
|
||||||
|
}
|
||||||
|
|
||||||
|
p.jsonServMu.Lock()
|
||||||
|
p.jsonServ = jsonServer
|
||||||
|
p.jsonClient = conn
|
||||||
|
p.jsonServMu.Unlock()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
log.Printf("started daemon JSON server: %v", jsonListener.address)
|
||||||
|
if err := jsonServer.Serve(jsonListener.Listener); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||||
|
log.Errorf("failed to serve daemon JSON requests: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
261
client/cmd/service_json_gateway_test.go
Normal file
261
client/cmd/service_json_gateway_test.go
Normal file
@@ -0,0 +1,261 @@
|
|||||||
|
//go:build !windows && !ios && !android
|
||||||
|
|
||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||||
|
"google.golang.org/grpc/credentials"
|
||||||
|
"google.golang.org/grpc/metadata"
|
||||||
|
"google.golang.org/grpc/peer"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The JSON gateway runs inside the daemon and re-dials it locally, so every JSON
|
||||||
|
// request reaches a handler with the daemon's own identity as the transport peer.
|
||||||
|
// The gateway therefore forwards its HTTP client's identity as metadata, and the
|
||||||
|
// daemon authorizes that instead of itself. These tests drive the real wiring
|
||||||
|
// (jsonConnContext, forwardIdentity, jsonIncomingHeaderMatcher) and check the
|
||||||
|
// identity a handler would end up authorizing.
|
||||||
|
|
||||||
|
// daemonSideCtx is what a handler sees for a gateway-relayed call. The transport
|
||||||
|
// peer must be this process's own identity: the gateway is the daemon, so the two
|
||||||
|
// cannot differ, and hardcoding root here instead would describe a state that
|
||||||
|
// never occurs.
|
||||||
|
func daemonSideCtx(t *testing.T, md metadata.MD) context.Context {
|
||||||
|
t.Helper()
|
||||||
|
self, err := ipcauth.CurrentProcessIdentity()
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("cannot read this process's identity: %v", err)
|
||||||
|
}
|
||||||
|
ctx := peer.NewContext(context.Background(), &peer.Peer{
|
||||||
|
AuthInfo: ipcauth.AuthInfo{
|
||||||
|
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
|
||||||
|
Identity: self,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
return metadata.NewIncomingContext(ctx, md)
|
||||||
|
}
|
||||||
|
|
||||||
|
// gatewayMetadata reproduces what the daemon receives for a JSON request: the
|
||||||
|
// mux annotates the context from the request's headers, then the interceptor on the
|
||||||
|
// gateway's client connection stamps the caller's identity. The order matters,
|
||||||
|
// since the interceptor must win over anything a header put there.
|
||||||
|
func gatewayMetadata(t *testing.T, req *http.Request, ctx context.Context) metadata.MD {
|
||||||
|
t.Helper()
|
||||||
|
mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher))
|
||||||
|
annotated, err := runtime.AnnotateContext(ctx, mux, req,
|
||||||
|
"/daemon.DaemonService/SetConfig",
|
||||||
|
runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("annotate: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
md, ok := metadata.FromOutgoingContext(forwardIdentity(annotated))
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("the interceptor produced no metadata")
|
||||||
|
}
|
||||||
|
return md
|
||||||
|
}
|
||||||
|
|
||||||
|
// clientCtx is the connection context jsonConnContext would have produced for an
|
||||||
|
// HTTP client whose identity the gateway could read.
|
||||||
|
func clientCtx(id ipcauth.Identity, known bool) context.Context {
|
||||||
|
return context.WithValue(context.Background(), jsonPeerIdentity{},
|
||||||
|
jsonPeerIdentityValue{id: id, known: known})
|
||||||
|
}
|
||||||
|
|
||||||
|
// An HTTP client must not be able to name its own identity. grpc-gateway turns
|
||||||
|
// Grpc-Metadata-<key> headers into gRPC metadata, so without the header filter and
|
||||||
|
// the interceptor overwriting the reserved keys, this request would authorize as
|
||||||
|
// uid 0.
|
||||||
|
func TestJSONGateway_ForgedIdentityHeaderIsDropped(t *testing.T) {
|
||||||
|
req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Uid", "0")
|
||||||
|
req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Gid", "0")
|
||||||
|
req.Header.Set("Grpc-Metadata-X-Netbird-Fwd", "1")
|
||||||
|
req.Header.Set("Grpc-Metadata-X-Netbird-Fwd-Sid", "S-1-5-18")
|
||||||
|
|
||||||
|
caller := ipcauth.Identity{UID: 31000, GID: 31000}
|
||||||
|
md := gatewayMetadata(t, req, clientCtx(caller, true))
|
||||||
|
|
||||||
|
id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md))
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("the forwarded identity should be usable")
|
||||||
|
}
|
||||||
|
if id.IsPrivileged() {
|
||||||
|
t.Errorf("forged header was believed: authorized as %v", id)
|
||||||
|
}
|
||||||
|
if id.UID != caller.UID {
|
||||||
|
t.Errorf("authorized as uid %d, want the real client %d", id.UID, caller.UID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A request with no headers at all (HTTP/1.0 needs no Host, and a unix socket
|
||||||
|
// yields no host:port) makes grpc-gateway produce no metadata whatsoever and skip
|
||||||
|
// its annotators: "if len(pairs) == 0 { return ctx, nil, nil }" in
|
||||||
|
// runtime/context.go. That is why the identity is stamped by an interceptor
|
||||||
|
// instead. This is the case that previously reached the gate as the daemon itself.
|
||||||
|
func TestJSONGateway_HeaderlessRequestIsStillMarkedForwarded(t *testing.T) {
|
||||||
|
req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
req.Header = http.Header{}
|
||||||
|
req.Host = ""
|
||||||
|
|
||||||
|
caller := ipcauth.Identity{UID: 31000, GID: 31000}
|
||||||
|
ctx := clientCtx(caller, true)
|
||||||
|
|
||||||
|
// Pin the skip path itself: if grpc-gateway ever produced a pair here, this
|
||||||
|
// test would still pass below while no longer covering what it was written for.
|
||||||
|
mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher))
|
||||||
|
annotated, err := runtime.AnnotateContext(ctx, mux, req,
|
||||||
|
"/daemon.DaemonService/SetConfig",
|
||||||
|
runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("annotate: %v", err)
|
||||||
|
}
|
||||||
|
if md, ok := metadata.FromOutgoingContext(annotated); ok {
|
||||||
|
t.Fatalf("grpc-gateway produced metadata %v for a headerless request; "+
|
||||||
|
"this test no longer covers the annotator-skip path", md)
|
||||||
|
}
|
||||||
|
|
||||||
|
md := gatewayMetadata(t, req, ctx)
|
||||||
|
|
||||||
|
id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md))
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("the forwarded identity should be usable")
|
||||||
|
}
|
||||||
|
if id.UID != caller.UID || id.IsPrivileged() {
|
||||||
|
t.Errorf("authorized as %v, want the real client uid %d", id, caller.UID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// When the gateway cannot read its client's identity (a TCP JSON socket, say) it
|
||||||
|
// forwards the marker alone. The daemon must then report "unidentified" so the
|
||||||
|
// privileged operations refuse, rather than falling back to the gateway's own
|
||||||
|
// identity.
|
||||||
|
func TestJSONGateway_UnreadableClientIdentityIsUnidentified(t *testing.T) {
|
||||||
|
req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
md := gatewayMetadata(t, req, clientCtx(ipcauth.Identity{}, false))
|
||||||
|
|
||||||
|
if id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)); ok {
|
||||||
|
t.Errorf("a request with no client identity was authorized as %v", id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A request that never passed through jsonConnContext (no stashed identity) must
|
||||||
|
// also come out unidentified rather than as the daemon.
|
||||||
|
func TestJSONGateway_MissingConnContextIsUnidentified(t *testing.T) {
|
||||||
|
req, err := http.NewRequest(http.MethodPost, "http://localhost/daemon.DaemonService/SetConfig", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
md := gatewayMetadata(t, req, context.Background())
|
||||||
|
|
||||||
|
if id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, md)); ok {
|
||||||
|
t.Errorf("a request with no connection context was authorized as %v", id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// End to end over a real unix socket: the gateway reads the connecting client's
|
||||||
|
// identity from the socket itself, so a client cannot present anything else.
|
||||||
|
func TestJSONGateway_IdentityComesFromTheSocket(t *testing.T) {
|
||||||
|
mux := runtime.NewServeMux(runtime.WithIncomingHeaderMatcher(jsonIncomingHeaderMatcher))
|
||||||
|
|
||||||
|
type observed struct {
|
||||||
|
md metadata.MD
|
||||||
|
}
|
||||||
|
seen := make(chan observed, 1)
|
||||||
|
|
||||||
|
srv := &http.Server{
|
||||||
|
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ctx, err := runtime.AnnotateContext(r.Context(), mux, r,
|
||||||
|
"/daemon.DaemonService/SetConfig",
|
||||||
|
runtime.WithHTTPPathPattern("/daemon.DaemonService/SetConfig"))
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("annotate: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
md, _ := metadata.FromOutgoingContext(forwardIdentity(ctx))
|
||||||
|
seen <- observed{md: md}
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
}),
|
||||||
|
ReadHeaderTimeout: 5 * time.Second,
|
||||||
|
ConnContext: jsonConnContext,
|
||||||
|
}
|
||||||
|
|
||||||
|
sock := filepath.Join(t.TempDir(), "http.sock")
|
||||||
|
ln, err := net.Listen("unix", sock)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := srv.Close(); err != nil {
|
||||||
|
t.Logf("close server: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
go func() {
|
||||||
|
if err := srv.Serve(ln); err != nil && err != http.ErrServerClosed {
|
||||||
|
t.Logf("serve: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
conn, err := net.Dial("unix", sock)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := conn.Close(); err != nil {
|
||||||
|
t.Logf("close conn: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
// Forge the identity headers on the wire as well.
|
||||||
|
request := "POST /daemon.DaemonService/SetConfig HTTP/1.1\r\n" +
|
||||||
|
"Host: localhost\r\n" +
|
||||||
|
"Grpc-Metadata-X-Netbird-Fwd: 1\r\n" +
|
||||||
|
"Grpc-Metadata-X-Netbird-Fwd-Uid: 0\r\n" +
|
||||||
|
"Content-Length: 0\r\n\r\n"
|
||||||
|
if _, err := conn.Write([]byte(request)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case got := <-seen:
|
||||||
|
self, err := ipcauth.CurrentProcessIdentity()
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("cannot read this process's identity: %v", err)
|
||||||
|
}
|
||||||
|
// The socket peer is this test process, so that is the identity the
|
||||||
|
// gateway must forward, not the uid 0 the request asked for.
|
||||||
|
if uids := got.md.Get("x-netbird-fwd-uid"); len(uids) != 1 {
|
||||||
|
t.Fatalf("x-netbird-fwd-uid = %v, want exactly the gateway's own value", uids)
|
||||||
|
}
|
||||||
|
id, ok := ipcauth.CallerIdentity(daemonSideCtx(t, got.md))
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("the forwarded identity should be usable")
|
||||||
|
}
|
||||||
|
if id.UID != self.UID {
|
||||||
|
t.Errorf("authorized as uid %d, want the socket peer %d", id.UID, self.UID)
|
||||||
|
}
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("the gateway never handled the request")
|
||||||
|
}
|
||||||
|
}
|
||||||
176
client/cmd/service_json_socket_test.go
Normal file
176
client/cmd/service_json_socket_test.go
Normal file
@@ -0,0 +1,176 @@
|
|||||||
|
//go:build !ios && !android
|
||||||
|
|
||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
"github.com/spf13/pflag"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func preserveJSONSocketTestState(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
origJSONSocket := jsonSocket
|
||||||
|
origEnableJSONSocket := enableJSONSocket
|
||||||
|
origChanged := map[string]bool{}
|
||||||
|
serviceCmd.PersistentFlags().VisitAll(func(flag *pflag.Flag) {
|
||||||
|
origChanged[flag.Name] = flag.Changed
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
jsonSocket = origJSONSocket
|
||||||
|
enableJSONSocket = origEnableJSONSocket
|
||||||
|
serviceCmd.PersistentFlags().VisitAll(func(flag *pflag.Flag) {
|
||||||
|
flag.Changed = origChanged[flag.Name]
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONSocketFlagsArePositiveEnableOnly(t *testing.T) {
|
||||||
|
assert.NotNil(t, serviceCmd.PersistentFlags().Lookup("enable-json-socket"))
|
||||||
|
assert.NotNil(t, serviceCmd.PersistentFlags().Lookup("json-socket"))
|
||||||
|
assert.Nil(t, serviceCmd.PersistentFlags().Lookup("disable-json-socket"))
|
||||||
|
assert.Equal(t, "false", serviceCmd.PersistentFlags().Lookup("enable-json-socket").DefValue)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildServiceArgumentsDefaultDisablesJSONSocket(t *testing.T) {
|
||||||
|
preserveJSONSocketTestState(t)
|
||||||
|
|
||||||
|
enableJSONSocket = false
|
||||||
|
jsonSocket = "tcp://127.0.0.1:8080"
|
||||||
|
|
||||||
|
args := buildServiceArguments()
|
||||||
|
|
||||||
|
assert.NotContains(t, args, "--enable-json-socket")
|
||||||
|
assert.NotContains(t, args, "--json-socket")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildServiceArgumentsIncludesJSONSocketWhenEnabled(t *testing.T) {
|
||||||
|
preserveJSONSocketTestState(t)
|
||||||
|
|
||||||
|
enableJSONSocket = true
|
||||||
|
jsonSocket = "tcp://127.0.0.1:8080"
|
||||||
|
|
||||||
|
args := buildServiceArguments()
|
||||||
|
|
||||||
|
enableIndex := indexOfArg(args, "--enable-json-socket")
|
||||||
|
jsonIndex := indexOfArg(args, "--json-socket")
|
||||||
|
require.NotEqual(t, -1, enableIndex)
|
||||||
|
require.NotEqual(t, -1, jsonIndex)
|
||||||
|
require.Less(t, enableIndex, jsonIndex)
|
||||||
|
require.Less(t, jsonIndex+1, len(args))
|
||||||
|
assert.Equal(t, "tcp://127.0.0.1:8080", args[jsonIndex+1])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONSocketWithoutEnableValidation(t *testing.T) {
|
||||||
|
preserveJSONSocketTestState(t)
|
||||||
|
|
||||||
|
enableJSONSocket = false
|
||||||
|
require.NoError(t, serviceCmd.PersistentFlags().Set("json-socket", "tcp://127.0.0.1:8080"))
|
||||||
|
|
||||||
|
err := validateJSONSocketFlags()
|
||||||
|
|
||||||
|
require.Error(t, err)
|
||||||
|
assert.Contains(t, err.Error(), "--enable-json-socket")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONSocketWithEnableValidation(t *testing.T) {
|
||||||
|
preserveJSONSocketTestState(t)
|
||||||
|
|
||||||
|
require.NoError(t, serviceCmd.PersistentFlags().Set("enable-json-socket", "true"))
|
||||||
|
require.NoError(t, serviceCmd.PersistentFlags().Set("json-socket", "tcp://127.0.0.1:8080"))
|
||||||
|
|
||||||
|
assert.NoError(t, validateJSONSocketFlags())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestJSONSocketServiceParamsPersistEnableAndAddress(t *testing.T) {
|
||||||
|
preserveJSONSocketTestState(t)
|
||||||
|
serviceCmd.PersistentFlags().VisitAll(func(flag *pflag.Flag) {
|
||||||
|
flag.Changed = false
|
||||||
|
})
|
||||||
|
|
||||||
|
enableJSONSocket = true
|
||||||
|
jsonSocket = "tcp://127.0.0.1:8080"
|
||||||
|
|
||||||
|
params := currentServiceParams()
|
||||||
|
require.True(t, params.EnableJSONSocket)
|
||||||
|
require.Equal(t, "tcp://127.0.0.1:8080", params.JSONSocket)
|
||||||
|
|
||||||
|
enableJSONSocket = false
|
||||||
|
jsonSocket = defaultJSONSocket
|
||||||
|
applyServiceParams(testServiceEnvCommand(), params)
|
||||||
|
|
||||||
|
assert.True(t, enableJSONSocket)
|
||||||
|
assert.Equal(t, "tcp://127.0.0.1:8080", jsonSocket)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoveStaleUnixSocketDoesNotRemoveRegularFile(t *testing.T) {
|
||||||
|
path := filepath.Join(t.TempDir(), "netbird-http.sock")
|
||||||
|
require.NoError(t, os.WriteFile(path, []byte("not a socket"), 0600))
|
||||||
|
|
||||||
|
removeStaleUnixSocket(path)
|
||||||
|
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, []byte("not a socket"), data)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoveStaleUnixSocketRemovesSocket(t *testing.T) {
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
t.Skip("unix sockets are not available on Windows")
|
||||||
|
}
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "netbird-http.sock")
|
||||||
|
addr := &net.UnixAddr{Name: path, Net: "unix"}
|
||||||
|
listener, err := net.ListenUnix("unix", addr)
|
||||||
|
require.NoError(t, err)
|
||||||
|
listener.SetUnlinkOnClose(false)
|
||||||
|
require.NoError(t, listener.Close())
|
||||||
|
|
||||||
|
_, err = os.Lstat(path)
|
||||||
|
require.NoError(t, err, "test setup must leave a stale Unix socket path")
|
||||||
|
|
||||||
|
removeStaleUnixSocket(path)
|
||||||
|
|
||||||
|
_, err = os.Lstat(path)
|
||||||
|
assert.True(t, os.IsNotExist(err), "expected stale Unix socket to be removed, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRemoveStaleUnixSocketDoesNotRemoveLiveSocket(t *testing.T) {
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
t.Skip("unix sockets are not available on Windows")
|
||||||
|
}
|
||||||
|
|
||||||
|
path := filepath.Join(t.TempDir(), "netbird-http.sock")
|
||||||
|
listener, err := net.Listen("unix", path)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer listener.Close()
|
||||||
|
|
||||||
|
removeStaleUnixSocket(path)
|
||||||
|
|
||||||
|
_, err = os.Lstat(path)
|
||||||
|
assert.NoError(t, err, "expected live Unix socket to be preserved")
|
||||||
|
}
|
||||||
|
|
||||||
|
func testServiceEnvCommand() *cobra.Command {
|
||||||
|
cmd := &cobra.Command{}
|
||||||
|
cmd.Flags().StringSlice("service-env", nil, "")
|
||||||
|
return cmd
|
||||||
|
}
|
||||||
|
|
||||||
|
func indexOfArg(args []string, arg string) int {
|
||||||
|
for i, candidate := range args {
|
||||||
|
if candidate == arg {
|
||||||
|
return i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/configs"
|
"github.com/netbirdio/netbird/client/configs"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/daemonaddr"
|
||||||
"github.com/netbirdio/netbird/util"
|
"github.com/netbirdio/netbird/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -23,6 +24,7 @@ const serviceParamsFile = "service.json"
|
|||||||
type serviceParams struct {
|
type serviceParams struct {
|
||||||
LogLevel string `json:"log_level"`
|
LogLevel string `json:"log_level"`
|
||||||
DaemonAddr string `json:"daemon_addr"`
|
DaemonAddr string `json:"daemon_addr"`
|
||||||
|
JSONSocket string `json:"json_socket"`
|
||||||
ManagementURL string `json:"management_url,omitempty"`
|
ManagementURL string `json:"management_url,omitempty"`
|
||||||
ConfigPath string `json:"config_path,omitempty"`
|
ConfigPath string `json:"config_path,omitempty"`
|
||||||
LogFiles []string `json:"log_files,omitempty"`
|
LogFiles []string `json:"log_files,omitempty"`
|
||||||
@@ -30,6 +32,7 @@ type serviceParams struct {
|
|||||||
DisableUpdateSettings bool `json:"disable_update_settings,omitempty"`
|
DisableUpdateSettings bool `json:"disable_update_settings,omitempty"`
|
||||||
EnableCapture bool `json:"enable_capture,omitempty"`
|
EnableCapture bool `json:"enable_capture,omitempty"`
|
||||||
DisableNetworks bool `json:"disable_networks,omitempty"`
|
DisableNetworks bool `json:"disable_networks,omitempty"`
|
||||||
|
EnableJSONSocket bool `json:"enable_json_socket,omitempty"`
|
||||||
ServiceEnvVars map[string]string `json:"service_env_vars,omitempty"`
|
ServiceEnvVars map[string]string `json:"service_env_vars,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -75,6 +78,7 @@ func currentServiceParams() *serviceParams {
|
|||||||
params := &serviceParams{
|
params := &serviceParams{
|
||||||
LogLevel: logLevel,
|
LogLevel: logLevel,
|
||||||
DaemonAddr: daemonAddr,
|
DaemonAddr: daemonAddr,
|
||||||
|
JSONSocket: jsonSocket,
|
||||||
ManagementURL: managementURL,
|
ManagementURL: managementURL,
|
||||||
ConfigPath: configPath,
|
ConfigPath: configPath,
|
||||||
LogFiles: logFiles,
|
LogFiles: logFiles,
|
||||||
@@ -82,6 +86,7 @@ func currentServiceParams() *serviceParams {
|
|||||||
DisableUpdateSettings: updateSettingsDisabled,
|
DisableUpdateSettings: updateSettingsDisabled,
|
||||||
EnableCapture: captureEnabled,
|
EnableCapture: captureEnabled,
|
||||||
DisableNetworks: networksDisabled,
|
DisableNetworks: networksDisabled,
|
||||||
|
EnableJSONSocket: enableJSONSocket,
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(serviceEnvVars) > 0 {
|
if len(serviceEnvVars) > 0 {
|
||||||
@@ -113,15 +118,29 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// For fields with non-empty defaults (log-level, daemon-addr), keep the
|
// For fields with non-empty defaults, keep the != "" guard so that an older
|
||||||
// != "" guard so that an older service.json missing the field doesn't
|
// service.json missing the field doesn't clobber the default with an empty string.
|
||||||
// clobber the default with an empty string.
|
|
||||||
if !rootCmd.PersistentFlags().Changed("log-level") && params.LogLevel != "" {
|
if !rootCmd.PersistentFlags().Changed("log-level") && params.LogLevel != "" {
|
||||||
logLevel = params.LogLevel
|
logLevel = params.LogLevel
|
||||||
}
|
}
|
||||||
|
|
||||||
if !rootCmd.PersistentFlags().Changed("daemon-addr") && params.DaemonAddr != "" {
|
if !rootCmd.PersistentFlags().Changed("daemon-addr") && params.DaemonAddr != "" {
|
||||||
daemonAddr = params.DaemonAddr
|
daemonAddr = params.DaemonAddr
|
||||||
|
// An install that predates named-pipe support has the loopback TCP
|
||||||
|
// address saved. Callers carry no identity over TCP, so move it to the
|
||||||
|
// pipe instead of restoring a socket the daemon cannot authorize on.
|
||||||
|
if migrated, ok := daemonaddr.MigrateLegacy(daemonAddr); ok {
|
||||||
|
cmd.Printf("Moving the saved daemon address from %s to %s so the daemon can identify its callers\n", daemonAddr, migrated)
|
||||||
|
daemonAddr = migrated
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !serviceCmd.PersistentFlags().Changed("json-socket") && params.JSONSocket != "" {
|
||||||
|
jsonSocket = params.JSONSocket
|
||||||
|
}
|
||||||
|
|
||||||
|
if !serviceCmd.PersistentFlags().Changed("enable-json-socket") {
|
||||||
|
enableJSONSocket = params.EnableJSONSocket
|
||||||
}
|
}
|
||||||
|
|
||||||
// For optional fields where empty means "use default", always apply so
|
// For optional fields where empty means "use default", always apply so
|
||||||
|
|||||||
@@ -41,6 +41,8 @@ func TestSaveAndLoadServiceParams(t *testing.T) {
|
|||||||
params := &serviceParams{
|
params := &serviceParams{
|
||||||
LogLevel: "debug",
|
LogLevel: "debug",
|
||||||
DaemonAddr: "unix:///var/run/netbird.sock",
|
DaemonAddr: "unix:///var/run/netbird.sock",
|
||||||
|
JSONSocket: "tcp://127.0.0.1:8080",
|
||||||
|
EnableJSONSocket: true,
|
||||||
ManagementURL: "https://my.server.com",
|
ManagementURL: "https://my.server.com",
|
||||||
ConfigPath: "/etc/netbird/config.json",
|
ConfigPath: "/etc/netbird/config.json",
|
||||||
LogFiles: []string{"/var/log/netbird/client.log", "console"},
|
LogFiles: []string{"/var/log/netbird/client.log", "console"},
|
||||||
@@ -63,6 +65,8 @@ func TestSaveAndLoadServiceParams(t *testing.T) {
|
|||||||
|
|
||||||
assert.Equal(t, params.LogLevel, loaded.LogLevel)
|
assert.Equal(t, params.LogLevel, loaded.LogLevel)
|
||||||
assert.Equal(t, params.DaemonAddr, loaded.DaemonAddr)
|
assert.Equal(t, params.DaemonAddr, loaded.DaemonAddr)
|
||||||
|
assert.Equal(t, params.JSONSocket, loaded.JSONSocket)
|
||||||
|
assert.Equal(t, params.EnableJSONSocket, loaded.EnableJSONSocket)
|
||||||
assert.Equal(t, params.ManagementURL, loaded.ManagementURL)
|
assert.Equal(t, params.ManagementURL, loaded.ManagementURL)
|
||||||
assert.Equal(t, params.ConfigPath, loaded.ConfigPath)
|
assert.Equal(t, params.ConfigPath, loaded.ConfigPath)
|
||||||
assert.Equal(t, params.LogFiles, loaded.LogFiles)
|
assert.Equal(t, params.LogFiles, loaded.LogFiles)
|
||||||
@@ -101,6 +105,8 @@ func TestLoadServiceParams_InvalidJSON(t *testing.T) {
|
|||||||
func TestCurrentServiceParams(t *testing.T) {
|
func TestCurrentServiceParams(t *testing.T) {
|
||||||
origLogLevel := logLevel
|
origLogLevel := logLevel
|
||||||
origDaemonAddr := daemonAddr
|
origDaemonAddr := daemonAddr
|
||||||
|
origJSONSocket := jsonSocket
|
||||||
|
origEnableJSONSocket := enableJSONSocket
|
||||||
origManagementURL := managementURL
|
origManagementURL := managementURL
|
||||||
origConfigPath := configPath
|
origConfigPath := configPath
|
||||||
origLogFiles := logFiles
|
origLogFiles := logFiles
|
||||||
@@ -110,6 +116,8 @@ func TestCurrentServiceParams(t *testing.T) {
|
|||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
logLevel = origLogLevel
|
logLevel = origLogLevel
|
||||||
daemonAddr = origDaemonAddr
|
daemonAddr = origDaemonAddr
|
||||||
|
jsonSocket = origJSONSocket
|
||||||
|
enableJSONSocket = origEnableJSONSocket
|
||||||
managementURL = origManagementURL
|
managementURL = origManagementURL
|
||||||
configPath = origConfigPath
|
configPath = origConfigPath
|
||||||
logFiles = origLogFiles
|
logFiles = origLogFiles
|
||||||
@@ -120,6 +128,8 @@ func TestCurrentServiceParams(t *testing.T) {
|
|||||||
|
|
||||||
logLevel = "trace"
|
logLevel = "trace"
|
||||||
daemonAddr = "tcp://127.0.0.1:9999"
|
daemonAddr = "tcp://127.0.0.1:9999"
|
||||||
|
jsonSocket = "tcp://127.0.0.1:8080"
|
||||||
|
enableJSONSocket = true
|
||||||
managementURL = "https://mgmt.example.com"
|
managementURL = "https://mgmt.example.com"
|
||||||
configPath = "/tmp/test-config.json"
|
configPath = "/tmp/test-config.json"
|
||||||
logFiles = []string{"/tmp/test.log"}
|
logFiles = []string{"/tmp/test.log"}
|
||||||
@@ -131,6 +141,8 @@ func TestCurrentServiceParams(t *testing.T) {
|
|||||||
|
|
||||||
assert.Equal(t, "trace", params.LogLevel)
|
assert.Equal(t, "trace", params.LogLevel)
|
||||||
assert.Equal(t, "tcp://127.0.0.1:9999", params.DaemonAddr)
|
assert.Equal(t, "tcp://127.0.0.1:9999", params.DaemonAddr)
|
||||||
|
assert.Equal(t, "tcp://127.0.0.1:8080", params.JSONSocket)
|
||||||
|
assert.True(t, params.EnableJSONSocket)
|
||||||
assert.Equal(t, "https://mgmt.example.com", params.ManagementURL)
|
assert.Equal(t, "https://mgmt.example.com", params.ManagementURL)
|
||||||
assert.Equal(t, "/tmp/test-config.json", params.ConfigPath)
|
assert.Equal(t, "/tmp/test-config.json", params.ConfigPath)
|
||||||
assert.Equal(t, []string{"/tmp/test.log"}, params.LogFiles)
|
assert.Equal(t, []string{"/tmp/test.log"}, params.LogFiles)
|
||||||
@@ -142,6 +154,8 @@ func TestCurrentServiceParams(t *testing.T) {
|
|||||||
func TestApplyServiceParams_OnlyUnchangedFlags(t *testing.T) {
|
func TestApplyServiceParams_OnlyUnchangedFlags(t *testing.T) {
|
||||||
origLogLevel := logLevel
|
origLogLevel := logLevel
|
||||||
origDaemonAddr := daemonAddr
|
origDaemonAddr := daemonAddr
|
||||||
|
origJSONSocket := jsonSocket
|
||||||
|
origEnableJSONSocket := enableJSONSocket
|
||||||
origManagementURL := managementURL
|
origManagementURL := managementURL
|
||||||
origConfigPath := configPath
|
origConfigPath := configPath
|
||||||
origLogFiles := logFiles
|
origLogFiles := logFiles
|
||||||
@@ -151,6 +165,8 @@ func TestApplyServiceParams_OnlyUnchangedFlags(t *testing.T) {
|
|||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
logLevel = origLogLevel
|
logLevel = origLogLevel
|
||||||
daemonAddr = origDaemonAddr
|
daemonAddr = origDaemonAddr
|
||||||
|
jsonSocket = origJSONSocket
|
||||||
|
enableJSONSocket = origEnableJSONSocket
|
||||||
managementURL = origManagementURL
|
managementURL = origManagementURL
|
||||||
configPath = origConfigPath
|
configPath = origConfigPath
|
||||||
logFiles = origLogFiles
|
logFiles = origLogFiles
|
||||||
@@ -162,6 +178,8 @@ func TestApplyServiceParams_OnlyUnchangedFlags(t *testing.T) {
|
|||||||
// Reset all flags to defaults.
|
// Reset all flags to defaults.
|
||||||
logLevel = "info"
|
logLevel = "info"
|
||||||
daemonAddr = "unix:///var/run/netbird.sock"
|
daemonAddr = "unix:///var/run/netbird.sock"
|
||||||
|
jsonSocket = defaultJSONSocket
|
||||||
|
enableJSONSocket = false
|
||||||
managementURL = ""
|
managementURL = ""
|
||||||
configPath = "/etc/netbird/config.json"
|
configPath = "/etc/netbird/config.json"
|
||||||
logFiles = []string{"/var/log/netbird/client.log"}
|
logFiles = []string{"/var/log/netbird/client.log"}
|
||||||
@@ -184,6 +202,8 @@ func TestApplyServiceParams_OnlyUnchangedFlags(t *testing.T) {
|
|||||||
saved := &serviceParams{
|
saved := &serviceParams{
|
||||||
LogLevel: "debug",
|
LogLevel: "debug",
|
||||||
DaemonAddr: "tcp://127.0.0.1:5555",
|
DaemonAddr: "tcp://127.0.0.1:5555",
|
||||||
|
JSONSocket: "tcp://127.0.0.1:8080",
|
||||||
|
EnableJSONSocket: true,
|
||||||
ManagementURL: "https://saved.example.com",
|
ManagementURL: "https://saved.example.com",
|
||||||
ConfigPath: "/saved/config.json",
|
ConfigPath: "/saved/config.json",
|
||||||
LogFiles: []string{"/saved/client.log"},
|
LogFiles: []string{"/saved/client.log"},
|
||||||
@@ -201,6 +221,8 @@ func TestApplyServiceParams_OnlyUnchangedFlags(t *testing.T) {
|
|||||||
|
|
||||||
// All other fields were not Changed, so they should use saved values.
|
// All other fields were not Changed, so they should use saved values.
|
||||||
assert.Equal(t, "tcp://127.0.0.1:5555", daemonAddr)
|
assert.Equal(t, "tcp://127.0.0.1:5555", daemonAddr)
|
||||||
|
assert.Equal(t, "tcp://127.0.0.1:8080", jsonSocket)
|
||||||
|
assert.True(t, enableJSONSocket)
|
||||||
assert.Equal(t, "https://saved.example.com", managementURL)
|
assert.Equal(t, "https://saved.example.com", managementURL)
|
||||||
assert.Equal(t, "/saved/config.json", configPath)
|
assert.Equal(t, "/saved/config.json", configPath)
|
||||||
assert.Equal(t, []string{"/saved/client.log"}, logFiles)
|
assert.Equal(t, []string{"/saved/client.log"}, logFiles)
|
||||||
@@ -212,14 +234,17 @@ func TestApplyServiceParams_OnlyUnchangedFlags(t *testing.T) {
|
|||||||
func TestApplyServiceParams_BooleanRevertToFalse(t *testing.T) {
|
func TestApplyServiceParams_BooleanRevertToFalse(t *testing.T) {
|
||||||
origProfilesDisabled := profilesDisabled
|
origProfilesDisabled := profilesDisabled
|
||||||
origUpdateSettingsDisabled := updateSettingsDisabled
|
origUpdateSettingsDisabled := updateSettingsDisabled
|
||||||
|
origEnableJSONSocket := enableJSONSocket
|
||||||
t.Cleanup(func() {
|
t.Cleanup(func() {
|
||||||
profilesDisabled = origProfilesDisabled
|
profilesDisabled = origProfilesDisabled
|
||||||
updateSettingsDisabled = origUpdateSettingsDisabled
|
updateSettingsDisabled = origUpdateSettingsDisabled
|
||||||
|
enableJSONSocket = origEnableJSONSocket
|
||||||
})
|
})
|
||||||
|
|
||||||
// Simulate current state where booleans are true (e.g. set by previous install).
|
// Simulate current state where booleans are true (e.g. set by previous install).
|
||||||
profilesDisabled = true
|
profilesDisabled = true
|
||||||
updateSettingsDisabled = true
|
updateSettingsDisabled = true
|
||||||
|
enableJSONSocket = true
|
||||||
|
|
||||||
// Reset Changed state so flags appear unset.
|
// Reset Changed state so flags appear unset.
|
||||||
serviceCmd.PersistentFlags().VisitAll(func(f *pflag.Flag) {
|
serviceCmd.PersistentFlags().VisitAll(func(f *pflag.Flag) {
|
||||||
@@ -238,6 +263,7 @@ func TestApplyServiceParams_BooleanRevertToFalse(t *testing.T) {
|
|||||||
|
|
||||||
assert.False(t, profilesDisabled, "saved false should override current true")
|
assert.False(t, profilesDisabled, "saved false should override current true")
|
||||||
assert.False(t, updateSettingsDisabled, "saved false should override current true")
|
assert.False(t, updateSettingsDisabled, "saved false should override current true")
|
||||||
|
assert.False(t, enableJSONSocket, "saved false should override current true")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestApplyServiceParams_ClearManagementURL(t *testing.T) {
|
func TestApplyServiceParams_ClearManagementURL(t *testing.T) {
|
||||||
@@ -530,6 +556,7 @@ func fieldToGlobalVar(field string) string {
|
|||||||
m := map[string]string{
|
m := map[string]string{
|
||||||
"LogLevel": "logLevel",
|
"LogLevel": "logLevel",
|
||||||
"DaemonAddr": "daemonAddr",
|
"DaemonAddr": "daemonAddr",
|
||||||
|
"JSONSocket": "jsonSocket",
|
||||||
"ManagementURL": "managementURL",
|
"ManagementURL": "managementURL",
|
||||||
"ConfigPath": "configPath",
|
"ConfigPath": "configPath",
|
||||||
"LogFiles": "logFiles",
|
"LogFiles": "logFiles",
|
||||||
@@ -537,6 +564,7 @@ func fieldToGlobalVar(field string) string {
|
|||||||
"DisableUpdateSettings": "updateSettingsDisabled",
|
"DisableUpdateSettings": "updateSettingsDisabled",
|
||||||
"EnableCapture": "captureEnabled",
|
"EnableCapture": "captureEnabled",
|
||||||
"DisableNetworks": "networksDisabled",
|
"DisableNetworks": "networksDisabled",
|
||||||
|
"EnableJSONSocket": "enableJSONSocket",
|
||||||
"ServiceEnvVars": "serviceEnvVars",
|
"ServiceEnvVars": "serviceEnvVars",
|
||||||
}
|
}
|
||||||
if v, ok := m[field]; ok {
|
if v, ok := m[field]; ok {
|
||||||
|
|||||||
14
client/cmd/service_pipe_other.go
Normal file
14
client/cmd/service_pipe_other.go
Normal file
@@ -0,0 +1,14 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
)
|
||||||
|
|
||||||
|
// listenNamedPipe is Windows-only: no other platform serves the daemon on a
|
||||||
|
// named pipe.
|
||||||
|
func listenNamedPipe(string) (net.Listener, string, error) {
|
||||||
|
return nil, "", fmt.Errorf("named pipes are only supported on Windows")
|
||||||
|
}
|
||||||
41
client/cmd/service_pipe_windows.go
Normal file
41
client/cmd/service_pipe_windows.go
Normal file
@@ -0,0 +1,41 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"github.com/Microsoft/go-winio"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/daemonaddr"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||||
|
)
|
||||||
|
|
||||||
|
// listenNamedPipe creates the daemon control pipe and reports the path it ended
|
||||||
|
// up on. The security descriptor lets any local caller connect, as a Unix socket
|
||||||
|
// at 0666 does, and the privileged operations are authorized separately from the
|
||||||
|
// caller's token.
|
||||||
|
//
|
||||||
|
// The protected name comes first so that an unprivileged process cannot take the
|
||||||
|
// name before the service does. Creating it requires being an administrator or
|
||||||
|
// LocalSystem, so a daemon an ordinary user runs themselves, as in netstack mode,
|
||||||
|
// falls back to the plain name; clients try both and check who serves them.
|
||||||
|
func listenNamedPipe(name string) (net.Listener, string, error) {
|
||||||
|
var errs []error
|
||||||
|
for _, path := range daemonaddr.PipePaths(name) {
|
||||||
|
listener, err := winio.ListenPipe(path, &winio.PipeConfig{
|
||||||
|
SecurityDescriptor: ipcauth.DefaultPipeSDDL(),
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("not serving the daemon on %s: %v", path, err)
|
||||||
|
errs = append(errs, fmt.Errorf("%s: %w", path, err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return listener, path, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, "", errors.Join(errs...)
|
||||||
|
}
|
||||||
119
client/cmd/service_socket.go
Normal file
119
client/cmd/service_socket.go
Normal file
@@ -0,0 +1,119 @@
|
|||||||
|
//go:build !ios && !android
|
||||||
|
|
||||||
|
package cmd
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
)
|
||||||
|
|
||||||
|
type socketListener struct {
|
||||||
|
net.Listener
|
||||||
|
network string
|
||||||
|
address string
|
||||||
|
}
|
||||||
|
|
||||||
|
func listenOnAddress(addr string) (*socketListener, error) {
|
||||||
|
network, address, err := parseListenAddress(addr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if network == "npipe" {
|
||||||
|
listener, path, err := listenNamedPipe(address)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &socketListener{Listener: listener, network: network, address: path}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if network == "unix" {
|
||||||
|
removeStaleUnixSocket(address)
|
||||||
|
}
|
||||||
|
|
||||||
|
listener, err := net.Listen(network, address)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &socketListener{Listener: listener, network: network, address: address}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseListenAddress(addr string) (string, string, error) {
|
||||||
|
network, address, ok := strings.Cut(addr, "://")
|
||||||
|
if !ok || network == "" || address == "" {
|
||||||
|
return "", "", fmt.Errorf("address must be in [unix|tcp|npipe]://[path|host:port|name] format: %q", addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
switch network {
|
||||||
|
case "unix", "tcp", "npipe":
|
||||||
|
return network, address, nil
|
||||||
|
default:
|
||||||
|
return "", "", fmt.Errorf("unsupported daemon address protocol: %v", network)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func removeStaleUnixSocket(path string) {
|
||||||
|
stat, err := os.Lstat(path)
|
||||||
|
if err != nil {
|
||||||
|
if !os.IsNotExist(err) {
|
||||||
|
log.Debugf("stat socket file: %v", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if stat.Mode()&os.ModeSocket == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !isStaleUnixSocket(path) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.Remove(path); err != nil {
|
||||||
|
log.Debugf("remove socket file: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func isStaleUnixSocket(path string) bool {
|
||||||
|
conn, err := net.DialTimeout("unix", path, 100*time.Millisecond)
|
||||||
|
if err == nil {
|
||||||
|
if closeErr := conn.Close(); closeErr != nil {
|
||||||
|
log.Debugf("close unix socket probe: %v", closeErr)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if os.IsNotExist(err) || os.IsPermission(err) || os.IsTimeout(err) {
|
||||||
|
log.Debugf("not removing unix socket %s after probe error: %v", path, err)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
return errors.Is(err, syscall.ECONNREFUSED)
|
||||||
|
}
|
||||||
|
|
||||||
|
func removeStaleUnixSocketForAddress(addr string) {
|
||||||
|
network, address, err := parseListenAddress(addr)
|
||||||
|
if err != nil || network != "unix" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
removeStaleUnixSocket(address)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (l *socketListener) chmodUnixSocket(description string) error {
|
||||||
|
if l == nil || l.network != "unix" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := os.Chmod(l.address, 0666); err != nil {
|
||||||
|
return fmt.Errorf("failed setting %s permissions for %s: %w", description, l.address, err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
"google.golang.org/grpc/status"
|
"google.golang.org/grpc/status"
|
||||||
@@ -115,6 +116,11 @@ func statusFunc(cmd *cobra.Command, args []string) error {
|
|||||||
// manager only knows the active profile ID, not its display name.
|
// manager only knows the active profile ID, not its display name.
|
||||||
profName := getActiveProfileName(ctx)
|
profName := getActiveProfileName(ctx)
|
||||||
|
|
||||||
|
var sessionExpiresAt time.Time
|
||||||
|
if ts := resp.GetSessionExpiresAt(); ts.IsValid() {
|
||||||
|
sessionExpiresAt = ts.AsTime().UTC()
|
||||||
|
}
|
||||||
|
|
||||||
var outputInformationHolder = nbstatus.ConvertToStatusOutputOverview(resp.GetFullStatus(), nbstatus.ConvertOptions{
|
var outputInformationHolder = nbstatus.ConvertToStatusOutputOverview(resp.GetFullStatus(), nbstatus.ConvertOptions{
|
||||||
Anonymize: anonymizeFlag,
|
Anonymize: anonymizeFlag,
|
||||||
DaemonVersion: resp.GetDaemonVersion(),
|
DaemonVersion: resp.GetDaemonVersion(),
|
||||||
@@ -125,6 +131,7 @@ func statusFunc(cmd *cobra.Command, args []string) error {
|
|||||||
IPsFilter: ipsFilterMap,
|
IPsFilter: ipsFilterMap,
|
||||||
ConnectionTypeFilter: connectionTypeFilter,
|
ConnectionTypeFilter: connectionTypeFilter,
|
||||||
ProfileName: profName,
|
ProfileName: profName,
|
||||||
|
SessionExpiresAt: sessionExpiresAt,
|
||||||
})
|
})
|
||||||
var statusOutputString string
|
var statusOutputString string
|
||||||
switch {
|
switch {
|
||||||
|
|||||||
@@ -22,6 +22,8 @@ import (
|
|||||||
"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/proto"
|
"github.com/netbirdio/netbird/client/proto"
|
||||||
|
nbnet "github.com/netbirdio/netbird/client/net"
|
||||||
|
"github.com/netbirdio/netbird/client/server"
|
||||||
"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"
|
||||||
"github.com/netbirdio/netbird/util"
|
"github.com/netbirdio/netbird/util"
|
||||||
@@ -229,6 +231,24 @@ func runInForegroundMode(ctx context.Context, cmd *cobra.Command, activeProf *pr
|
|||||||
|
|
||||||
_, _ = profilemanager.UpdateOldManagementURL(ctx, config, configFilePath)
|
_, _ = profilemanager.UpdateOldManagementURL(ctx, config, configFilePath)
|
||||||
|
|
||||||
|
// Restore residual state left by a previous run that did not shut down
|
||||||
|
// cleanly, mirroring what the daemon does before connecting: it recovers
|
||||||
|
// DNS config (a stale resolv.conf takeover can make the management
|
||||||
|
// hostname unresolvable), firewall rules, ssh config and legacy routing.
|
||||||
|
// Route cleanup itself happens at engine start; nbnet.Init() below lets
|
||||||
|
// the management dial bypass a leftover fwmark rule until then.
|
||||||
|
// Foreground mode is particularly exposed in containers: a crashed
|
||||||
|
// container restarts inside the same (pod) network namespace, so stale
|
||||||
|
// state survives while the process does not.
|
||||||
|
if err := server.RestoreResidualState(ctx, profilemanager.NewServiceManager(configPath).GetStatePath()); err != nil {
|
||||||
|
log.Warnf("failed to restore residual state: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Enable advanced routing (as the daemon does on startup) so the
|
||||||
|
// management dial bypasses a leftover fwmark rule instead of being
|
||||||
|
// shunted into a stale routing table.
|
||||||
|
nbnet.Init()
|
||||||
|
|
||||||
err = foregroundLogin(ctx, cmd, config, providedSetupKey, activeProf.ID)
|
err = foregroundLogin(ctx, cmd, config, providedSetupKey, activeProf.ID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("foreground login failed: %v", err)
|
return fmt.Errorf("foreground login failed: %v", err)
|
||||||
@@ -305,7 +325,7 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
|
|||||||
if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unavailable {
|
if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unavailable {
|
||||||
log.Warnf("setConfig method is not available in the daemon: %s", st.Message())
|
log.Warnf("setConfig method is not available in the daemon: %s", st.Message())
|
||||||
} else {
|
} else {
|
||||||
return fmt.Errorf("call service setConfig method: %v", err)
|
return daemonCallError("call service setConfig method", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -359,7 +379,7 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ
|
|||||||
}
|
}
|
||||||
|
|
||||||
if loginErr != nil {
|
if loginErr != nil {
|
||||||
return fmt.Errorf("login failed: %v", loginErr)
|
return daemonCallError("login failed", loginErr)
|
||||||
}
|
}
|
||||||
|
|
||||||
if loginResp.NeedsSSOLogin {
|
if loginResp.NeedsSSOLogin {
|
||||||
@@ -372,7 +392,7 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ
|
|||||||
ProfileName: &profileID,
|
ProfileName: &profileID,
|
||||||
Username: &username,
|
Username: &username,
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
return fmt.Errorf("call service up method: %v", err)
|
return daemonCallError("call service up method", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -470,7 +470,7 @@ func (c *Client) Status() (peer.FullStatus, error) {
|
|||||||
if connect != nil {
|
if connect != nil {
|
||||||
engine := connect.Engine()
|
engine := connect.Engine()
|
||||||
if engine != nil {
|
if engine != nil {
|
||||||
_ = engine.RunHealthProbes(false)
|
_ = engine.RunHealthProbes(context.Background(), false)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -121,6 +121,7 @@ type Manager struct {
|
|||||||
udpTracker *conntrack.UDPTracker
|
udpTracker *conntrack.UDPTracker
|
||||||
icmpTracker *conntrack.ICMPTracker
|
icmpTracker *conntrack.ICMPTracker
|
||||||
tcpTracker *conntrack.TCPTracker
|
tcpTracker *conntrack.TCPTracker
|
||||||
|
fragments *fragmentTracker
|
||||||
forwarder atomic.Pointer[forwarder.Forwarder]
|
forwarder atomic.Pointer[forwarder.Forwarder]
|
||||||
pendingCapture atomic.Pointer[forwarder.PacketCapture]
|
pendingCapture atomic.Pointer[forwarder.PacketCapture]
|
||||||
logger *nblog.Logger
|
logger *nblog.Logger
|
||||||
@@ -183,6 +184,41 @@ func (d *decoder) decodePacket(data []byte) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// decodeTransport decodes the transport header of a first fragment (which
|
||||||
|
// gopacket leaves undecoded) into the decoder and appends its layer type to
|
||||||
|
// decoded, so the ACL pipeline can evaluate it like a normal packet. It returns
|
||||||
|
// false if the protocol is unsupported or the header is truncated.
|
||||||
|
func (d *decoder) decodeTransport(proto layers.IPProtocol, payload []byte) bool {
|
||||||
|
var l4 gopacket.DecodingLayer
|
||||||
|
var layerType gopacket.LayerType
|
||||||
|
var minLen int
|
||||||
|
switch proto {
|
||||||
|
case layers.IPProtocolTCP:
|
||||||
|
l4, layerType, minLen = &d.tcp, layers.LayerTypeTCP, 20
|
||||||
|
case layers.IPProtocolUDP:
|
||||||
|
l4, layerType, minLen = &d.udp, layers.LayerTypeUDP, 8
|
||||||
|
case layers.IPProtocolICMPv4:
|
||||||
|
l4, layerType, minLen = &d.icmp4, layers.LayerTypeICMPv4, 8
|
||||||
|
case layers.IPProtocolICMPv6:
|
||||||
|
l4, layerType, minLen = &d.icmp6, layers.LayerTypeICMPv6, 8
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reject a fragment too small to hold the full transport header before
|
||||||
|
// decoding: it can't be ACL-evaluated (tiny-fragment attack), and skipping
|
||||||
|
// the decode avoids gopacket allocating an error on the drop path.
|
||||||
|
if len(payload) < minLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := l4.DecodeFromBytes(payload, gopacket.NilDecodeFeedback); err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
d.decoded = append(d.decoded, layerType)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
// Create userspace firewall manager constructor
|
// Create userspace firewall manager constructor
|
||||||
func Create(iface common.IFaceMapper, disableServerRoutes bool, flowLogger nftypes.FlowLogger, mtu uint16) (*Manager, error) {
|
func Create(iface common.IFaceMapper, disableServerRoutes bool, flowLogger nftypes.FlowLogger, mtu uint16) (*Manager, error) {
|
||||||
return create(iface, nil, disableServerRoutes, flowLogger, mtu)
|
return create(iface, nil, disableServerRoutes, flowLogger, mtu)
|
||||||
@@ -286,6 +322,8 @@ func create(iface common.IFaceMapper, nativeFirewall firewall.Manager, disableSe
|
|||||||
if err := m.localipmanager.UpdateLocalIPs(iface); err != nil {
|
if err := m.localipmanager.UpdateLocalIPs(iface); err != nil {
|
||||||
return nil, fmt.Errorf("update local IPs: %w", err)
|
return nil, fmt.Errorf("update local IPs: %w", err)
|
||||||
}
|
}
|
||||||
|
m.fragments = newFragmentTracker(m.logger)
|
||||||
|
|
||||||
if disableConntrack {
|
if disableConntrack {
|
||||||
log.Info("conntrack is disabled")
|
log.Info("conntrack is disabled")
|
||||||
} else {
|
} else {
|
||||||
@@ -299,6 +337,7 @@ func create(iface common.IFaceMapper, nativeFirewall firewall.Manager, disableSe
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := iface.SetFilter(m); err != nil {
|
if err := iface.SetFilter(m); err != nil {
|
||||||
|
m.fragments.Close()
|
||||||
return nil, fmt.Errorf("set filter: %w", err)
|
return nil, fmt.Errorf("set filter: %w", err)
|
||||||
}
|
}
|
||||||
return m, nil
|
return m, nil
|
||||||
@@ -694,6 +733,10 @@ func (m *Manager) resetState() {
|
|||||||
m.tcpTracker.Close()
|
m.tcpTracker.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if m.fragments != nil {
|
||||||
|
m.fragments.Close()
|
||||||
|
}
|
||||||
|
|
||||||
if fwder := m.forwarder.Load(); fwder != nil {
|
if fwder := m.forwarder.Load(); fwder != nil {
|
||||||
fwder.SetCapture(nil)
|
fwder.SetCapture(nil)
|
||||||
fwder.Stop()
|
fwder.Stop()
|
||||||
@@ -1046,19 +1089,20 @@ func (m *Manager) filterInbound(packetData []byte, size int) bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// TODO: pass fragments of routed packets to forwarder
|
// gopacket does not decode the transport header of any IP fragment, so
|
||||||
|
// fragments take a dedicated path: the first fragment's header is decoded
|
||||||
|
// and ACL-evaluated here, and the remaining fragments inherit its verdict.
|
||||||
if fragment {
|
if fragment {
|
||||||
if m.logger.Enabled(nblog.LevelTrace) {
|
return m.filterInboundFragment(d, srcIP, dstIP, size)
|
||||||
if d.decoded[0] == layers.LayerTypeIPv4 {
|
|
||||||
m.logger.Trace4("packet is a fragment: src=%v dst=%v id=%v flags=%v",
|
|
||||||
srcIP, dstIP, d.ip4.Id, d.ip4.Flags)
|
|
||||||
} else {
|
|
||||||
m.logger.Trace2("packet is an IPv6 fragment: src=%v dst=%v", srcIP, dstIP)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return m.filterInboundDecoded(d, srcIP, dstIP, packetData, size)
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterInboundDecoded runs the ACL, DNAT and conntrack pipeline on a fully
|
||||||
|
// decoded (non-fragment) inbound packet. It returns true if the packet should
|
||||||
|
// be dropped.
|
||||||
|
func (m *Manager) filterInboundDecoded(d *decoder, srcIP, dstIP netip.Addr, packetData []byte, size int) bool {
|
||||||
// TODO: optimize port DNAT by caching matched rules in conntrack
|
// TODO: optimize port DNAT by caching matched rules in conntrack
|
||||||
if translated := m.translateInboundPortDNAT(packetData, d, srcIP, dstIP); translated {
|
if translated := m.translateInboundPortDNAT(packetData, d, srcIP, dstIP); translated {
|
||||||
// Re-decode after port DNAT translation to update port information
|
// Re-decode after port DNAT translation to update port information
|
||||||
@@ -1089,33 +1133,226 @@ func (m *Manager) filterInbound(packetData []byte, size int) bool {
|
|||||||
return m.handleRoutedTraffic(d, srcIP, dstIP, packetData, size)
|
return m.handleRoutedTraffic(d, srcIP, dstIP, packetData, size)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// fragmentMeta holds the reassembly identity and layout of an IP fragment,
|
||||||
|
// extracted uniformly for IPv4 and IPv6.
|
||||||
|
type fragmentMeta struct {
|
||||||
|
key fragmentKey
|
||||||
|
// offset is the fragment offset in 8-byte units (zero for the first
|
||||||
|
// fragment).
|
||||||
|
offset uint16
|
||||||
|
// moreFragments is the More Fragments bit. A first fragment with it unset is
|
||||||
|
// an IPv6 atomic fragment (a complete datagram, RFC 6946): it has no trailing
|
||||||
|
// fragments to inherit a verdict, so it must not be recorded.
|
||||||
|
moreFragments bool
|
||||||
|
proto layers.IPProtocol
|
||||||
|
// l4payload is the fragmentable payload of this fragment. For the first
|
||||||
|
// fragment it starts with the transport header.
|
||||||
|
l4payload []byte
|
||||||
|
// headerEndOctets is the first fragment's payload length in 8-byte units:
|
||||||
|
// the smallest offset a trailing fragment may start at without overlapping
|
||||||
|
// the inspected transport header.
|
||||||
|
headerEndOctets uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
// fragmentMetadata extracts the fragment identity and layout from a decoded IP
|
||||||
|
// fragment. It returns false for fragments it can't interpret (e.g. an IPv6
|
||||||
|
// fragment header shorter than 8 bytes), which are then dropped.
|
||||||
|
func fragmentMetadata(d *decoder, srcIP, dstIP netip.Addr) (fragmentMeta, bool) {
|
||||||
|
switch d.decoded[0] {
|
||||||
|
case layers.LayerTypeIPv4:
|
||||||
|
payload := d.ip4.Payload
|
||||||
|
return fragmentMeta{
|
||||||
|
key: fragmentKey{srcIP: srcIP, dstIP: dstIP, id: uint32(d.ip4.Id), proto: uint8(d.ip4.Protocol)},
|
||||||
|
offset: d.ip4.FragOffset,
|
||||||
|
moreFragments: d.ip4.Flags&layers.IPv4MoreFragments != 0,
|
||||||
|
proto: d.ip4.Protocol,
|
||||||
|
l4payload: payload,
|
||||||
|
headerEndOctets: octets(len(payload)),
|
||||||
|
}, true
|
||||||
|
|
||||||
|
case layers.LayerTypeIPv6:
|
||||||
|
// IPv6 fragment extension header: 8 bytes, followed by the fragmentable
|
||||||
|
// payload. Layout: next header (1), reserved (1), offset+flags (2), id (4).
|
||||||
|
payload := d.ip6.Payload
|
||||||
|
if len(payload) < 8 {
|
||||||
|
return fragmentMeta{}, false
|
||||||
|
}
|
||||||
|
nextHeader := layers.IPProtocol(payload[0])
|
||||||
|
offsetFlags := binary.BigEndian.Uint16(payload[2:4])
|
||||||
|
id := binary.BigEndian.Uint32(payload[4:8])
|
||||||
|
l4 := payload[8:]
|
||||||
|
return fragmentMeta{
|
||||||
|
key: fragmentKey{srcIP: srcIP, dstIP: dstIP, id: id, proto: uint8(nextHeader)},
|
||||||
|
offset: offsetFlags >> 3,
|
||||||
|
moreFragments: offsetFlags&1 != 0,
|
||||||
|
proto: nextHeader,
|
||||||
|
l4payload: l4,
|
||||||
|
headerEndOctets: octets(len(l4)),
|
||||||
|
}, true
|
||||||
|
|
||||||
|
default:
|
||||||
|
return fragmentMeta{}, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// octets rounds a byte length up to whole 8-byte units, the granularity of the
|
||||||
|
// IP fragment offset field.
|
||||||
|
func octets(nbytes int) uint16 {
|
||||||
|
return uint16((nbytes + 7) / 8)
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterInboundFragment decides the fate of an IP fragment. gopacket stops
|
||||||
|
// decoding at the network layer for every fragment, so the first fragment's
|
||||||
|
// transport header is decoded and ACL-evaluated here and its verdict recorded;
|
||||||
|
// the remaining (headerless) fragments inherit that verdict. Anything that
|
||||||
|
// cannot be tied to an allowed, non-overlapping first fragment is dropped.
|
||||||
|
func (m *Manager) filterInboundFragment(d *decoder, srcIP, dstIP netip.Addr, size int) bool {
|
||||||
|
meta, ok := fragmentMetadata(d, srcIP, dstIP)
|
||||||
|
if !ok {
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace2("dropping unsupported fragment: src=%v dst=%v", srcIP, dstIP)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if meta.offset != 0 {
|
||||||
|
return m.filterTrailingFragment(meta, srcIP, dstIP)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A new first fragment supersedes any recorded verdict for this datagram, so
|
||||||
|
// a re-sent or overlapping offset-zero fragment can't inherit the old one.
|
||||||
|
m.fragments.poison(meta.key)
|
||||||
|
|
||||||
|
// First fragment: decode its transport header so the ACL can evaluate it. A
|
||||||
|
// decode failure means the fragment is too small to hold the full transport
|
||||||
|
// header (RFC 1858 §3 tiny-fragment attack); it can't be evaluated, so drop it.
|
||||||
|
if !d.decodeTransport(meta.proto, meta.l4payload) {
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace3("dropping first fragment without full L4 header: src=%v dst=%v id=%v",
|
||||||
|
srcIP, dstIP, meta.key.id)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return m.filterFirstFragment(d, meta, srcIP, dstIP, size)
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterTrailingFragment applies a recorded first-fragment verdict to a
|
||||||
|
// non-first fragment.
|
||||||
|
func (m *Manager) filterTrailingFragment(meta fragmentMeta, srcIP, dstIP netip.Addr) bool {
|
||||||
|
switch m.fragments.verdict(meta.key, meta.offset) {
|
||||||
|
case fragmentAllow:
|
||||||
|
return false
|
||||||
|
case fragmentOverlap:
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace3("dropping overlapping fragment rewriting inspected header: src=%v dst=%v id=%v",
|
||||||
|
srcIP, dstIP, meta.key.id)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace3("dropping fragment with no allowed first fragment: src=%v dst=%v id=%v",
|
||||||
|
srcIP, dstIP, meta.key.id)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// filterFirstFragment runs the verdict part of the inbound pipeline on a first
|
||||||
|
// fragment with its transport header decoded. It mirrors filterInboundDecoded
|
||||||
|
// but skips DNAT (port rewriting on fragments is unsupported) and forwarder
|
||||||
|
// injection (fragments are left to the stack to reassemble, not forwarded).
|
||||||
|
// Allowed fragments have their verdict recorded so the datagram's trailing
|
||||||
|
// fragments inherit it.
|
||||||
|
func (m *Manager) filterFirstFragment(d *decoder, meta fragmentMeta, srcIP, dstIP netip.Addr, size int) bool {
|
||||||
|
if m.stateful && m.isValidTrackedConnection(d, srcIP, dstIP, size) {
|
||||||
|
m.recordFirstFragment(meta)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if m.localipmanager.IsLocalIP(dstIP) {
|
||||||
|
ruleID, blocked := m.peerACLsBlock(srcIP, d, nil)
|
||||||
|
if blocked {
|
||||||
|
m.storeDropFlow("Dropping local first fragment (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||||
|
d, srcIP, dstIP, ruleID, size)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
m.trackInbound(d, srcIP, dstIP, ruleID, size)
|
||||||
|
m.recordFirstFragment(meta)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if !m.routingEnabled.Load() {
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace2("Dropping routed fragment (routing disabled): src=%s dst=%s", srcIP, dstIP)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if m.nativeRouter.Load() {
|
||||||
|
m.trackInbound(d, srcIP, dstIP, nil, size)
|
||||||
|
m.recordFirstFragment(meta)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: pass fragments of routed packets to the forwarder; until then
|
||||||
|
// allowed routed fragments go to the native stack.
|
||||||
|
srcPort, dstPort := getPortsFromPacket(d)
|
||||||
|
ruleID, pass := m.routeACLsPass(srcIP, dstIP, d.decoded[1], srcPort, dstPort)
|
||||||
|
if !pass {
|
||||||
|
m.storeDropFlow("Dropping routed first fragment (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||||
|
d, srcIP, dstIP, ruleID, size)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
m.recordFirstFragment(meta)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// recordFirstFragment caches an allowed first fragment's verdict for its
|
||||||
|
// trailing fragments to inherit. Atomic fragments (no More Fragments bit) are
|
||||||
|
// complete datagrams with no trailing fragments, so they are not cached and
|
||||||
|
// cannot exhaust the verdict table.
|
||||||
|
func (m *Manager) recordFirstFragment(meta fragmentMeta) {
|
||||||
|
if !meta.moreFragments {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
m.fragments.recordAllowed(meta.key, meta.headerEndOctets)
|
||||||
|
}
|
||||||
|
|
||||||
|
// storeDropFlow logs and records a netflow drop event for an inbound packet
|
||||||
|
// denied by the ACLs. msg is the trace format taking rule id, protocol, source
|
||||||
|
// and destination.
|
||||||
|
func (m *Manager) storeDropFlow(msg string, d *decoder, srcIP, dstIP netip.Addr, ruleID []byte, size int) {
|
||||||
|
pnum := getProtocolFromPacket(d)
|
||||||
|
srcPort, dstPort := getPortsFromPacket(d)
|
||||||
|
|
||||||
|
if m.logger.Enabled(nblog.LevelTrace) {
|
||||||
|
m.logger.Trace6(msg, ruleID, pnum, srcIP, srcPort, dstIP, dstPort)
|
||||||
|
}
|
||||||
|
|
||||||
|
m.flowLogger.StoreEvent(nftypes.EventFields{
|
||||||
|
FlowID: uuid.New(),
|
||||||
|
Type: nftypes.TypeDrop,
|
||||||
|
RuleID: ruleID,
|
||||||
|
Direction: nftypes.Ingress,
|
||||||
|
Protocol: pnum,
|
||||||
|
SourceIP: srcIP,
|
||||||
|
DestIP: dstIP,
|
||||||
|
SourcePort: srcPort,
|
||||||
|
DestPort: dstPort,
|
||||||
|
// TODO: icmp type/code
|
||||||
|
RxPackets: 1,
|
||||||
|
RxBytes: uint64(size),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// handleLocalTraffic handles local traffic.
|
// handleLocalTraffic handles local traffic.
|
||||||
// If it returns true, the packet should be dropped.
|
// If it returns true, the packet should be dropped.
|
||||||
func (m *Manager) handleLocalTraffic(d *decoder, srcIP, dstIP netip.Addr, packetData []byte, size int) bool {
|
func (m *Manager) handleLocalTraffic(d *decoder, srcIP, dstIP netip.Addr, packetData []byte, size int) bool {
|
||||||
ruleID, blocked := m.peerACLsBlock(srcIP, d, packetData)
|
ruleID, blocked := m.peerACLsBlock(srcIP, d, packetData)
|
||||||
if blocked {
|
if blocked {
|
||||||
pnum := getProtocolFromPacket(d)
|
m.storeDropFlow("Dropping local packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||||
srcPort, dstPort := getPortsFromPacket(d)
|
d, srcIP, dstIP, ruleID, size)
|
||||||
|
|
||||||
if m.logger.Enabled(nblog.LevelTrace) {
|
|
||||||
m.logger.Trace6("Dropping local packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
|
||||||
ruleID, pnum, srcIP, srcPort, dstIP, dstPort)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.flowLogger.StoreEvent(nftypes.EventFields{
|
|
||||||
FlowID: uuid.New(),
|
|
||||||
Type: nftypes.TypeDrop,
|
|
||||||
RuleID: ruleID,
|
|
||||||
Direction: nftypes.Ingress,
|
|
||||||
Protocol: pnum,
|
|
||||||
SourceIP: srcIP,
|
|
||||||
DestIP: dstIP,
|
|
||||||
SourcePort: srcPort,
|
|
||||||
DestPort: dstPort,
|
|
||||||
// TODO: icmp type/code
|
|
||||||
RxPackets: 1,
|
|
||||||
RxBytes: uint64(size),
|
|
||||||
})
|
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1168,27 +1405,8 @@ func (m *Manager) handleRoutedTraffic(d *decoder, srcIP, dstIP netip.Addr, packe
|
|||||||
|
|
||||||
ruleID, pass := m.routeACLsPass(srcIP, dstIP, protoLayer, srcPort, dstPort)
|
ruleID, pass := m.routeACLsPass(srcIP, dstIP, protoLayer, srcPort, dstPort)
|
||||||
if !pass {
|
if !pass {
|
||||||
proto := getProtocolFromPacket(d)
|
m.storeDropFlow("Dropping routed packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||||
|
d, srcIP, dstIP, ruleID, size)
|
||||||
if m.logger.Enabled(nblog.LevelTrace) {
|
|
||||||
m.logger.Trace6("Dropping routed packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
|
||||||
ruleID, proto, srcIP, srcPort, dstIP, dstPort)
|
|
||||||
}
|
|
||||||
|
|
||||||
m.flowLogger.StoreEvent(nftypes.EventFields{
|
|
||||||
FlowID: uuid.New(),
|
|
||||||
Type: nftypes.TypeDrop,
|
|
||||||
RuleID: ruleID,
|
|
||||||
Direction: nftypes.Ingress,
|
|
||||||
Protocol: proto,
|
|
||||||
SourceIP: srcIP,
|
|
||||||
DestIP: dstIP,
|
|
||||||
SourcePort: srcPort,
|
|
||||||
DestPort: dstPort,
|
|
||||||
// TODO: icmp type/code
|
|
||||||
RxPackets: 1,
|
|
||||||
RxBytes: uint64(size),
|
|
||||||
})
|
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,9 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"strconv"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -31,6 +33,11 @@ const (
|
|||||||
defaultMaxInFlight = 1024
|
defaultMaxInFlight = 1024
|
||||||
iosReceiveWindow = 16384
|
iosReceiveWindow = 16384
|
||||||
iosMaxInFlight = 256
|
iosMaxInFlight = 256
|
||||||
|
|
||||||
|
// envForceTCPRACK overrides the platform default for gVisor's RACK loss
|
||||||
|
// detection. Set to a truthy value to force RACK on, or a falsy value to
|
||||||
|
// force it off, on any platform.
|
||||||
|
envForceTCPRACK = "NB_FORCE_TCP_RACK"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Forwarder struct {
|
type Forwarder struct {
|
||||||
@@ -152,6 +159,8 @@ func New(iface common.IFaceMapper, logger *nblog.Logger, flowLogger nftypes.Flow
|
|||||||
maxInFlight = iosMaxInFlight
|
maxInFlight = iosMaxInFlight
|
||||||
}
|
}
|
||||||
|
|
||||||
|
configureTCPRecovery(s)
|
||||||
|
|
||||||
tcpForwarder := tcp.NewForwarder(s, receiveWindow, maxInFlight, f.handleTCP)
|
tcpForwarder := tcp.NewForwarder(s, receiveWindow, maxInFlight, f.handleTCP)
|
||||||
s.SetTransportProtocolHandler(tcp.ProtocolNumber, tcpForwarder.HandlePacket)
|
s.SetTransportProtocolHandler(tcp.ProtocolNumber, tcpForwarder.HandlePacket)
|
||||||
|
|
||||||
@@ -466,3 +475,31 @@ func probeRawICMP(network, addr string, logger *nblog.Logger) bool {
|
|||||||
logger.Debug1("forwarder: raw %s socket access available", network)
|
logger.Debug1("forwarder: raw %s socket access available", network)
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// configureTCPRecovery disables gVisor's RACK loss detection on Windows, where
|
||||||
|
// it interacts poorly with the host and collapses throughput on routed TCP
|
||||||
|
// connections (gVisor issue #9778). Other platforms keep the default. The
|
||||||
|
// EnvForceTCPRACK environment variable overrides the platform default.
|
||||||
|
func configureTCPRecovery(s *stack.Stack) {
|
||||||
|
disableRACK := runtime.GOOS == "windows"
|
||||||
|
|
||||||
|
if val := os.Getenv(envForceTCPRACK); val != "" {
|
||||||
|
force, err := strconv.ParseBool(val)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("parse %s: %v", envForceTCPRACK, err)
|
||||||
|
} else {
|
||||||
|
disableRACK = !force
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !disableRACK {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
opt := tcpip.TCPRecovery(0)
|
||||||
|
if err := s.SetTransportProtocolOption(tcp.ProtocolNumber, &opt); err != nil {
|
||||||
|
log.Warnf("disable TCP RACK loss detection: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Info("forwarder: TCP RACK loss detection disabled")
|
||||||
|
}
|
||||||
|
|||||||
204
client/firewall/uspfilter/fragment.go
Normal file
204
client/firewall/uspfilter/fragment.go
Normal file
@@ -0,0 +1,204 @@
|
|||||||
|
package uspfilter
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
nblog "github.com/netbirdio/netbird/client/firewall/uspfilter/log"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// defaultFragmentTimeout bounds how long a first-fragment verdict is kept
|
||||||
|
// while the remaining fragments arrive. It mirrors the Linux IP reassembly
|
||||||
|
// timeout (net.ipv4.ipfrag_time).
|
||||||
|
defaultFragmentTimeout = 30 * time.Second
|
||||||
|
// fragmentCleanupInterval is how often expired verdicts are purged.
|
||||||
|
fragmentCleanupInterval = 10 * time.Second
|
||||||
|
// defaultMaxFragmentEntries caps the number of concurrently tracked
|
||||||
|
// fragmented datagrams. The table stays bounded because each datagram is a
|
||||||
|
// single small entry regardless of how many fragments it is split into, and
|
||||||
|
// the 13-bit IPv4 fragment-offset field limits any datagram to 64 KiB.
|
||||||
|
defaultMaxFragmentEntries = 16384
|
||||||
|
|
||||||
|
// EnvFragmentMaxEntries overrides defaultMaxFragmentEntries.
|
||||||
|
EnvFragmentMaxEntries = "NB_FRAGMENT_MAX_ENTRIES"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fragmentVerdict is the decision for a trailing (headerless) fragment.
|
||||||
|
type fragmentVerdict int
|
||||||
|
|
||||||
|
const (
|
||||||
|
// fragmentDeny drops the fragment: no allowed first fragment is on record.
|
||||||
|
fragmentDeny fragmentVerdict = iota
|
||||||
|
// fragmentAllow passes the fragment: it belongs to an allowed datagram and
|
||||||
|
// does not overlap the already-inspected transport header.
|
||||||
|
fragmentAllow
|
||||||
|
// fragmentOverlap drops the fragment and poisons its datagram: it overlaps
|
||||||
|
// the transport header the ACL inspected (RFC 1858 §4, RFC 3128; RFC 5722
|
||||||
|
// requires discarding the whole datagram on overlap for IPv6).
|
||||||
|
fragmentOverlap
|
||||||
|
)
|
||||||
|
|
||||||
|
// fragmentKey identifies a fragmented datagram. It matches the RFC 791 / RFC
|
||||||
|
// 8200 reassembly key: source, destination, protocol and identification. The id
|
||||||
|
// is 32-bit to hold both the IPv4 (16-bit) and IPv6 (32-bit) identification.
|
||||||
|
type fragmentKey struct {
|
||||||
|
srcIP netip.Addr
|
||||||
|
dstIP netip.Addr
|
||||||
|
id uint32
|
||||||
|
proto uint8
|
||||||
|
}
|
||||||
|
|
||||||
|
// fragmentEntry records the verdict of an allowed first fragment.
|
||||||
|
type fragmentEntry struct {
|
||||||
|
// headerEndOctets is the offset, in 8-byte units, at which the first
|
||||||
|
// fragment's payload ended. A trailing fragment starting before this
|
||||||
|
// overlaps bytes the ACL already inspected and is rejected.
|
||||||
|
headerEndOctets uint16
|
||||||
|
// recordedAt is when the first fragment was accepted. The verdict expires a
|
||||||
|
// fixed timeout later and is not refreshed, mirroring the kernel reassembly
|
||||||
|
// timer so a trailing-fragment flood can't keep a datagram alive.
|
||||||
|
recordedAt time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// fragmentTracker records the ACL verdict of a datagram's first fragment so the
|
||||||
|
// remaining fragments, which carry no L4 header, can inherit the decision
|
||||||
|
// without reassembling the datagram. Only allowed first fragments are stored;
|
||||||
|
// anything that cannot be tied to an allowed, non-overlapping first fragment is
|
||||||
|
// dropped (fail closed).
|
||||||
|
type fragmentTracker struct {
|
||||||
|
logger *nblog.Logger
|
||||||
|
mutex sync.Mutex
|
||||||
|
entries map[fragmentKey]fragmentEntry
|
||||||
|
timeout time.Duration
|
||||||
|
// maxEntries caps the table; atCapacity dedups the capacity warning until
|
||||||
|
// the table drains below the cap again.
|
||||||
|
maxEntries int
|
||||||
|
atCapacity bool
|
||||||
|
cleanupTicker *time.Ticker
|
||||||
|
cancel context.CancelFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFragmentTracker(logger *nblog.Logger) *fragmentTracker {
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
t := &fragmentTracker{
|
||||||
|
logger: logger,
|
||||||
|
entries: make(map[fragmentKey]fragmentEntry),
|
||||||
|
timeout: defaultFragmentTimeout,
|
||||||
|
maxEntries: fragmentMaxEntries(logger),
|
||||||
|
cleanupTicker: time.NewTicker(fragmentCleanupInterval),
|
||||||
|
cancel: cancel,
|
||||||
|
}
|
||||||
|
go t.cleanupRoutine(ctx)
|
||||||
|
return t
|
||||||
|
}
|
||||||
|
|
||||||
|
func fragmentMaxEntries(logger *nblog.Logger) int {
|
||||||
|
v := os.Getenv(EnvFragmentMaxEntries)
|
||||||
|
if v == "" {
|
||||||
|
return defaultMaxFragmentEntries
|
||||||
|
}
|
||||||
|
n, err := strconv.Atoi(v)
|
||||||
|
if err != nil || n <= 0 {
|
||||||
|
logger.Warn2("invalid %s=%q, using default", EnvFragmentMaxEntries, v)
|
||||||
|
return defaultMaxFragmentEntries
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
// recordAllowed stores the verdict of an allowed first fragment. headerEndOctets
|
||||||
|
// is the first fragment's payload length in 8-byte units. When the table is full
|
||||||
|
// the record is dropped, which fails closed: the datagram's trailing fragments
|
||||||
|
// will be denied.
|
||||||
|
func (t *fragmentTracker) recordAllowed(key fragmentKey, headerEndOctets uint16) {
|
||||||
|
t.mutex.Lock()
|
||||||
|
defer t.mutex.Unlock()
|
||||||
|
|
||||||
|
if t.entries == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if _, ok := t.entries[key]; !ok && len(t.entries) >= t.maxEntries {
|
||||||
|
if !t.atCapacity {
|
||||||
|
t.atCapacity = true
|
||||||
|
t.logger.Warn2("fragment verdict table at capacity (%d/%d): trailing fragments of new datagrams will be dropped",
|
||||||
|
len(t.entries), t.maxEntries)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
t.entries[key] = fragmentEntry{
|
||||||
|
headerEndOctets: headerEndOctets,
|
||||||
|
recordedAt: time.Now(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// poison drops any recorded verdict for a datagram, so its later fragments are
|
||||||
|
// denied until a new allowed first fragment is recorded. Called on every
|
||||||
|
// offset-zero fragment to defeat offset-zero overlap rewrites (RFC 3128).
|
||||||
|
func (t *fragmentTracker) poison(key fragmentKey) {
|
||||||
|
t.mutex.Lock()
|
||||||
|
defer t.mutex.Unlock()
|
||||||
|
delete(t.entries, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// verdict decides the fate of a trailing fragment at fragOffsetOctets (the IPv4
|
||||||
|
// fragment offset, in 8-byte units). A fragment overlapping the inspected
|
||||||
|
// header poisons the datagram: the entry is removed so all further fragments of
|
||||||
|
// that datagram are denied too.
|
||||||
|
func (t *fragmentTracker) verdict(key fragmentKey, fragOffsetOctets uint16) fragmentVerdict {
|
||||||
|
t.mutex.Lock()
|
||||||
|
defer t.mutex.Unlock()
|
||||||
|
|
||||||
|
entry, ok := t.entries[key]
|
||||||
|
if !ok {
|
||||||
|
return fragmentDeny
|
||||||
|
}
|
||||||
|
if time.Since(entry.recordedAt) > t.timeout {
|
||||||
|
delete(t.entries, key)
|
||||||
|
return fragmentDeny
|
||||||
|
}
|
||||||
|
if fragOffsetOctets < entry.headerEndOctets {
|
||||||
|
delete(t.entries, key)
|
||||||
|
return fragmentOverlap
|
||||||
|
}
|
||||||
|
return fragmentAllow
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *fragmentTracker) cleanupRoutine(ctx context.Context) {
|
||||||
|
defer t.cleanupTicker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-t.cleanupTicker.C:
|
||||||
|
t.cleanup()
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *fragmentTracker) cleanup() {
|
||||||
|
t.mutex.Lock()
|
||||||
|
defer t.mutex.Unlock()
|
||||||
|
|
||||||
|
for key, entry := range t.entries {
|
||||||
|
if time.Since(entry.recordedAt) > t.timeout {
|
||||||
|
delete(t.entries, key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(t.entries) < t.maxEntries {
|
||||||
|
t.atCapacity = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close stops the cleanup routine and releases resources.
|
||||||
|
func (t *fragmentTracker) Close() {
|
||||||
|
t.cancel()
|
||||||
|
|
||||||
|
t.mutex.Lock()
|
||||||
|
t.entries = nil
|
||||||
|
t.mutex.Unlock()
|
||||||
|
}
|
||||||
115
client/firewall/uspfilter/fragment_bench_test.go
Normal file
115
client/firewall/uspfilter/fragment_bench_test.go
Normal file
@@ -0,0 +1,115 @@
|
|||||||
|
package uspfilter
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// benchFilterInbound drives filterInbound over a fixed packet in a tight loop.
|
||||||
|
// Packets are built once, outside the timed region, so the benchmark measures
|
||||||
|
// only pipeline cost, which is what an attacker can amplify.
|
||||||
|
func benchFilterInbound(b *testing.B, pkt []byte) {
|
||||||
|
b.Helper()
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
m := benchManager
|
||||||
|
m.filterInbound(pkt, len(pkt))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// benchManager is a package-level manager reused across fragment benchmarks so
|
||||||
|
// setup cost stays out of the timed region.
|
||||||
|
var benchManager *Manager
|
||||||
|
|
||||||
|
func setupBenchManager(b *testing.B) *Manager {
|
||||||
|
b.Helper()
|
||||||
|
m := newFragmentTestManager(b)
|
||||||
|
allowUDP(b, m, 8080)
|
||||||
|
// Disable conntrack so the allowed-first-fragment path measures transport
|
||||||
|
// decode + ACL every iteration instead of matching the connection tracked
|
||||||
|
// on the first iteration.
|
||||||
|
m.stateful = false
|
||||||
|
benchManager = m
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_NormalPacket is the baseline: a full, non-fragmented UDP
|
||||||
|
// packet that passes the ACL. Fragment paths should stay comparable to this.
|
||||||
|
func BenchmarkInbound_NormalPacket(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := normalUDPPacket(b, 8080, 32)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_FirstFragmentAllowed measures the first-fragment path:
|
||||||
|
// transport decode + ACL evaluation + verdict record.
|
||||||
|
func BenchmarkInbound_FirstFragmentAllowed(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := firstFragmentUDP(b, 0x2000, 8080, 32)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_TrailingFragmentAllowed measures the common trailing-fragment
|
||||||
|
// path: a single map lookup after the first fragment is on record.
|
||||||
|
func BenchmarkInbound_TrailingFragmentAllowed(b *testing.B) {
|
||||||
|
m := setupBenchManager(b)
|
||||||
|
first := firstFragmentUDP(b, 0x3000, 8080, 32)
|
||||||
|
m.filterInbound(first, len(first))
|
||||||
|
pkt := trailingFragment(b, 0x3000, 5, false, 24)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_TrailingFragmentNoFirst is the primary DoS vector: an
|
||||||
|
// attacker floods trailing fragments with no first fragment on record. Each is
|
||||||
|
// a map miss and must be cheap.
|
||||||
|
func BenchmarkInbound_TrailingFragmentNoFirst(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := trailingFragment(b, 0x4000, 185, false, 40)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_TinyFirstFragment measures the tiny-fragment drop path: a
|
||||||
|
// first fragment too small to decode a transport header.
|
||||||
|
func BenchmarkInbound_TinyFirstFragment(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := trailingFragment(b, 0x5000, 0, true, 4)
|
||||||
|
benchFilterInbound(b, pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_TrailingFragmentDistinctIDs is the worst case for the
|
||||||
|
// verdict table: an attacker varies the datagram id on every packet so no first
|
||||||
|
// fragment ever matches. Verdict lookups always miss and nothing is recorded,
|
||||||
|
// so the table cannot grow. Each iteration rewrites the id field in place.
|
||||||
|
func BenchmarkInbound_TrailingFragmentDistinctIDs(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := trailingFragment(b, 0x6000, 185, false, 40)
|
||||||
|
m := benchManager
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
// IPv4 identification field is at bytes 4:6.
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(i))
|
||||||
|
m.filterInbound(pkt, len(pkt))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkInbound_FirstFragmentDistinctIDs measures sustained first-fragment
|
||||||
|
// pressure with distinct ids: transport decode + ACL + verdict insert until the
|
||||||
|
// table caps, exercising the map growth and capacity guard.
|
||||||
|
func BenchmarkInbound_FirstFragmentDistinctIDs(b *testing.B) {
|
||||||
|
setupBenchManager(b)
|
||||||
|
pkt := firstFragmentUDP(b, 0x7000, 8080, 32)
|
||||||
|
m := benchManager
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkt)))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(i))
|
||||||
|
m.filterInbound(pkt, len(pkt))
|
||||||
|
}
|
||||||
|
}
|
||||||
554
client/firewall/uspfilter/fragment_test.go
Normal file
554
client/firewall/uspfilter/fragment_test.go
Normal file
@@ -0,0 +1,554 @@
|
|||||||
|
package uspfilter
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/gopacket"
|
||||||
|
"github.com/google/gopacket/layers"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
fw "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
nbiface "github.com/netbirdio/netbird/client/iface"
|
||||||
|
"github.com/netbirdio/netbird/client/iface/device"
|
||||||
|
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
fragTestSrc = "100.10.0.1"
|
||||||
|
fragTestDst = "100.10.0.100"
|
||||||
|
fragTestSrcV6 = "fd00::1"
|
||||||
|
fragTestDstV6 = "fd00::100"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newFragmentTestManager(tb testing.TB) *Manager {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ifaceMock := &IFaceMock{
|
||||||
|
SetFilterFunc: func(device.PacketFilter) error { return nil },
|
||||||
|
AddressFunc: func() wgaddr.Address {
|
||||||
|
return wgaddr.Address{
|
||||||
|
IP: netip.MustParseAddr(fragTestDst),
|
||||||
|
Network: netip.MustParsePrefix("100.10.0.0/16"),
|
||||||
|
IPv6: netip.MustParseAddr(fragTestDstV6),
|
||||||
|
IPv6Net: netip.MustParsePrefix("fd00::/64"),
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
m, err := Create(ifaceMock, false, flowLogger, nbiface.DefaultMTU)
|
||||||
|
require.NoError(tb, err)
|
||||||
|
require.NoError(tb, m.UpdateLocalIPs())
|
||||||
|
tb.Cleanup(func() { require.NoError(tb, m.Close(nil)) })
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// firstFragmentUDPTo builds the first fragment of a fragmented UDP datagram to
|
||||||
|
// the given destination: it carries the full UDP header plus payloadLen bytes
|
||||||
|
// of data, with the More Fragments flag set and offset zero.
|
||||||
|
func firstFragmentUDPTo(tb testing.TB, dst string, id uint16, dstPort uint16, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: id,
|
||||||
|
Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrc),
|
||||||
|
DstIP: net.ParseIP(dst),
|
||||||
|
Flags: layers.IPv4MoreFragments,
|
||||||
|
}
|
||||||
|
udp := &layers.UDP{SrcPort: 40000, DstPort: layers.UDPPort(dstPort)}
|
||||||
|
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, payloadLen))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func firstFragmentUDP(tb testing.TB, id uint16, dstPort uint16, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
return firstFragmentUDPTo(tb, fragTestDst, id, dstPort, payloadLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
// firstFragmentTCP builds the first fragment of a fragmented TCP datagram: the
|
||||||
|
// full 20-byte TCP header plus 12 bytes of data, with the More Fragments flag
|
||||||
|
// set and offset zero.
|
||||||
|
func firstFragmentTCP(tb testing.TB, id uint16, dstPort uint16) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: id,
|
||||||
|
Protocol: layers.IPProtocolTCP,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrc),
|
||||||
|
DstIP: net.ParseIP(fragTestDst),
|
||||||
|
Flags: layers.IPv4MoreFragments,
|
||||||
|
}
|
||||||
|
tcp := &layers.TCP{SrcPort: 40000, DstPort: layers.TCPPort(dstPort), SYN: true, Window: 64240}
|
||||||
|
require.NoError(tb, tcp.SetNetworkLayerForChecksum(ip))
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, tcp, gopacket.Payload(make([]byte, 12))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
// trailingFragmentTo builds a non-first fragment to the given destination: an
|
||||||
|
// IPv4 header at the given fragment offset (in 8-byte units) carrying raw
|
||||||
|
// payload and no L4 header.
|
||||||
|
func trailingFragmentTo(tb testing.TB, dst string, proto layers.IPProtocol, id uint16, fragOffsetOctets uint16, moreFragments bool, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: id,
|
||||||
|
Protocol: proto,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrc),
|
||||||
|
DstIP: net.ParseIP(dst),
|
||||||
|
FragOffset: fragOffsetOctets,
|
||||||
|
}
|
||||||
|
if moreFragments {
|
||||||
|
ip.Flags = layers.IPv4MoreFragments
|
||||||
|
}
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, gopacket.Payload(make([]byte, payloadLen))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func trailingFragment(tb testing.TB, id uint16, fragOffsetOctets uint16, moreFragments bool, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
return trailingFragmentTo(tb, fragTestDst, layers.IPProtocolUDP, id, fragOffsetOctets, moreFragments, payloadLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
// outboundUDPPacket builds a complete outbound UDP packet from the local
|
||||||
|
// address, used to establish conntrack state for reply-direction tests.
|
||||||
|
func outboundUDPPacket(tb testing.TB, srcPort, dstPort uint16) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: 1,
|
||||||
|
Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: net.ParseIP(fragTestDst),
|
||||||
|
DstIP: net.ParseIP(fragTestSrc),
|
||||||
|
}
|
||||||
|
udp := &layers.UDP{SrcPort: layers.UDPPort(srcPort), DstPort: layers.UDPPort(dstPort)}
|
||||||
|
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, 16))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
// normalUDPPacket builds a complete, non-fragmented UDP packet for baseline
|
||||||
|
// comparisons against the fragment paths.
|
||||||
|
func normalUDPPacket(tb testing.TB, dstPort uint16, payloadLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv4{
|
||||||
|
Version: 4,
|
||||||
|
TTL: 64,
|
||||||
|
Id: 1,
|
||||||
|
Protocol: layers.IPProtocolUDP,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrc),
|
||||||
|
DstIP: net.ParseIP(fragTestDst),
|
||||||
|
}
|
||||||
|
udp := &layers.UDP{SrcPort: 40000, DstPort: layers.UDPPort(dstPort)}
|
||||||
|
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, payloadLen))))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func allowUDP(tb testing.TB, m *Manager, dstPort uint16) {
|
||||||
|
tb.Helper()
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolUDP, nil,
|
||||||
|
&fw.Port{Values: []uint16{dstPort}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(tb, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_TrailingWithoutFirstDropped is the core bypass repro: a trailing
|
||||||
|
// fragment with no allowed first fragment on record must be dropped. Before the
|
||||||
|
// fix, filterInbound returned false (allow) for any fragment.
|
||||||
|
func TestFragment_TrailingWithoutFirstDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
|
||||||
|
frag := trailingFragment(t, 0x1234, 185, false, 40)
|
||||||
|
require.True(t, m.filterInbound(frag, len(frag)),
|
||||||
|
"trailing fragment without an allowed first fragment must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_AllowedFirstPassesTrailing verifies that once a first fragment
|
||||||
|
// passes the ACL, its trailing fragments inherit the allow verdict.
|
||||||
|
func TestFragment_AllowedFirstPassesTrailing(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
// First fragment: UDP header (8) + 32 payload = 40 octets -> headerEnd = 5.
|
||||||
|
first := firstFragmentUDP(t, 0x2222, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"allowed first fragment should pass and be recorded")
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0x2222, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of an allowed datagram should pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_DeniedFirstDropsTrailing verifies that a first fragment blocked
|
||||||
|
// by the ACL leaves no verdict, so its trailing fragments are dropped.
|
||||||
|
func TestFragment_DeniedFirstDropsTrailing(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
// No accept rule: local traffic defaults to deny.
|
||||||
|
|
||||||
|
first := firstFragmentUDP(t, 0x3333, 9999, 32)
|
||||||
|
require.True(t, m.filterInbound(first, len(first)),
|
||||||
|
"first fragment to a blocked port should be dropped by the ACL")
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0x3333, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of a denied datagram must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_OverlappingHeaderDropped covers the RFC 1858 §4 / RFC 3128
|
||||||
|
// overlapping-fragment rewrite: a trailing fragment starting inside the range
|
||||||
|
// the ACL already inspected is dropped and poisons the datagram. TCP is used so
|
||||||
|
// the overlap lands on real header bytes (the flags at byte 13).
|
||||||
|
func TestFragment_OverlappingHeaderDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// First fragment: TCP header (20) + 12 data = 32 bytes -> headerEnd = 4 octets.
|
||||||
|
first := firstFragmentTCP(t, 0x4444, 8080)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)))
|
||||||
|
|
||||||
|
// Overlapping fragment at offset 1 (byte 8) falls inside the inspected TCP
|
||||||
|
// header, so it could rewrite the flags or port on reassembly.
|
||||||
|
overlap := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x4444, 1, true, 32)
|
||||||
|
require.True(t, m.filterInbound(overlap, len(overlap)),
|
||||||
|
"fragment overlapping the inspected header must be dropped")
|
||||||
|
|
||||||
|
// The datagram is now poisoned: a later, non-overlapping fragment is also
|
||||||
|
// dropped because the verdict was removed.
|
||||||
|
later := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x4444, 4, false, 24)
|
||||||
|
require.True(t, m.filterInbound(later, len(later)),
|
||||||
|
"fragments after an overlap must be dropped (datagram poisoned)")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_OffsetZeroOverlapPoisons covers the RFC 3128 offset-zero rewrite:
|
||||||
|
// an allowed first fragment followed by a denied offset-zero fragment for the
|
||||||
|
// same datagram must not leave the earlier allow verdict in place.
|
||||||
|
func TestFragment_OffsetZeroOverlapPoisons(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
allowed := firstFragmentUDP(t, 0x5A5A, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(allowed, len(allowed)),
|
||||||
|
"allowed first fragment should pass and be recorded")
|
||||||
|
|
||||||
|
// A second offset-zero fragment to a denied port supersedes the datagram's
|
||||||
|
// verdict; it is dropped and must not leave the allow in place.
|
||||||
|
denied := firstFragmentUDP(t, 0x5A5A, 9999, 32)
|
||||||
|
require.True(t, m.filterInbound(denied, len(denied)),
|
||||||
|
"denied offset-zero fragment must be dropped")
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0x5A5A, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment must be denied after the datagram was poisoned")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_TinyFirstDropped covers the tiny-fragment attack: a first
|
||||||
|
// fragment too small to contain the full transport header can't be
|
||||||
|
// ACL-evaluated and must be dropped.
|
||||||
|
func TestFragment_TinyFirstDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
// IPv4 header + 4 raw bytes, MF set, offset 0: too small for the 8-byte UDP
|
||||||
|
// header, so it decodes to L3 only.
|
||||||
|
tiny := trailingFragment(t, 0x5555, 0, true, 4)
|
||||||
|
require.True(t, m.filterInbound(tiny, len(tiny)),
|
||||||
|
"tiny first fragment without a full L4 header must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_TCPFirstFragment verifies the TCP arm of the transport decode: a
|
||||||
|
// first fragment carrying the full 20-byte TCP header is ACL-evaluated and its
|
||||||
|
// trailing fragments inherit the verdict.
|
||||||
|
func TestFragment_TCPFirstFragment(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// TCP header (20) + 12 data = 32 bytes -> headerEnd = 4 octets.
|
||||||
|
first := firstFragmentTCP(t, 0x6666, 8080)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"allowed TCP first fragment should pass and be recorded")
|
||||||
|
|
||||||
|
trailing := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x6666, 4, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of an allowed TCP datagram should pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_TCPTinyFirstDropped verifies the TCP minimum header length: 12
|
||||||
|
// bytes would satisfy a UDP header but falls short of the 20-byte TCP header.
|
||||||
|
func TestFragment_TCPTinyFirstDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
tiny := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x7777, 0, true, 12)
|
||||||
|
require.True(t, m.filterInbound(tiny, len(tiny)),
|
||||||
|
"first fragment shorter than the TCP header must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_ConntrackAllowsFirstFragment verifies the conntrack branch: reply
|
||||||
|
// fragments of an outbound-established UDP flow pass without any inbound rule.
|
||||||
|
func TestFragment_ConntrackAllowsFirstFragment(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
|
||||||
|
out := outboundUDPPacket(t, 12345, 40000)
|
||||||
|
require.False(t, m.filterOutbound(out, len(out)))
|
||||||
|
|
||||||
|
first := firstFragmentUDP(t, 0x8888, 12345, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"reply first fragment should pass via conntrack")
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0x8888, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of a tracked flow should pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_RoutingDisabledDropsFragment verifies routed first fragments are
|
||||||
|
// dropped when routing is disabled.
|
||||||
|
func TestFragment_RoutingDisabledDropsFragment(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
m.routingEnabled.Store(false)
|
||||||
|
|
||||||
|
first := firstFragmentUDPTo(t, "198.51.100.10", 0x9999, 8080, 32)
|
||||||
|
require.True(t, m.filterInbound(first, len(first)),
|
||||||
|
"routed first fragment must be dropped when routing is disabled")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_RouteACL verifies the route-ACL branch: fragments to a non-local
|
||||||
|
// destination follow the route rules, allowed datagrams pass their trailing
|
||||||
|
// fragments and denied ones don't.
|
||||||
|
func TestFragment_RouteACL(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
m.routingEnabled.Store(true)
|
||||||
|
m.nativeRouter.Store(false)
|
||||||
|
|
||||||
|
_, err := m.AddRouteFiltering(
|
||||||
|
[]byte("rt-1"),
|
||||||
|
[]netip.Prefix{netip.MustParsePrefix("100.10.0.0/16")},
|
||||||
|
fw.Network{Prefix: netip.MustParsePrefix("198.51.100.0/24")},
|
||||||
|
fw.ProtocolUDP,
|
||||||
|
nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}},
|
||||||
|
fw.ActionAccept,
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
first := firstFragmentUDPTo(t, "198.51.100.10", 0xAAAA, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"route-ACL-allowed first fragment should pass")
|
||||||
|
trailing := trailingFragmentTo(t, "198.51.100.10", layers.IPProtocolUDP, 0xAAAA, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of an allowed routed datagram should pass")
|
||||||
|
|
||||||
|
denied := firstFragmentUDPTo(t, "198.51.100.10", 0xBBBB, 9999, 32)
|
||||||
|
require.True(t, m.filterInbound(denied, len(denied)),
|
||||||
|
"route-ACL-denied first fragment must be dropped")
|
||||||
|
deniedTrailing := trailingFragmentTo(t, "198.51.100.10", layers.IPProtocolUDP, 0xBBBB, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(deniedTrailing, len(deniedTrailing)),
|
||||||
|
"trailing fragment of a denied routed datagram must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_ExpiredVerdictDropsTrailing verifies a verdict older than the
|
||||||
|
// tracker timeout no longer admits trailing fragments.
|
||||||
|
func TestFragment_ExpiredVerdictDropsTrailing(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
first := firstFragmentUDP(t, 0xCCCC, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)))
|
||||||
|
|
||||||
|
m.fragments.mutex.Lock()
|
||||||
|
for key, entry := range m.fragments.entries {
|
||||||
|
entry.recordedAt = time.Now().Add(-defaultFragmentTimeout - time.Second)
|
||||||
|
m.fragments.entries[key] = entry
|
||||||
|
}
|
||||||
|
m.fragments.mutex.Unlock()
|
||||||
|
|
||||||
|
trailing := trailingFragment(t, 0xCCCC, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment after verdict expiry must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragment_CapacityFailsClosed verifies the table cap: at capacity, new
|
||||||
|
// datagram verdicts are not recorded (their trailing fragments are dropped)
|
||||||
|
// while already-recorded datagrams keep working.
|
||||||
|
func TestFragment_CapacityFailsClosed(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
allowUDP(t, m, 8080)
|
||||||
|
|
||||||
|
m.fragments.mutex.Lock()
|
||||||
|
m.fragments.maxEntries = 1
|
||||||
|
m.fragments.mutex.Unlock()
|
||||||
|
|
||||||
|
first1 := firstFragmentUDP(t, 0x0101, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first1, len(first1)))
|
||||||
|
|
||||||
|
first2 := firstFragmentUDP(t, 0x0202, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first2, len(first2)),
|
||||||
|
"first fragment itself still passes at capacity")
|
||||||
|
|
||||||
|
trailing2 := trailingFragment(t, 0x0202, 5, false, 24)
|
||||||
|
require.True(t, m.filterInbound(trailing2, len(trailing2)),
|
||||||
|
"trailing fragment of an unrecorded datagram must be dropped at capacity")
|
||||||
|
|
||||||
|
trailing1 := trailingFragment(t, 0x0101, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing1, len(trailing1)),
|
||||||
|
"already-recorded datagram should keep passing at capacity")
|
||||||
|
}
|
||||||
|
|
||||||
|
// v6FragmentHeader builds the 8-byte IPv6 fragment extension header for the
|
||||||
|
// given inner protocol, offset (8-byte units), More Fragments bit and id.
|
||||||
|
func v6FragmentHeader(proto layers.IPProtocol, offsetOctets uint16, moreFragments bool, id uint32) []byte {
|
||||||
|
offsetFlags := offsetOctets << 3
|
||||||
|
if moreFragments {
|
||||||
|
offsetFlags |= 1
|
||||||
|
}
|
||||||
|
hdr := make([]byte, 8)
|
||||||
|
hdr[0] = uint8(proto)
|
||||||
|
binary.BigEndian.PutUint16(hdr[2:4], offsetFlags)
|
||||||
|
binary.BigEndian.PutUint32(hdr[4:8], id)
|
||||||
|
return hdr
|
||||||
|
}
|
||||||
|
|
||||||
|
func v6UDPHeader(dstPort uint16, dataLen int) []byte {
|
||||||
|
hdr := make([]byte, 8)
|
||||||
|
binary.BigEndian.PutUint16(hdr[0:2], 40000)
|
||||||
|
binary.BigEndian.PutUint16(hdr[2:4], dstPort)
|
||||||
|
binary.BigEndian.PutUint16(hdr[4:6], uint16(8+dataLen))
|
||||||
|
return hdr
|
||||||
|
}
|
||||||
|
|
||||||
|
// firstFragmentUDPv6 builds the first fragment of a fragmented IPv6 UDP
|
||||||
|
// datagram: fragment header (offset 0, More Fragments set) + full UDP header +
|
||||||
|
// data.
|
||||||
|
func firstFragmentUDPv6(tb testing.TB, id uint32, dstPort uint16, dataLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
return fragmentUDPv6(tb, id, dstPort, dataLen, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// fragmentUDPv6 builds an offset-zero IPv6 UDP fragment. With moreFragments
|
||||||
|
// false it is an atomic fragment (a complete datagram, RFC 6946).
|
||||||
|
func fragmentUDPv6(tb testing.TB, id uint32, dstPort uint16, dataLen int, moreFragments bool) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv6{
|
||||||
|
Version: 6,
|
||||||
|
NextHeader: layers.IPProtocolIPv6Fragment,
|
||||||
|
HopLimit: 64,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrcV6),
|
||||||
|
DstIP: net.ParseIP(fragTestDstV6),
|
||||||
|
}
|
||||||
|
payload := append(v6FragmentHeader(layers.IPProtocolUDP, 0, moreFragments, id), v6UDPHeader(dstPort, dataLen)...)
|
||||||
|
payload = append(payload, make([]byte, dataLen)...)
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true}, ip, gopacket.Payload(payload)))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
// trailingFragmentV6 builds a non-first IPv6 fragment: fragment header at the
|
||||||
|
// given offset carrying raw data and no transport header.
|
||||||
|
func trailingFragmentV6(tb testing.TB, id uint32, offsetOctets uint16, moreFragments bool, dataLen int) []byte {
|
||||||
|
tb.Helper()
|
||||||
|
|
||||||
|
ip := &layers.IPv6{
|
||||||
|
Version: 6,
|
||||||
|
NextHeader: layers.IPProtocolIPv6Fragment,
|
||||||
|
HopLimit: 64,
|
||||||
|
SrcIP: net.ParseIP(fragTestSrcV6),
|
||||||
|
DstIP: net.ParseIP(fragTestDstV6),
|
||||||
|
}
|
||||||
|
payload := append(v6FragmentHeader(layers.IPProtocolUDP, offsetOctets, moreFragments, id), make([]byte, dataLen)...)
|
||||||
|
|
||||||
|
buf := gopacket.NewSerializeBuffer()
|
||||||
|
require.NoError(tb, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true}, ip, gopacket.Payload(payload)))
|
||||||
|
return buf.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragmentV6_TrailingWithoutFirstDropped verifies the IPv6 bypass is closed:
|
||||||
|
// a trailing fragment with no allowed first fragment is dropped.
|
||||||
|
func TestFragmentV6_TrailingWithoutFirstDropped(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
|
||||||
|
frag := trailingFragmentV6(t, 0xAABBCCDD, 100, false, 40)
|
||||||
|
require.True(t, m.filterInbound(frag, len(frag)),
|
||||||
|
"IPv6 trailing fragment without an allowed first fragment must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragmentV6_AllowedFirstPassesTrailing verifies IPv6 fragments are
|
||||||
|
// evaluated like IPv4: an allowed first fragment lets its trailing fragments
|
||||||
|
// through.
|
||||||
|
func TestFragmentV6_AllowedFirstPassesTrailing(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrcV6), fw.ProtocolUDP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// First fragment: UDP header (8) + 32 data = 40 octets -> headerEnd = 5.
|
||||||
|
first := firstFragmentUDPv6(t, 0xAABBCCDD, 8080, 32)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)),
|
||||||
|
"allowed IPv6 first fragment should pass and be recorded")
|
||||||
|
|
||||||
|
trailing := trailingFragmentV6(t, 0xAABBCCDD, 5, false, 24)
|
||||||
|
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||||
|
"trailing fragment of an allowed IPv6 datagram should pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestFragmentV6_AtomicNotCached verifies an IPv6 atomic fragment (fragment
|
||||||
|
// header with offset 0 and no More Fragments, a complete datagram per RFC 6946)
|
||||||
|
// is evaluated but not recorded, so a flood of allowed atomic fragments can't
|
||||||
|
// exhaust the verdict table.
|
||||||
|
func TestFragmentV6_AtomicNotCached(t *testing.T) {
|
||||||
|
m := newFragmentTestManager(t)
|
||||||
|
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrcV6), fw.ProtocolUDP, nil,
|
||||||
|
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
atomic := fragmentUDPv6(t, 0xA70301C, 8080, 16, false)
|
||||||
|
require.False(t, m.filterInbound(atomic, len(atomic)),
|
||||||
|
"allowed IPv6 atomic fragment should pass")
|
||||||
|
|
||||||
|
m.fragments.mutex.Lock()
|
||||||
|
n := len(m.fragments.entries)
|
||||||
|
m.fragments.mutex.Unlock()
|
||||||
|
require.Zero(t, n, "atomic fragment must not create a verdict entry")
|
||||||
|
|
||||||
|
// A genuine fragmented datagram (More Fragments set) is still recorded.
|
||||||
|
first := fragmentUDPv6(t, 0xBEEF, 8080, 32, true)
|
||||||
|
require.False(t, m.filterInbound(first, len(first)))
|
||||||
|
m.fragments.mutex.Lock()
|
||||||
|
n = len(m.fragments.entries)
|
||||||
|
m.fragments.mutex.Unlock()
|
||||||
|
require.Equal(t, 1, n, "genuine first fragment must record a verdict")
|
||||||
|
}
|
||||||
@@ -464,6 +464,8 @@ func Test_RemovePeer(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func Test_ConnectPeers(t *testing.T) {
|
func Test_ConnectPeers(t *testing.T) {
|
||||||
|
t.Setenv("NB_DISABLE_EBPF_WG_PROXY", "true")
|
||||||
|
|
||||||
peer1ifaceName := fmt.Sprintf("utun%d", WgIntNumber+400)
|
peer1ifaceName := fmt.Sprintf("utun%d", WgIntNumber+400)
|
||||||
peer1wgIP := netip.MustParsePrefix("10.99.99.17/30")
|
peer1wgIP := netip.MustParsePrefix("10.99.99.17/30")
|
||||||
peer1Key, _ := wgtypes.GeneratePrivateKey()
|
peer1Key, _ := wgtypes.GeneratePrivateKey()
|
||||||
@@ -505,12 +507,8 @@ func Test_ConnectPeers(t *testing.T) {
|
|||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
localIP, err := getLocalIP()
|
localIP1 := "127.0.0.1"
|
||||||
if err != nil {
|
peer1endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP1, peer1wgPort))
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
peer1endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP, peer1wgPort))
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
@@ -546,7 +544,8 @@ func Test_ConnectPeers(t *testing.T) {
|
|||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
peer2endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP, peer2wgPort))
|
localIP2 := "127.0.0.1"
|
||||||
|
peer2endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP2, peer2wgPort))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
@@ -569,17 +568,17 @@ func Test_ConnectPeers(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
// todo: investigate why in some tests execution we need 30s
|
// The peers use userspace WireGuard (stdnet transport). A tight busy-loop
|
||||||
|
// here starves the wireguard-go goroutines that process the handshake, so
|
||||||
|
// poll on a ticker instead and yield the CPU between checks. WireGuard also
|
||||||
|
// only retries a lost handshake initiation every REKEY_TIMEOUT (5s), which
|
||||||
|
// is why the overall wait can occasionally stretch to tens of seconds.
|
||||||
timeout := 30 * time.Second
|
timeout := 30 * time.Second
|
||||||
timeoutChannel := time.After(timeout)
|
timeoutChannel := time.After(timeout)
|
||||||
|
ticker := time.NewTicker(500 * time.Millisecond)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
|
||||||
case <-timeoutChannel:
|
|
||||||
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String())
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
|
|
||||||
peer, gpErr := getPeer(peer1ifaceName, peer2Key.PublicKey().String())
|
peer, gpErr := getPeer(peer1ifaceName, peer2Key.PublicKey().String())
|
||||||
if gpErr != nil {
|
if gpErr != nil {
|
||||||
t.Fatal(gpErr)
|
t.Fatal(gpErr)
|
||||||
@@ -588,6 +587,12 @@ func Test_ConnectPeers(t *testing.T) {
|
|||||||
t.Log("peers successfully handshake")
|
t.Log("peers successfully handshake")
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-timeoutChannel:
|
||||||
|
t.Fatalf("waiting for peer handshake timeout after %s", timeout.String())
|
||||||
|
case <-ticker.C:
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
@@ -615,28 +620,3 @@ func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
|
|||||||
}
|
}
|
||||||
return wgtypes.Peer{}, fmt.Errorf("peer not found")
|
return wgtypes.Peer{}, fmt.Errorf("peer not found")
|
||||||
}
|
}
|
||||||
|
|
||||||
func getLocalIP() (string, error) {
|
|
||||||
// Get all interfaces
|
|
||||||
addrs, err := net.InterfaceAddrs()
|
|
||||||
if err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, addr := range addrs {
|
|
||||||
ipNet, ok := addr.(*net.IPNet)
|
|
||||||
if !ok {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if ipNet.IP.IsLoopback() {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if ipNet.IP.To4() == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
return ipNet.IP.String(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return "", fmt.Errorf("no local IP found")
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -3,14 +3,31 @@
|
|||||||
package netstack
|
package netstack
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
)
|
)
|
||||||
|
|
||||||
const EnvUseNetstackMode = "NB_USE_NETSTACK_MODE"
|
const (
|
||||||
|
EnvUseNetstackMode = "NB_USE_NETSTACK_MODE"
|
||||||
|
|
||||||
|
// EnvSocks5ListenerPort overrides the port the SOCKS5 proxy listens on.
|
||||||
|
EnvSocks5ListenerPort = "NB_SOCKS5_LISTENER_PORT"
|
||||||
|
|
||||||
|
// EnvSocks5ListenerAddress overrides the host/IP the SOCKS5 proxy binds to.
|
||||||
|
// The proxy is a bridge for local host applications into the userspace
|
||||||
|
// WireGuard netstack, so it binds to loopback by default. Override this only
|
||||||
|
// when the proxy must be reachable from other hosts (e.g. a container
|
||||||
|
// gateway); doing so exposes an unauthenticated SOCKS5 proxy on that
|
||||||
|
// address.
|
||||||
|
EnvSocks5ListenerAddress = "NB_SOCKS5_LISTENER_ADDRESS"
|
||||||
|
|
||||||
|
// defaultSocks5Host is the loopback address the SOCKS5 proxy binds to unless
|
||||||
|
// overridden via EnvSocks5ListenerAddress.
|
||||||
|
defaultSocks5Host = "127.0.0.1"
|
||||||
|
)
|
||||||
|
|
||||||
// IsEnabled todo: move these function to cmd layer
|
// IsEnabled todo: move these function to cmd layer
|
||||||
func IsEnabled() bool {
|
func IsEnabled() bool {
|
||||||
@@ -18,24 +35,40 @@ func IsEnabled() bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func ListenAddr() string {
|
func ListenAddr() string {
|
||||||
sPort := os.Getenv("NB_SOCKS5_LISTENER_PORT")
|
return net.JoinHostPort(listenHost(), strconv.Itoa(listenPort()))
|
||||||
|
}
|
||||||
|
|
||||||
|
// listenHost returns the host/IP the SOCKS5 proxy binds to. It defaults to
|
||||||
|
// loopback and only honors EnvSocks5ListenerAddress when it holds a valid IP.
|
||||||
|
func listenHost() string {
|
||||||
|
addr := os.Getenv(EnvSocks5ListenerAddress)
|
||||||
|
if addr == "" {
|
||||||
|
return defaultSocks5Host
|
||||||
|
}
|
||||||
|
if net.ParseIP(addr) == nil {
|
||||||
|
log.Warnf("invalid socks5 listener address %q, falling back to default: %s", addr, defaultSocks5Host)
|
||||||
|
return defaultSocks5Host
|
||||||
|
}
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
|
||||||
|
// listenPort returns the port the SOCKS5 proxy binds to, defaulting to
|
||||||
|
// DefaultSocks5Port when EnvSocks5ListenerPort is unset or invalid.
|
||||||
|
func listenPort() int {
|
||||||
|
sPort := os.Getenv(EnvSocks5ListenerPort)
|
||||||
if sPort == "" {
|
if sPort == "" {
|
||||||
return listenAddr(DefaultSocks5Port)
|
return DefaultSocks5Port
|
||||||
}
|
}
|
||||||
|
|
||||||
port, err := strconv.Atoi(sPort)
|
port, err := strconv.Atoi(sPort)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Warnf("invalid socks5 listener port, unable to convert it to int, falling back to default: %d", DefaultSocks5Port)
|
log.Warnf("invalid socks5 listener port, unable to convert it to int, falling back to default: %d", DefaultSocks5Port)
|
||||||
return listenAddr(DefaultSocks5Port)
|
return DefaultSocks5Port
|
||||||
}
|
}
|
||||||
if port < 1 || port > 65535 {
|
if port < 1 || port > 65535 {
|
||||||
log.Warnf("invalid socks5 listener port, it should be in the range 1-65535, falling back to default: %d", DefaultSocks5Port)
|
log.Warnf("invalid socks5 listener port, it should be in the range 1-65535, falling back to default: %d", DefaultSocks5Port)
|
||||||
return listenAddr(DefaultSocks5Port)
|
return DefaultSocks5Port
|
||||||
}
|
}
|
||||||
|
|
||||||
return listenAddr(port)
|
return port
|
||||||
}
|
|
||||||
|
|
||||||
func listenAddr(port int) string {
|
|
||||||
return fmt.Sprintf("0.0.0.0:%d", port)
|
|
||||||
}
|
}
|
||||||
|
|||||||
63
client/iface/netstack/env_test.go
Normal file
63
client/iface/netstack/env_test.go
Normal file
@@ -0,0 +1,63 @@
|
|||||||
|
//go:build !js
|
||||||
|
|
||||||
|
package netstack
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestListenAddr_DefaultsToLoopback(t *testing.T) {
|
||||||
|
// No env overrides: must bind loopback, never all interfaces.
|
||||||
|
got := ListenAddr()
|
||||||
|
want := net.JoinHostPort("127.0.0.1", strconv.Itoa(DefaultSocks5Port))
|
||||||
|
if got != want {
|
||||||
|
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListenAddr_AddressOverride(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
env string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{name: "valid override honored", env: "0.0.0.0", want: "0.0.0.0"},
|
||||||
|
{name: "valid specific ip honored", env: "10.0.0.5", want: "10.0.0.5"},
|
||||||
|
{name: "ipv6 loopback bracketed", env: "::1", want: "::1"},
|
||||||
|
{name: "invalid falls back to loopback", env: "not-an-ip", want: "127.0.0.1"},
|
||||||
|
{name: "empty falls back to loopback", env: "", want: "127.0.0.1"},
|
||||||
|
}
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Setenv(EnvSocks5ListenerAddress, tc.env)
|
||||||
|
want := net.JoinHostPort(tc.want, strconv.Itoa(DefaultSocks5Port))
|
||||||
|
if got := ListenAddr(); got != want {
|
||||||
|
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListenAddr_PortOverride(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
env string
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{name: "valid port honored", env: "1081", want: 1081},
|
||||||
|
{name: "non-numeric falls back", env: "abc", want: DefaultSocks5Port},
|
||||||
|
{name: "out of range falls back", env: "70000", want: DefaultSocks5Port},
|
||||||
|
{name: "zero falls back", env: "0", want: DefaultSocks5Port},
|
||||||
|
}
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Setenv(EnvSocks5ListenerPort, tc.env)
|
||||||
|
want := net.JoinHostPort("127.0.0.1", strconv.Itoa(tc.want))
|
||||||
|
if got := ListenAddr(); got != want {
|
||||||
|
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -6,7 +6,7 @@
|
|||||||
!define DESCRIPTION "Connect your devices into a secure WireGuard-based overlay network with SSO, MFA, and granular access controls."
|
!define DESCRIPTION "Connect your devices into a secure WireGuard-based overlay network with SSO, MFA, and granular access controls."
|
||||||
!define INSTALLER_NAME "netbird-installer.exe"
|
!define INSTALLER_NAME "netbird-installer.exe"
|
||||||
!define MAIN_APP_EXE "Netbird"
|
!define MAIN_APP_EXE "Netbird"
|
||||||
!define ICON "ui\\assets\\netbird.ico"
|
!define ICON "ui\\build\\windows\\icon.ico"
|
||||||
!define BANNER "ui\\build\\banner.bmp"
|
!define BANNER "ui\\build\\banner.bmp"
|
||||||
!define LICENSE_DATA "..\\LICENSE"
|
!define LICENSE_DATA "..\\LICENSE"
|
||||||
|
|
||||||
@@ -79,8 +79,6 @@ ShowInstDetails Show
|
|||||||
|
|
||||||
!insertmacro MUI_PAGE_DIRECTORY
|
!insertmacro MUI_PAGE_DIRECTORY
|
||||||
|
|
||||||
Page custom AutostartPage AutostartPageLeave
|
|
||||||
|
|
||||||
!insertmacro MUI_PAGE_INSTFILES
|
!insertmacro MUI_PAGE_INSTFILES
|
||||||
|
|
||||||
!insertmacro MUI_PAGE_FINISH
|
!insertmacro MUI_PAGE_FINISH
|
||||||
@@ -97,40 +95,12 @@ UninstPage custom un.DeleteDataPage un.DeleteDataPageLeave
|
|||||||
|
|
||||||
!insertmacro MUI_LANGUAGE "English"
|
!insertmacro MUI_LANGUAGE "English"
|
||||||
|
|
||||||
; Variables for autostart option
|
|
||||||
Var AutostartCheckbox
|
|
||||||
Var AutostartEnabled
|
|
||||||
|
|
||||||
; Variables for uninstall data deletion option
|
; Variables for uninstall data deletion option
|
||||||
Var DeleteDataCheckbox
|
Var DeleteDataCheckbox
|
||||||
Var DeleteDataEnabled
|
Var DeleteDataEnabled
|
||||||
|
|
||||||
######################################################################
|
######################################################################
|
||||||
|
|
||||||
; Function to create the autostart options page
|
|
||||||
Function AutostartPage
|
|
||||||
!insertmacro MUI_HEADER_TEXT "Startup Options" "Configure how ${APP_NAME} launches with Windows."
|
|
||||||
|
|
||||||
nsDialogs::Create 1018
|
|
||||||
Pop $0
|
|
||||||
|
|
||||||
${If} $0 == error
|
|
||||||
Abort
|
|
||||||
${EndIf}
|
|
||||||
|
|
||||||
${NSD_CreateCheckbox} 0 20u 100% 10u "Start ${APP_NAME} UI automatically when Windows starts"
|
|
||||||
Pop $AutostartCheckbox
|
|
||||||
${NSD_Check} $AutostartCheckbox
|
|
||||||
StrCpy $AutostartEnabled "1"
|
|
||||||
|
|
||||||
nsDialogs::Show
|
|
||||||
FunctionEnd
|
|
||||||
|
|
||||||
; Function to handle leaving the autostart page
|
|
||||||
Function AutostartPageLeave
|
|
||||||
${NSD_GetState} $AutostartCheckbox $AutostartEnabled
|
|
||||||
FunctionEnd
|
|
||||||
|
|
||||||
; Function to create the uninstall data deletion page
|
; Function to create the uninstall data deletion page
|
||||||
Function un.DeleteDataPage
|
Function un.DeleteDataPage
|
||||||
!insertmacro MUI_HEADER_TEXT "Uninstall Options" "Choose whether to delete ${APP_NAME} data."
|
!insertmacro MUI_HEADER_TEXT "Uninstall Options" "Choose whether to delete ${APP_NAME} data."
|
||||||
@@ -201,8 +171,6 @@ Pop $0
|
|||||||
|
|
||||||
Function .onInit
|
Function .onInit
|
||||||
StrCpy $INSTDIR "${INSTALL_DIR}"
|
StrCpy $INSTDIR "${INSTALL_DIR}"
|
||||||
; Default autostart to enabled so silent installs (/S) match the interactive default
|
|
||||||
StrCpy $AutostartEnabled "1"
|
|
||||||
|
|
||||||
; Pre-0.70.1 installers ran without SetRegView, so their uninstall keys live
|
; Pre-0.70.1 installers ran without SetRegView, so their uninstall keys live
|
||||||
; in the 32-bit view. Fall back to it so upgrades still find them.
|
; in the 32-bit view. Fall back to it so upgrades still find them.
|
||||||
@@ -260,17 +228,12 @@ WriteRegStr ${REG_ROOT} "${UNINSTALL_PATH}" "Publisher" "${COMP_NAME}"
|
|||||||
|
|
||||||
WriteRegStr ${REG_ROOT} "${UI_REG_APP_PATH}" "" "$INSTDIR\${UI_APP_EXE}"
|
WriteRegStr ${REG_ROOT} "${UI_REG_APP_PATH}" "" "$INSTDIR\${UI_APP_EXE}"
|
||||||
|
|
||||||
; Create autostart registry entry based on checkbox
|
; Autostart is owned by the UI's per-user setting (HKCU\...\Run via Wails),
|
||||||
DetailPrint "Autostart enabled: $AutostartEnabled"
|
; not the installer. Drop the machine-wide entry older installers wrote so the
|
||||||
${If} $AutostartEnabled == "1"
|
; toggle is the single source of truth. HKCU is left untouched -- it may hold
|
||||||
WriteRegStr HKLM "${AUTOSTART_REG_KEY}" "${APP_NAME}" '"$INSTDIR\${UI_APP_EXE}.exe"'
|
; the user's own toggle state, which must survive upgrades.
|
||||||
DetailPrint "Added autostart registry entry: $INSTDIR\${UI_APP_EXE}.exe"
|
DetailPrint "Removing installer-managed autostart registry entry if present..."
|
||||||
${Else}
|
DeleteRegValue HKLM "${AUTOSTART_REG_KEY}" "${APP_NAME}"
|
||||||
DeleteRegValue HKLM "${AUTOSTART_REG_KEY}" "${APP_NAME}"
|
|
||||||
; Legacy: pre-HKLM installs wrote to HKCU; clean that up too.
|
|
||||||
DeleteRegValue HKCU "${AUTOSTART_REG_KEY}" "${APP_NAME}"
|
|
||||||
DetailPrint "Autostart not enabled by user"
|
|
||||||
${EndIf}
|
|
||||||
|
|
||||||
EnVar::SetHKLM
|
EnVar::SetHKLM
|
||||||
EnVar::AddValueEx "path" "$INSTDIR"
|
EnVar::AddValueEx "path" "$INSTDIR"
|
||||||
@@ -280,6 +243,43 @@ CreateShortCut "$SMPROGRAMS\${APP_NAME}.lnk" "$INSTDIR\${UI_APP_EXE}"
|
|||||||
CreateShortCut "$DESKTOP\${APP_NAME}.lnk" "$INSTDIR\${UI_APP_EXE}"
|
CreateShortCut "$DESKTOP\${APP_NAME}.lnk" "$INSTDIR\${UI_APP_EXE}"
|
||||||
SectionEnd
|
SectionEnd
|
||||||
|
|
||||||
|
# Install the Microsoft Edge WebView2 runtime if it isn't already present.
|
||||||
|
# Macro adapted from Wails3's NSIS template (wails_tools.nsh): a registry
|
||||||
|
# probe followed by a silent install of the embedded evergreen bootstrapper.
|
||||||
|
# The MicrosoftEdgeWebview2Setup.exe payload is staged next to this script
|
||||||
|
# by the sign-pipelines build step (`wails3 generate webview2bootstrapper`).
|
||||||
|
!macro nb.webview2runtime
|
||||||
|
SetRegView 64
|
||||||
|
# Per-machine install marker — populated when the runtime ships with
|
||||||
|
# Edge or has been installed by an admin previously.
|
||||||
|
ReadRegStr $0 HKLM "SOFTWARE\WOW6432Node\Microsoft\EdgeUpdate\Clients\{F3017226-FE2A-4295-8BDF-00C3A9A7E4C5}" "pv"
|
||||||
|
${If} $0 != ""
|
||||||
|
Goto webview2_ok
|
||||||
|
${EndIf}
|
||||||
|
# Per-user fallback for HKCU installs.
|
||||||
|
ReadRegStr $0 HKCU "Software\Microsoft\EdgeUpdate\Clients\{F3017226-FE2A-4295-8BDF-00C3A9A7E4C5}" "pv"
|
||||||
|
${If} $0 != ""
|
||||||
|
Goto webview2_ok
|
||||||
|
${EndIf}
|
||||||
|
|
||||||
|
SetDetailsPrint both
|
||||||
|
DetailPrint "Installing: WebView2 Runtime"
|
||||||
|
SetDetailsPrint listonly
|
||||||
|
|
||||||
|
InitPluginsDir
|
||||||
|
CreateDirectory "$pluginsdir\webview2bootstrapper"
|
||||||
|
SetOutPath "$pluginsdir\webview2bootstrapper"
|
||||||
|
File "MicrosoftEdgeWebview2Setup.exe"
|
||||||
|
ExecWait '"$pluginsdir\webview2bootstrapper\MicrosoftEdgeWebview2Setup.exe" /silent /install'
|
||||||
|
|
||||||
|
SetDetailsPrint both
|
||||||
|
webview2_ok:
|
||||||
|
!macroend
|
||||||
|
|
||||||
|
Section -WebView2
|
||||||
|
!insertmacro nb.webview2runtime
|
||||||
|
SectionEnd
|
||||||
|
|
||||||
Section -Post
|
Section -Post
|
||||||
ExecWait '"$INSTDIR\${MAIN_APP_EXE}" service install'
|
ExecWait '"$INSTDIR\${MAIN_APP_EXE}" service install'
|
||||||
ExecWait '"$INSTDIR\${MAIN_APP_EXE}" service start'
|
ExecWait '"$INSTDIR\${MAIN_APP_EXE}" service start'
|
||||||
@@ -299,11 +299,14 @@ ExecWait '"$INSTDIR\${MAIN_APP_EXE}" service uninstall'
|
|||||||
DetailPrint "Terminating Netbird UI process..."
|
DetailPrint "Terminating Netbird UI process..."
|
||||||
ExecWait `taskkill /im ${UI_APP_EXE}.exe /f`
|
ExecWait `taskkill /im ${UI_APP_EXE}.exe /f`
|
||||||
|
|
||||||
; Remove autostart registry entry
|
; Remove autostart registry entries
|
||||||
DetailPrint "Removing autostart registry entry if exists..."
|
DetailPrint "Removing autostart registry entries if they exist..."
|
||||||
|
; Legacy machine-wide entry written by older installers.
|
||||||
DeleteRegValue HKLM "${AUTOSTART_REG_KEY}" "${APP_NAME}"
|
DeleteRegValue HKLM "${AUTOSTART_REG_KEY}" "${APP_NAME}"
|
||||||
; Legacy: pre-HKLM installs wrote to HKCU; clean that up too.
|
; Per-user entry the UI toggle writes via Wails (value name is the lowercase
|
||||||
|
; app-name slug). Uninstall removes the app, so drop it too.
|
||||||
DeleteRegValue HKCU "${AUTOSTART_REG_KEY}" "${APP_NAME}"
|
DeleteRegValue HKCU "${AUTOSTART_REG_KEY}" "${APP_NAME}"
|
||||||
|
DeleteRegValue HKCU "${AUTOSTART_REG_KEY}" "netbird"
|
||||||
|
|
||||||
; Handle data deletion based on checkbox
|
; Handle data deletion based on checkbox
|
||||||
DetailPrint "Checking if user requested data deletion..."
|
DetailPrint "Checking if user requested data deletion..."
|
||||||
@@ -326,9 +329,9 @@ DetailPrint "Deleting application files..."
|
|||||||
Delete "$INSTDIR\${UI_APP_EXE}"
|
Delete "$INSTDIR\${UI_APP_EXE}"
|
||||||
Delete "$INSTDIR\${MAIN_APP_EXE}"
|
Delete "$INSTDIR\${MAIN_APP_EXE}"
|
||||||
Delete "$INSTDIR\wintun.dll"
|
Delete "$INSTDIR\wintun.dll"
|
||||||
!if ${ARCH} == "amd64"
|
# Legacy: pre-Wails installs shipped opengl32.dll (Mesa3D for Fyne); remove
|
||||||
|
# any leftover copy on uninstall so old upgrades don't leave it behind.
|
||||||
Delete "$INSTDIR\opengl32.dll"
|
Delete "$INSTDIR\opengl32.dll"
|
||||||
!endif
|
|
||||||
DetailPrint "Removing application directory..."
|
DetailPrint "Removing application directory..."
|
||||||
RmDir /r "$INSTDIR"
|
RmDir /r "$INSTDIR"
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package auth
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -21,6 +22,25 @@ import (
|
|||||||
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// peerLoginExpiredMsg is the exact phrase the management server returns
|
||||||
|
// when a previously SSO-enrolled peer's login has expired. Sourced from
|
||||||
|
// shared/management/status/error.go (NewPeerLoginExpiredError). Matched
|
||||||
|
// by substring so a future server-side rewording that keeps the phrase
|
||||||
|
// still triggers the friendly fallback in Login().
|
||||||
|
const peerLoginExpiredMsg = "peer login has expired"
|
||||||
|
|
||||||
|
// errSetupKeyOnSSOExpiredPeer replaces the raw management error when the
|
||||||
|
// user runs `netbird login -k <setup-key>` against a peer that was
|
||||||
|
// originally enrolled via SSO. Wrapped in a PermissionDenied gRPC status
|
||||||
|
// so callers' existing isPermissionDenied / isAuthError checks still
|
||||||
|
// classify it correctly (early-exit from retry backoff, StatusNeedsLogin
|
||||||
|
// in the server state machine).
|
||||||
|
var errSetupKeyOnSSOExpiredPeer = status.Error(
|
||||||
|
codes.PermissionDenied,
|
||||||
|
"this peer was originally enrolled via SSO and its session has expired. "+
|
||||||
|
"Setup keys can only enrol new peers — run `netbird up` (interactive SSO) to re-login.",
|
||||||
|
)
|
||||||
|
|
||||||
// Auth manages authentication operations with the management server
|
// Auth manages authentication operations with the management server
|
||||||
// It maintains a long-lived connection and automatically handles reconnection with backoff
|
// It maintains a long-lived connection and automatically handles reconnection with backoff
|
||||||
type Auth struct {
|
type Auth struct {
|
||||||
@@ -184,6 +204,15 @@ func (a *Auth) Login(ctx context.Context, setupKey string, jwtToken string) (err
|
|||||||
log.Debugf("peer registration required")
|
log.Debugf("peer registration required")
|
||||||
_, err = a.registerPeer(client, ctx, setupKey, jwtToken, pubSSHKey)
|
_, err = a.registerPeer(client, ctx, setupKey, jwtToken, pubSSHKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
// The peer pub-key is already on file with the management
|
||||||
|
// server (originally enrolled via SSO) and the session has
|
||||||
|
// expired. The setup-key path can only enrol new peers, so
|
||||||
|
// retrying with -k will keep failing. Replace the raw mgm
|
||||||
|
// message with an actionable hint that tells the user to
|
||||||
|
// re-authenticate via SSO instead.
|
||||||
|
if setupKey != "" && jwtToken == "" && isPeerLoginExpired(err) {
|
||||||
|
err = errSetupKeyOnSSOExpiredPeer
|
||||||
|
}
|
||||||
isAuthError = isPermissionDenied(err)
|
isAuthError = isPermissionDenied(err)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -322,6 +351,7 @@ func (a *Auth) setSystemInfoFlags(info *system.Info) {
|
|||||||
a.config.BlockLANAccess,
|
a.config.BlockLANAccess,
|
||||||
a.config.BlockInbound,
|
a.config.BlockInbound,
|
||||||
a.config.DisableIPv6,
|
a.config.DisableIPv6,
|
||||||
|
a.config.SyncMessageVersion,
|
||||||
a.config.EnableSSHRoot,
|
a.config.EnableSSHRoot,
|
||||||
a.config.EnableSSHSFTP,
|
a.config.EnableSSHSFTP,
|
||||||
a.config.EnableSSHLocalPortForwarding,
|
a.config.EnableSSHLocalPortForwarding,
|
||||||
@@ -473,3 +503,16 @@ func isLoginNeeded(err error) bool {
|
|||||||
func isRegistrationNeeded(err error) bool {
|
func isRegistrationNeeded(err error) bool {
|
||||||
return isPermissionDenied(err)
|
return isPermissionDenied(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isPeerLoginExpired reports whether err is the management server's
|
||||||
|
// "peer login has expired" PermissionDenied response. Used by Login to
|
||||||
|
// detect the case where the caller passed a setup-key but the peer is
|
||||||
|
// actually an SSO-enrolled record whose session needs refreshing — the
|
||||||
|
// setup-key path cannot help there.
|
||||||
|
func isPeerLoginExpired(err error) bool {
|
||||||
|
if !isPermissionDenied(err) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
s, _ := status.FromError(err)
|
||||||
|
return strings.Contains(s.Message(), peerLoginExpiredMsg)
|
||||||
|
}
|
||||||
|
|||||||
80
client/internal/auth/auth_test.go
Normal file
80
client/internal/auth/auth_test.go
Normal file
@@ -0,0 +1,80 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"google.golang.org/grpc/codes"
|
||||||
|
"google.golang.org/grpc/status"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestIsPeerLoginExpired(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
err error
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "nil",
|
||||||
|
err: nil,
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "plain error (not a gRPC status)",
|
||||||
|
err: errors.New("network read: connection reset"),
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "PermissionDenied with different message",
|
||||||
|
err: status.Error(codes.PermissionDenied, "user is blocked"),
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "Unauthenticated with the expected phrase",
|
||||||
|
// Wrong status code — must still return false.
|
||||||
|
err: status.Error(codes.Unauthenticated, "peer login has expired, please log in once more"),
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "exact server message",
|
||||||
|
err: status.Error(codes.PermissionDenied, "peer login has expired, please log in once more"),
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "phrase as substring",
|
||||||
|
// Future-proofing: if mgm reworords but keeps the phrase,
|
||||||
|
// the friendly fallback must still kick in.
|
||||||
|
err: status.Error(codes.PermissionDenied, "session refused: peer login has expired (account=foo)"),
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
if got := isPeerLoginExpired(tc.err); got != tc.want {
|
||||||
|
t.Fatalf("isPeerLoginExpired(%v) = %v, want %v", tc.err, got, tc.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestErrSetupKeyOnSSOExpiredPeer(t *testing.T) {
|
||||||
|
// Sentinel must surface as PermissionDenied so the upstream
|
||||||
|
// isPermissionDenied / isAuthError checks classify it correctly
|
||||||
|
// (short-circuit retry backoff, set StatusNeedsLogin).
|
||||||
|
if !isPermissionDenied(errSetupKeyOnSSOExpiredPeer) {
|
||||||
|
t.Fatalf("errSetupKeyOnSSOExpiredPeer must be a PermissionDenied gRPC error")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Message must actually mention SSO and `netbird up` so it is
|
||||||
|
// actionable for the end user. Loose substring checks keep the
|
||||||
|
// test resilient to copy edits.
|
||||||
|
s, _ := status.FromError(errSetupKeyOnSSOExpiredPeer)
|
||||||
|
msg := strings.ToLower(s.Message())
|
||||||
|
for _, want := range []string{"sso", "netbird up"} {
|
||||||
|
if !strings.Contains(msg, want) {
|
||||||
|
t.Errorf("sentinel message should contain %q, got %q", want, s.Message())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -259,12 +259,18 @@ func (d *DeviceAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowIn
|
|||||||
ticker := time.NewTicker(interval)
|
ticker := time.NewTicker(interval)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
log.Infof("device flow: waiting for user authorization, polling token endpoint every %s, code expires in %s", interval, timeout)
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
polls := 0
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-waitCtx.Done():
|
case <-waitCtx.Done():
|
||||||
return TokenInfo{}, waitCtx.Err()
|
return TokenInfo{}, waitCtx.Err()
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
|
|
||||||
|
polls++
|
||||||
tokenResponse, err := d.requestToken(info)
|
tokenResponse, err := d.requestToken(info)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return TokenInfo{}, fmt.Errorf("parsing token response failed with error: %v", err)
|
return TokenInfo{}, fmt.Errorf("parsing token response failed with error: %v", err)
|
||||||
@@ -272,10 +278,12 @@ func (d *DeviceAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowIn
|
|||||||
|
|
||||||
if tokenResponse.Error != "" {
|
if tokenResponse.Error != "" {
|
||||||
if tokenResponse.Error == "authorization_pending" {
|
if tokenResponse.Error == "authorization_pending" {
|
||||||
|
log.Tracef("device flow: authorization still pending after poll %d", polls)
|
||||||
continue
|
continue
|
||||||
} else if tokenResponse.Error == "slow_down" {
|
} else if tokenResponse.Error == "slow_down" {
|
||||||
interval += (3 * time.Second)
|
interval += (3 * time.Second)
|
||||||
ticker.Reset(interval)
|
ticker.Reset(interval)
|
||||||
|
log.Infof("device flow: IdP requested slow_down, polling interval increased to %s", interval)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -291,11 +299,12 @@ func (d *DeviceAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowIn
|
|||||||
UseIDToken: d.providerConfig.UseIDToken,
|
UseIDToken: d.providerConfig.UseIDToken,
|
||||||
}
|
}
|
||||||
|
|
||||||
err = isValidAccessToken(tokenInfo.GetTokenToUse(), d.providerConfig.Audience)
|
err = validateTokenAudience(tokenInfo.GetTokenToUse(), d.providerConfig.Audience)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return TokenInfo{}, fmt.Errorf("validate access token failed with error: %v", err)
|
return TokenInfo{}, fmt.Errorf("validate access token failed with error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
log.Infof("device flow: user authorization confirmed after %d polls in %s", polls, time.Since(start).Round(time.Second))
|
||||||
return tokenInfo, err
|
return tokenInfo, err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
89
client/internal/auth/pending_flow.go
Normal file
89
client/internal/auth/pending_flow.go
Normal file
@@ -0,0 +1,89 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PendingFlow stores an in-progress OAuth flow between the RPC that
|
||||||
|
// initiates it (returns the verification URI to the UI) and the RPC
|
||||||
|
// that waits for the user to complete it. The flow handle, the
|
||||||
|
// device-code info, and the absolute expiry are kept together so the
|
||||||
|
// waiting RPC can validate the device code and reuse the same flow.
|
||||||
|
//
|
||||||
|
// PendingFlow is safe for concurrent use; callers must not access the
|
||||||
|
// stored fields directly.
|
||||||
|
type PendingFlow struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
flow OAuthFlow
|
||||||
|
info AuthFlowInfo
|
||||||
|
expiresAt time.Time
|
||||||
|
waitCancel context.CancelFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewPendingFlow returns an empty PendingFlow ready to be populated by Set.
|
||||||
|
func NewPendingFlow() *PendingFlow {
|
||||||
|
return &PendingFlow{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set stores the flow and its authorization info, computing the absolute
|
||||||
|
// expiry from info.ExpiresIn (seconds, as returned by the IdP).
|
||||||
|
func (p *PendingFlow) Set(flow OAuthFlow, info AuthFlowInfo) {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
p.flow = flow
|
||||||
|
p.info = info
|
||||||
|
p.expiresAt = time.Now().Add(time.Duration(info.ExpiresIn) * time.Second)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get returns the stored flow, info, and whether a flow is currently
|
||||||
|
// pending. Returns (nil, zero, false) after Clear or before Set.
|
||||||
|
func (p *PendingFlow) Get() (OAuthFlow, AuthFlowInfo, bool) {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
if p.flow == nil {
|
||||||
|
return nil, AuthFlowInfo{}, false
|
||||||
|
}
|
||||||
|
return p.flow, p.info, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExpiresAt returns the absolute expiry of the pending flow. Returns
|
||||||
|
// the zero time when no flow is pending.
|
||||||
|
func (p *PendingFlow) ExpiresAt() time.Time {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
return p.expiresAt
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetWaitCancel records the cancel function for the goroutine currently
|
||||||
|
// blocked in WaitToken so a new RequestAuth can preempt it.
|
||||||
|
func (p *PendingFlow) SetWaitCancel(cancel context.CancelFunc) {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
p.waitCancel = cancel
|
||||||
|
}
|
||||||
|
|
||||||
|
// CancelWait invokes and clears the stored wait-cancel, if any. Safe to
|
||||||
|
// call when no wait is in progress.
|
||||||
|
func (p *PendingFlow) CancelWait() {
|
||||||
|
p.mu.Lock()
|
||||||
|
cancel := p.waitCancel
|
||||||
|
p.waitCancel = nil
|
||||||
|
p.mu.Unlock()
|
||||||
|
if cancel != nil {
|
||||||
|
cancel()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Clear resets the pending flow to empty. Any stored wait-cancel is
|
||||||
|
// dropped without being invoked — call CancelWait first if the waiting
|
||||||
|
// goroutine must be stopped.
|
||||||
|
func (p *PendingFlow) Clear() {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
p.flow = nil
|
||||||
|
p.info = AuthFlowInfo{}
|
||||||
|
p.expiresAt = time.Time{}
|
||||||
|
p.waitCancel = nil
|
||||||
|
}
|
||||||
@@ -188,6 +188,8 @@ func (p *PKCEAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowInfo
|
|||||||
waitCtx, cancel := context.WithTimeout(ctx, timeout)
|
waitCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
|
log.Infof("pkce flow: waiting for authorization callback on %s, timeout %s", p.oAuthConfig.RedirectURL, timeout)
|
||||||
|
|
||||||
tokenChan := make(chan *oauth2.Token, 1)
|
tokenChan := make(chan *oauth2.Token, 1)
|
||||||
errChan := make(chan error, 1)
|
errChan := make(chan error, 1)
|
||||||
|
|
||||||
@@ -221,6 +223,7 @@ func (p *PKCEAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowInfo
|
|||||||
func (p *PKCEAuthorizationFlow) startServer(server *http.Server, tokenChan chan<- *oauth2.Token, errChan chan<- error) {
|
func (p *PKCEAuthorizationFlow) startServer(server *http.Server, tokenChan chan<- *oauth2.Token, errChan chan<- error) {
|
||||||
mux := http.NewServeMux()
|
mux := http.NewServeMux()
|
||||||
mux.HandleFunc("/", func(w http.ResponseWriter, req *http.Request) {
|
mux.HandleFunc("/", func(w http.ResponseWriter, req *http.Request) {
|
||||||
|
log.Infof("pkce flow: received authorization callback from IdP")
|
||||||
cert := p.providerConfig.ClientCertPair
|
cert := p.providerConfig.ClientCertPair
|
||||||
if cert != nil {
|
if cert != nil {
|
||||||
tr := &http.Transport{
|
tr := &http.Transport{
|
||||||
@@ -271,11 +274,18 @@ func (p *PKCEAuthorizationFlow) handleRequest(req *http.Request) (*oauth2.Token,
|
|||||||
return nil, fmt.Errorf("authentication failed: missing code")
|
return nil, fmt.Errorf("authentication failed: missing code")
|
||||||
}
|
}
|
||||||
|
|
||||||
return p.oAuthConfig.Exchange(
|
exchangeStart := time.Now()
|
||||||
|
token, err := p.oAuthConfig.Exchange(
|
||||||
req.Context(),
|
req.Context(),
|
||||||
code,
|
code,
|
||||||
oauth2.SetAuthURLParam("code_verifier", p.codeVerifier),
|
oauth2.SetAuthURLParam("code_verifier", p.codeVerifier),
|
||||||
)
|
)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Infof("pkce flow: authorization code exchanged for token in %s", time.Since(exchangeStart).Round(time.Millisecond))
|
||||||
|
return token, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo, error) {
|
func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo, error) {
|
||||||
@@ -296,7 +306,7 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo,
|
|||||||
audience = p.providerConfig.ClientID
|
audience = p.providerConfig.ClientID
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := isValidAccessToken(tokenInfo.GetTokenToUse(), audience); err != nil {
|
if err := validateTokenAudience(tokenInfo.GetTokenToUse(), audience); err != nil {
|
||||||
return TokenInfo{}, fmt.Errorf("authentication failed: invalid access token - %w", err)
|
return TokenInfo{}, fmt.Errorf("authentication failed: invalid access token - %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -310,6 +320,11 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo,
|
|||||||
return tokenInfo, nil
|
return tokenInfo, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parseEmailFromIDToken extracts the email (or name) claim from an ID token
|
||||||
|
// without verifying its signature. The value is best-effort and used only as a
|
||||||
|
// UX convenience (login hint prefill and display); it never drives an
|
||||||
|
// authorization decision. The authoritative identity is established server-side
|
||||||
|
// from the signature-verified token.
|
||||||
func parseEmailFromIDToken(token string) (string, error) {
|
func parseEmailFromIDToken(token string) (string, error) {
|
||||||
parts := strings.Split(token, ".")
|
parts := strings.Split(token, ".")
|
||||||
if len(parts) < 2 {
|
if len(parts) < 2 {
|
||||||
|
|||||||
82
client/internal/auth/sessionwatch/event.go
Normal file
82
client/internal/auth/sessionwatch/event.go
Normal file
@@ -0,0 +1,82 @@
|
|||||||
|
package sessionwatch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strconv"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// internal event kinds are no longer exposed: the watcher drives the Sink
|
||||||
|
// directly (NotifyStateChange on deadline change/clear, PublishEvent at
|
||||||
|
// each warning lead). Tests use a mock Sink to observe what the watcher
|
||||||
|
// emits.
|
||||||
|
|
||||||
|
// Metadata keys attached by the daemon to session-warning SystemEvents.
|
||||||
|
// The UI tray reads these to build a locale-aware notification without
|
||||||
|
// relying on the daemon's locale-less UserMessage string, and to
|
||||||
|
// disambiguate the T-WarningLead notification from the T-FinalWarningLead
|
||||||
|
// fallback that auto-opens the SessionAboutToExpire dialog.
|
||||||
|
const (
|
||||||
|
// MetaSessionWarning is set to "true" on both warning events (T-10 and
|
||||||
|
// T-2) so the UI can detect a session-warning SystemEvent without
|
||||||
|
// matching on the message text. Use MetaSessionFinal to distinguish
|
||||||
|
// the two.
|
||||||
|
MetaSessionWarning = "session_warning"
|
||||||
|
// MetaSessionFinal is set to "true" on the T-FinalWarningLead event
|
||||||
|
// only. Consumers that need to auto-open the SessionAboutToExpire
|
||||||
|
// dialog gate on this; T-WarningLead events leave the field unset.
|
||||||
|
MetaSessionFinal = "session_final_warning"
|
||||||
|
// MetaSessionExpiresAt carries the absolute UTC deadline encoded with
|
||||||
|
// FormatExpiresAt; consumers must decode with ParseExpiresAt so a
|
||||||
|
// future format change stays a single edit.
|
||||||
|
MetaSessionExpiresAt = "session_expires_at"
|
||||||
|
// MetaSessionLeadMinutes carries the lead in whole minutes (WarningLead
|
||||||
|
// for the T-10 event, FinalWarningLead for the T-2 event) so the UI
|
||||||
|
// can show "expires in ~N minutes" without hardcoding either constant.
|
||||||
|
MetaSessionLeadMinutes = "lead_minutes"
|
||||||
|
// MetaSessionDeadlineRejected is attached to the ERROR/AUTHENTICATION
|
||||||
|
// SystemEvent the daemon emits when it discards a deadline from the
|
||||||
|
// management server (pre-epoch, too far in the future, or past the
|
||||||
|
// clock-skew tolerance). The value is the rejection reason string.
|
||||||
|
// userMessage is left empty; the UI detects the event via this key
|
||||||
|
// and builds a localized notification — same pattern as the session
|
||||||
|
// warnings above.
|
||||||
|
MetaSessionDeadlineRejected = "session_deadline_rejected"
|
||||||
|
)
|
||||||
|
|
||||||
|
// expiresAtLayout is the wire format used for MetaSessionExpiresAt.
|
||||||
|
// Producer and consumers both go through FormatExpiresAt/ParseExpiresAt
|
||||||
|
// so this layout stays a single source of truth.
|
||||||
|
const expiresAtLayout = time.RFC3339
|
||||||
|
|
||||||
|
// FormatExpiresAt encodes a deadline for MetaSessionExpiresAt. Always
|
||||||
|
// emits UTC so a consumer in another timezone reads the same wall-clock
|
||||||
|
// deadline.
|
||||||
|
func FormatExpiresAt(t time.Time) string {
|
||||||
|
return t.UTC().Format(expiresAtLayout)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseExpiresAt decodes the MetaSessionExpiresAt value back to a UTC
|
||||||
|
// time. Returns an error when the field is empty or malformed; the
|
||||||
|
// caller decides whether to fall back (zero value) or propagate.
|
||||||
|
func ParseExpiresAt(s string) (time.Time, error) {
|
||||||
|
t, err := time.Parse(expiresAtLayout, s)
|
||||||
|
if err != nil {
|
||||||
|
return time.Time{}, err
|
||||||
|
}
|
||||||
|
return t.UTC(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// FormatLeadMinutes encodes a lead duration for MetaSessionLeadMinutes
|
||||||
|
// as the integer count of whole minutes. Sub-minute residuals are
|
||||||
|
// truncated — the field is informational ("expires in ~N minutes") and
|
||||||
|
// fractional minutes don't change what the UI displays.
|
||||||
|
func FormatLeadMinutes(d time.Duration) string {
|
||||||
|
return strconv.Itoa(int(d / time.Minute))
|
||||||
|
}
|
||||||
|
|
||||||
|
// ParseLeadMinutes decodes a MetaSessionLeadMinutes value. Returns 0
|
||||||
|
// and the parse error for malformed input; consumers that prefer a
|
||||||
|
// silent fallback can simply ignore the error.
|
||||||
|
func ParseLeadMinutes(s string) (int, error) {
|
||||||
|
return strconv.Atoi(s)
|
||||||
|
}
|
||||||
382
client/internal/auth/sessionwatch/watcher.go
Normal file
382
client/internal/auth/sessionwatch/watcher.go
Normal file
@@ -0,0 +1,382 @@
|
|||||||
|
// Package sessionwatch tracks the SSO session expiry deadline that the
|
||||||
|
// management server publishes via LoginResponse / SyncResponse and fires
|
||||||
|
// two warning events at fixed lead times before expiry: an interactive
|
||||||
|
// T-WarningLead notification and a dismiss-gated T-FinalWarningLead
|
||||||
|
// fallback dialog.
|
||||||
|
//
|
||||||
|
// The watcher is idempotent: Update may be called as often as the network
|
||||||
|
// map snapshots arrive. Repeating the same deadline is a no-op; a new
|
||||||
|
// deadline reschedules the timers and arms a fresh warning cycle.
|
||||||
|
//
|
||||||
|
// Warning firing is edge-detected. Each unique deadline value fires each
|
||||||
|
// warning callback at most once.
|
||||||
|
package sessionwatch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
cProto "github.com/netbirdio/netbird/client/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
maxPastHorizon = 30 * 24 * time.Hour
|
||||||
|
|
||||||
|
// maxDeadlineHorizon caps how far in the future an accepted deadline
|
||||||
|
// can sit. A timestamp beyond this is almost certainly a protocol
|
||||||
|
// glitch, and silently arming a 100-year timer would hide the bug.
|
||||||
|
maxDeadlineHorizon = 10 * 365 * 24 * time.Hour
|
||||||
|
|
||||||
|
// WarningLead is how far before expiry the first (interactive)
|
||||||
|
// warning fires. Drives the T-10 OS notification with
|
||||||
|
// Extend/Dismiss actions.
|
||||||
|
WarningLead = 10 * time.Minute
|
||||||
|
|
||||||
|
// FinalWarningLead is how far before expiry the fallback final
|
||||||
|
// warning fires. Drives the auto-opened SessionAboutToExpire dialog,
|
||||||
|
// but only when the user has not dismissed the T-WarningLead warning
|
||||||
|
// for the same deadline. Must be strictly less than WarningLead.
|
||||||
|
FinalWarningLead = 2 * time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
// ErrDeadlineBeforeEpoch is returned by Update when the supplied
|
||||||
|
// deadline pre-dates 1970-01-01.
|
||||||
|
ErrDeadlineBeforeEpoch = errors.New("session deadline before unix epoch")
|
||||||
|
|
||||||
|
// ErrDeadlineTooFarFuture is returned by Update when the supplied
|
||||||
|
// deadline is more than maxDeadlineHorizon in the future.
|
||||||
|
ErrDeadlineTooFarFuture = errors.New("session deadline too far in the future")
|
||||||
|
|
||||||
|
// ErrDeadlineInPast is returned by Update when the supplied deadline
|
||||||
|
// is more than maxPastHorizon in the past.
|
||||||
|
ErrDeadlineInPast = errors.New("session deadline in the past")
|
||||||
|
)
|
||||||
|
|
||||||
|
// StatusRecorder is the side-effect surface the watcher drives on every
|
||||||
|
// state transition. Production wires this to peer.Status (SetSessionExpiresAt
|
||||||
|
// for deadline change/clear, PublishEvent for the two warnings); tests pass
|
||||||
|
// a fake recorder so the same surface is observable without an engine.
|
||||||
|
//
|
||||||
|
// While the watcher runs, it owns the deadline propagated to the recorder:
|
||||||
|
// every set, clear and sanity-check rejection routes the value through
|
||||||
|
// SetSessionExpiresAt, so the SubscribeStatus snapshot the UI reads can
|
||||||
|
// never drift from the watcher's timer state. (SetSessionExpiresAt fans
|
||||||
|
// out its own state-change notification, so no separate notify is needed.)
|
||||||
|
// The recorder is server-scoped and outlives this engine-scoped watcher;
|
||||||
|
// Close deliberately leaves the recorder value in place so transient engine
|
||||||
|
// restarts don't blank it — the client run loop clears it on real teardown.
|
||||||
|
//
|
||||||
|
// PublishEvent's signature mirrors peer.Status.PublishEvent: the watcher
|
||||||
|
// composes the metadata internally so the wire format (MetaSession*) is
|
||||||
|
// owned by sessionwatch, not the caller.
|
||||||
|
type StatusRecorder interface {
|
||||||
|
SetSessionExpiresAt(deadline time.Time)
|
||||||
|
PublishEvent(
|
||||||
|
severity cProto.SystemEvent_Severity,
|
||||||
|
category cProto.SystemEvent_Category,
|
||||||
|
message string,
|
||||||
|
userMessage string,
|
||||||
|
metadata map[string]string,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Watcher observes the latest session deadline and fires two warnings
|
||||||
|
// before it expires: the interactive T-WarningLead notification, and the
|
||||||
|
// fallback T-FinalWarningLead dialog (suppressed when the user dismissed
|
||||||
|
// the first one for the same deadline). Safe for concurrent use.
|
||||||
|
type Watcher struct {
|
||||||
|
lead time.Duration
|
||||||
|
finalLead time.Duration
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
current time.Time
|
||||||
|
timer *time.Timer
|
||||||
|
finalTimer *time.Timer
|
||||||
|
firedAt time.Time // deadline value the T-WarningLead callback last fired against
|
||||||
|
finalFiredAt time.Time // deadline value the T-FinalWarningLead callback last fired against
|
||||||
|
dismissedAt time.Time // deadline value the user dismissed via Dismiss(); gates fireFinal
|
||||||
|
closed bool
|
||||||
|
recorder StatusRecorder
|
||||||
|
}
|
||||||
|
|
||||||
|
// New returns a watcher with the package defaults WarningLead and
|
||||||
|
// FinalWarningLead. Pass nil for recorder to silence side effects (handy
|
||||||
|
// in unit tests that exercise sanity checks without observing the publish
|
||||||
|
// path).
|
||||||
|
func New(recorder StatusRecorder) *Watcher {
|
||||||
|
return NewWithLeads(WarningLead, FinalWarningLead, recorder)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewWithLeads returns a watcher with custom lead times. Useful for tests.
|
||||||
|
// final must be strictly less than lead; otherwise both timers fire in the
|
||||||
|
// wrong order or simultaneously and the UI flow breaks. A zero final lead
|
||||||
|
// disables the final-warning timer entirely (see armTimerLocked) so a
|
||||||
|
// millisecond-scale deadline doesn't flush both timers in one tick.
|
||||||
|
func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher {
|
||||||
|
return &Watcher{
|
||||||
|
lead: lead,
|
||||||
|
finalLead: final,
|
||||||
|
recorder: recorder,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update sets the latest deadline. Pass the zero time to clear (e.g. when
|
||||||
|
// a Sync push from the server omits the field because login expiration
|
||||||
|
// was disabled).
|
||||||
|
//
|
||||||
|
// Same-value updates are no-ops. A different non-zero value cancels any
|
||||||
|
// pending timer, resets the "already fired" guards, and — when the
|
||||||
|
// deadline lies in the future — arms fresh warning timers. A deadline
|
||||||
|
// already in the past (within maxPastHorizon) is recorded as-is with no
|
||||||
|
// timers: the session has expired and consumers render it that way.
|
||||||
|
//
|
||||||
|
// Returns one of the sentinel Err* values when the deadline fails the
|
||||||
|
// sanity checks (pre-epoch, far future, or past beyond maxPastHorizon).
|
||||||
|
// In every error case the watcher first clears its state so it stays
|
||||||
|
// consistent with what the caller will push into its other sinks (e.g.
|
||||||
|
// applySessionDeadline forces a zero deadline into the status recorder
|
||||||
|
// after a non-nil error).
|
||||||
|
func (w *Watcher) Update(deadline time.Time) error {
|
||||||
|
w.mu.Lock()
|
||||||
|
if w.closed {
|
||||||
|
w.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if deadline.IsZero() {
|
||||||
|
w.clearLocked()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
now := time.Now()
|
||||||
|
switch {
|
||||||
|
case deadline.Before(time.Unix(0, 0)):
|
||||||
|
w.clearLocked()
|
||||||
|
return fmt.Errorf("%w: %v", ErrDeadlineBeforeEpoch, deadline)
|
||||||
|
case deadline.After(now.Add(maxDeadlineHorizon)):
|
||||||
|
w.clearLocked()
|
||||||
|
return fmt.Errorf("%w: %v", ErrDeadlineTooFarFuture, deadline)
|
||||||
|
case deadline.Before(now.Add(-maxPastHorizon)):
|
||||||
|
w.clearLocked()
|
||||||
|
return fmt.Errorf("%w: %v (now=%v)", ErrDeadlineInPast, deadline, now)
|
||||||
|
}
|
||||||
|
|
||||||
|
if deadline.Equal(w.current) {
|
||||||
|
w.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
w.stopTimerLocked()
|
||||||
|
w.current = deadline
|
||||||
|
// Reset every per-deadline guard so a refreshed deadline arms a fresh
|
||||||
|
// warning cycle: both edge triggers and the user Dismiss decision
|
||||||
|
// (the user agreed to the old deadline expiring; a new deadline
|
||||||
|
// restarts the contract).
|
||||||
|
w.firedAt = time.Time{}
|
||||||
|
w.finalFiredAt = time.Time{}
|
||||||
|
w.dismissedAt = time.Time{}
|
||||||
|
|
||||||
|
if deadline.After(now) {
|
||||||
|
w.armTimerLocked(deadline)
|
||||||
|
}
|
||||||
|
recorder := w.recorder
|
||||||
|
w.mu.Unlock()
|
||||||
|
if recorder != nil {
|
||||||
|
recorder.SetSessionExpiresAt(deadline)
|
||||||
|
}
|
||||||
|
log.Infof("auth session deadline set to: %s (in %s)", deadline.Format(time.RFC3339), time.Until(deadline).Round(time.Second))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Deadline returns the most recently observed deadline. Zero when no
|
||||||
|
// deadline is currently tracked.
|
||||||
|
func (w *Watcher) Deadline() time.Time {
|
||||||
|
w.mu.Lock()
|
||||||
|
defer w.mu.Unlock()
|
||||||
|
return w.current
|
||||||
|
}
|
||||||
|
|
||||||
|
// Dismiss records the user's "Dismiss" action against the current deadline
|
||||||
|
// and suppresses the upcoming final-warning callback for that deadline.
|
||||||
|
// Idempotent: repeated calls are no-ops. A subsequent Update with a fresh
|
||||||
|
// deadline resets the dismissal so the final-warning cycle re-arms.
|
||||||
|
//
|
||||||
|
// No-op when the watcher holds no deadline or has been closed.
|
||||||
|
func (w *Watcher) Dismiss() {
|
||||||
|
w.mu.Lock()
|
||||||
|
defer w.mu.Unlock()
|
||||||
|
if w.closed || w.current.IsZero() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if w.dismissedAt.Equal(w.current) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.dismissedAt = w.current
|
||||||
|
// Cancel the armed final-warning timer eagerly. fireFinal would also
|
||||||
|
// gate on dismissedAt, but stopping the timer avoids a wakeup with
|
||||||
|
// nothing to do and makes the intent visible.
|
||||||
|
if w.finalTimer != nil {
|
||||||
|
w.finalTimer.Stop()
|
||||||
|
w.finalTimer = nil
|
||||||
|
}
|
||||||
|
log.Infof("auth session final-warning dismissed for deadline %s", w.current.Format(time.RFC3339))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close stops any pending timer. Update calls after Close are ignored.
|
||||||
|
// The recorder keeps its deadline: the watcher is engine-scoped and closes
|
||||||
|
// on every engine restart (network change, sleep/wake, stream errors)
|
||||||
|
// while the SSO deadline stays valid across those, so clearing here would
|
||||||
|
// blank the UI's "expires in" row on every transient reconnect. The
|
||||||
|
// client run loop clears the server-scoped recorder when it exits for
|
||||||
|
// real (Down, profile switch, permanent login failure).
|
||||||
|
func (w *Watcher) Close() {
|
||||||
|
w.mu.Lock()
|
||||||
|
defer w.mu.Unlock()
|
||||||
|
if w.closed {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.closed = true
|
||||||
|
w.stopTimerLocked()
|
||||||
|
w.current = time.Time{}
|
||||||
|
w.firedAt = time.Time{}
|
||||||
|
w.finalFiredAt = time.Time{}
|
||||||
|
w.dismissedAt = time.Time{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// clearLocked drops the tracked deadline and notifies the recorder so
|
||||||
|
// downstream consumers (SubscribeStatus stream, UI) drop their anchor.
|
||||||
|
// The caller must hold w.mu; this helper releases it before invoking
|
||||||
|
// the recorder.
|
||||||
|
func (w *Watcher) clearLocked() {
|
||||||
|
if w.current.IsZero() {
|
||||||
|
w.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.stopTimerLocked()
|
||||||
|
w.current = time.Time{}
|
||||||
|
w.firedAt = time.Time{}
|
||||||
|
w.finalFiredAt = time.Time{}
|
||||||
|
w.dismissedAt = time.Time{}
|
||||||
|
recorder := w.recorder
|
||||||
|
w.mu.Unlock()
|
||||||
|
if recorder != nil {
|
||||||
|
recorder.SetSessionExpiresAt(time.Time{})
|
||||||
|
}
|
||||||
|
log.Infof("auth session deadline cleared")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Watcher) stopTimerLocked() {
|
||||||
|
if w.timer != nil {
|
||||||
|
w.timer.Stop()
|
||||||
|
w.timer = nil
|
||||||
|
}
|
||||||
|
if w.finalTimer != nil {
|
||||||
|
w.finalTimer.Stop()
|
||||||
|
w.finalTimer = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Watcher) armTimerLocked(deadline time.Time) {
|
||||||
|
w.timer = armOneShotLocked(deadline.Add(-w.lead), func() { w.fire(deadline) })
|
||||||
|
// finalLead <= 0 disables the final-warning timer entirely. Used by
|
||||||
|
// tests that predate the final-warning fallback so a millisecond-scale
|
||||||
|
// deadline does not flush both timers at once.
|
||||||
|
if w.finalLead > 0 {
|
||||||
|
w.finalTimer = armOneShotLocked(deadline.Add(-w.finalLead), func() { w.fireFinal(deadline) })
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *Watcher) fire(armedFor time.Time) {
|
||||||
|
w.mu.Lock()
|
||||||
|
if w.closed || !w.current.Equal(armedFor) {
|
||||||
|
// Deadline moved while we were waiting (e.g. a successful extend).
|
||||||
|
// The reschedule path armed a fresh timer; this one is stale.
|
||||||
|
w.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !w.firedAt.IsZero() && w.firedAt.Equal(armedFor) {
|
||||||
|
w.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.firedAt = armedFor
|
||||||
|
recorder := w.recorder
|
||||||
|
w.mu.Unlock()
|
||||||
|
if recorder == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Infof("auth session expiry soon warning fired")
|
||||||
|
publishWarning(recorder, armedFor, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
// fireFinal mirrors fire for the T-FinalWarningLead timer with an extra
|
||||||
|
// dismiss-gate: if the user dismissed the T-WarningLead notification for
|
||||||
|
// this deadline, the final warning is suppressed entirely.
|
||||||
|
func (w *Watcher) fireFinal(armedFor time.Time) {
|
||||||
|
w.mu.Lock()
|
||||||
|
if w.closed || !w.current.Equal(armedFor) {
|
||||||
|
w.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !w.finalFiredAt.IsZero() && w.finalFiredAt.Equal(armedFor) {
|
||||||
|
w.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if w.dismissedAt.Equal(armedFor) {
|
||||||
|
w.mu.Unlock()
|
||||||
|
log.Infof("auth session final-warning skipped (dismissed by user)")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.finalFiredAt = armedFor
|
||||||
|
recorder := w.recorder
|
||||||
|
w.mu.Unlock()
|
||||||
|
if recorder == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
log.Infof("auth session final-warning fired")
|
||||||
|
publishWarning(recorder, armedFor, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
// armOneShotLocked schedules cb at fireAt. When fireAt is already in the
|
||||||
|
// past it dispatches on the next scheduler tick so a state-change recorder
|
||||||
|
// notification (invoked after w.mu is released) lands first. Caller must
|
||||||
|
// hold w.mu.
|
||||||
|
func armOneShotLocked(fireAt time.Time, cb func()) *time.Timer {
|
||||||
|
delay := time.Until(fireAt)
|
||||||
|
if delay <= 0 {
|
||||||
|
return time.AfterFunc(0, cb)
|
||||||
|
}
|
||||||
|
return time.AfterFunc(delay, cb)
|
||||||
|
}
|
||||||
|
|
||||||
|
// publishWarning composes the SystemEvent for a watcher-fired warning and
|
||||||
|
// pushes it through the recorder. Severity is CRITICAL on both — bypassing
|
||||||
|
// the user's Notifications toggle is deliberate: missing the warning
|
||||||
|
// window forces the post-mortem SessionExpired flow (tunnel torn down,
|
||||||
|
// lock icon, manual re-login), which is the UX we are trying to avoid.
|
||||||
|
func publishWarning(recorder StatusRecorder, deadline time.Time, final bool) {
|
||||||
|
lead := WarningLead
|
||||||
|
message := "session expiry warning"
|
||||||
|
meta := map[string]string{
|
||||||
|
MetaSessionWarning: "true",
|
||||||
|
MetaSessionExpiresAt: FormatExpiresAt(deadline),
|
||||||
|
}
|
||||||
|
if final {
|
||||||
|
lead = FinalWarningLead
|
||||||
|
message = "session expiry final warning"
|
||||||
|
meta[MetaSessionFinal] = "true"
|
||||||
|
}
|
||||||
|
meta[MetaSessionLeadMinutes] = FormatLeadMinutes(lead)
|
||||||
|
|
||||||
|
recorder.PublishEvent(
|
||||||
|
cProto.SystemEvent_CRITICAL,
|
||||||
|
cProto.SystemEvent_AUTHENTICATION,
|
||||||
|
message,
|
||||||
|
"",
|
||||||
|
meta,
|
||||||
|
)
|
||||||
|
}
|
||||||
529
client/internal/auth/sessionwatch/watcher_test.go
Normal file
529
client/internal/auth/sessionwatch/watcher_test.go
Normal file
@@ -0,0 +1,529 @@
|
|||||||
|
package sessionwatch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
cProto "github.com/netbirdio/netbird/client/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fakeRecorder satisfies StatusRecorder and records every call so tests
|
||||||
|
// can observe what the watcher emits. SetSessionExpiresAt and PublishEvent
|
||||||
|
// land in the same ordered events slice (with the Kind distinguishing
|
||||||
|
// them) so tests that care about ordering still work. lastDeadline holds
|
||||||
|
// the most recent value passed to SetSessionExpiresAt so tests can assert
|
||||||
|
// the recorder ended up cleared/set as expected.
|
||||||
|
type fakeRecorder struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
events []event
|
||||||
|
lastDeadline time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type eventKind int
|
||||||
|
|
||||||
|
const (
|
||||||
|
stateChange eventKind = iota
|
||||||
|
publish
|
||||||
|
)
|
||||||
|
|
||||||
|
type event struct {
|
||||||
|
kind eventKind
|
||||||
|
// Set only for publish events.
|
||||||
|
severity cProto.SystemEvent_Severity
|
||||||
|
category cProto.SystemEvent_Category
|
||||||
|
message string
|
||||||
|
meta map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetSessionExpiresAt mirrors peer.Status: a same-value write is a no-op,
|
||||||
|
// a real change records the new value and fans out a state-change (the
|
||||||
|
// production recorder calls notifyStateChange internally). The baseline
|
||||||
|
// is the zero time, so an initial clear before any deadline is set emits
|
||||||
|
// nothing — matching the real recorder.
|
||||||
|
func (r *fakeRecorder) SetSessionExpiresAt(deadline time.Time) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
if r.lastDeadline.Equal(deadline) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r.lastDeadline = deadline
|
||||||
|
r.events = append(r.events, event{kind: stateChange})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *fakeRecorder) deadline() time.Time {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
return r.lastDeadline
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *fakeRecorder) PublishEvent(
|
||||||
|
severity cProto.SystemEvent_Severity,
|
||||||
|
category cProto.SystemEvent_Category,
|
||||||
|
message string,
|
||||||
|
_ string,
|
||||||
|
metadata map[string]string,
|
||||||
|
) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.events = append(r.events, event{
|
||||||
|
kind: publish,
|
||||||
|
severity: severity,
|
||||||
|
category: category,
|
||||||
|
message: message,
|
||||||
|
meta: metadata,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *fakeRecorder) snapshot() []event {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
out := make([]event, len(r.events))
|
||||||
|
copy(out, r.events)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e event) isFinalWarning() bool {
|
||||||
|
return e.kind == publish && e.meta[MetaSessionFinal] == "true"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e event) isWarning() bool {
|
||||||
|
return e.kind == publish && e.meta[MetaSessionWarning] == "true" && e.meta[MetaSessionFinal] != "true"
|
||||||
|
}
|
||||||
|
|
||||||
|
func countWhere(events []event, pred func(event) bool) int {
|
||||||
|
n := 0
|
||||||
|
for _, e := range events {
|
||||||
|
if pred(e) {
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitForEvents(t *testing.T, r *fakeRecorder, want int) []event {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(500 * time.Millisecond)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if got := r.snapshot(); len(got) >= want {
|
||||||
|
return got
|
||||||
|
}
|
||||||
|
time.Sleep(5 * time.Millisecond)
|
||||||
|
}
|
||||||
|
got := r.snapshot()
|
||||||
|
t.Fatalf("timed out waiting for %d events, got %d: %+v", want, len(got), got)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// newWatcher builds a watcher with the final timer disabled (finalLead=0),
|
||||||
|
// matching the lead-only behaviour the pre-final-warning tests assume.
|
||||||
|
func newWatcher(lead time.Duration, r *fakeRecorder) *Watcher {
|
||||||
|
return NewWithLeads(lead, 0, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateZeroBeforeAnythingIsNoop(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
_ = w.Update(time.Time{})
|
||||||
|
|
||||||
|
if got := r.snapshot(); len(got) != 0 {
|
||||||
|
t.Fatalf("expected no events on initial zero, got %+v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateNonZeroFiresStateChange(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
d := time.Now().Add(time.Hour)
|
||||||
|
_ = w.Update(d)
|
||||||
|
|
||||||
|
events := waitForEvents(t, r, 1)
|
||||||
|
if events[0].kind != stateChange {
|
||||||
|
t.Fatalf("expected stateChange, got %+v", events[0])
|
||||||
|
}
|
||||||
|
if !w.Deadline().Equal(d) {
|
||||||
|
t.Fatalf("deadline mismatch: %v vs %v", w.Deadline(), d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSameDeadlineIsNoop(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
d := time.Now().Add(time.Hour)
|
||||||
|
_ = w.Update(d)
|
||||||
|
_ = w.Update(d)
|
||||||
|
_ = w.Update(d)
|
||||||
|
|
||||||
|
events := waitForEvents(t, r, 1)
|
||||||
|
if len(events) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 event for repeated same deadline, got %d: %+v", len(events), events)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWarningFiresOnceWithinLeadWindow(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
lead := 50 * time.Millisecond
|
||||||
|
w := newWatcher(lead, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
// Deadline 80ms out — warning should fire after ~30ms.
|
||||||
|
d := time.Now().Add(80 * time.Millisecond)
|
||||||
|
_ = w.Update(d)
|
||||||
|
|
||||||
|
events := waitForEvents(t, r, 2)
|
||||||
|
if events[0].kind != stateChange {
|
||||||
|
t.Fatalf("event[0] should be stateChange, got %+v", events[0])
|
||||||
|
}
|
||||||
|
if !events[1].isWarning() {
|
||||||
|
t.Fatalf("event[1] should be a warning publish, got %+v", events[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWarningFiresImmediatelyWhenAlreadyInsideWindow(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(time.Hour, r) // lead > delta => fire immediately
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
d := time.Now().Add(10 * time.Millisecond)
|
||||||
|
_ = w.Update(d)
|
||||||
|
|
||||||
|
events := waitForEvents(t, r, 2)
|
||||||
|
if !events[1].isWarning() {
|
||||||
|
t.Fatalf("expected immediate warning publish, got %+v", events[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewDeadlineCancelsPriorTimer(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
lead := 50 * time.Millisecond
|
||||||
|
w := newWatcher(lead, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
first := time.Now().Add(80 * time.Millisecond) // would fire warning ~30ms in
|
||||||
|
_ = w.Update(first)
|
||||||
|
|
||||||
|
// Replace with a far-future deadline before the warning fires.
|
||||||
|
time.Sleep(5 * time.Millisecond)
|
||||||
|
second := time.Now().Add(time.Hour)
|
||||||
|
_ = w.Update(second)
|
||||||
|
|
||||||
|
// Wait past when first's warning would have fired.
|
||||||
|
time.Sleep(80 * time.Millisecond)
|
||||||
|
|
||||||
|
if n := countWhere(r.snapshot(), event.isWarning); n != 0 {
|
||||||
|
t.Fatalf("warning fired for cancelled deadline: %+v", r.snapshot())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRefreshAfterFireArmsNewWarning(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
lead := 150 * time.Millisecond
|
||||||
|
w := newWatcher(lead, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
// Warning fires ~20ms in; the deadline itself stays 150ms away so the
|
||||||
|
// replacement below lands well before it.
|
||||||
|
first := time.Now().Add(170 * time.Millisecond)
|
||||||
|
_ = w.Update(first)
|
||||||
|
|
||||||
|
// Wait for stateChange + warning of the first cycle.
|
||||||
|
waitForEvents(t, r, 2)
|
||||||
|
|
||||||
|
// Simulate a successful extend: brand new deadline.
|
||||||
|
second := time.Now().Add(60 * time.Millisecond)
|
||||||
|
_ = w.Update(second)
|
||||||
|
|
||||||
|
// 4 events total: stateChange, warning (first), stateChange, warning (second).
|
||||||
|
events := waitForEvents(t, r, 4)
|
||||||
|
if events[2].kind != stateChange {
|
||||||
|
t.Fatalf("event[2] should be stateChange for the new deadline, got %+v", events[2])
|
||||||
|
}
|
||||||
|
if !events[3].isWarning() {
|
||||||
|
t.Fatalf("event[3] should be a warning publish for the new deadline, got %+v", events[3])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateZeroAfterNonZeroClearsState(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(time.Hour, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
d := time.Now().Add(2 * time.Hour)
|
||||||
|
_ = w.Update(d)
|
||||||
|
waitForEvents(t, r, 1)
|
||||||
|
|
||||||
|
_ = w.Update(time.Time{})
|
||||||
|
|
||||||
|
events := waitForEvents(t, r, 2)
|
||||||
|
if events[1].kind != stateChange {
|
||||||
|
t.Fatalf("expected stateChange on clear, got %+v", events[1])
|
||||||
|
}
|
||||||
|
if !w.Deadline().IsZero() {
|
||||||
|
t.Fatalf("Deadline should be zero after clear")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateRejectsBeforeEpoch(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
good := time.Now().Add(time.Hour)
|
||||||
|
if err := w.Update(good); err != nil {
|
||||||
|
t.Fatalf("seed Update: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err := w.Update(time.Unix(-100, 0))
|
||||||
|
if !errors.Is(err, ErrDeadlineBeforeEpoch) {
|
||||||
|
t.Fatalf("want ErrDeadlineBeforeEpoch, got %v", err)
|
||||||
|
}
|
||||||
|
if !w.Deadline().IsZero() {
|
||||||
|
t.Fatalf("rejected pre-epoch update must clear deadline; got %v", w.Deadline())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateRejectsTooFarFuture(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
good := time.Now().Add(time.Hour)
|
||||||
|
if err := w.Update(good); err != nil {
|
||||||
|
t.Fatalf("seed Update: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err := w.Update(time.Now().Add(50 * 365 * 24 * time.Hour))
|
||||||
|
if !errors.Is(err, ErrDeadlineTooFarFuture) {
|
||||||
|
t.Fatalf("want ErrDeadlineTooFarFuture, got %v", err)
|
||||||
|
}
|
||||||
|
if !w.Deadline().IsZero() {
|
||||||
|
t.Fatalf("rejected far-future update must clear deadline; got %v", w.Deadline())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateRecentPastRecordedAsExpired(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
d := time.Now().Add(-1 * time.Hour)
|
||||||
|
if err := w.Update(d); err != nil {
|
||||||
|
t.Fatalf("recent-past Update should succeed, got %v", err)
|
||||||
|
}
|
||||||
|
if !w.Deadline().Equal(d) {
|
||||||
|
t.Fatalf("expected deadline to be recorded, got %v want %v", w.Deadline(), d)
|
||||||
|
}
|
||||||
|
if got := r.deadline(); !got.Equal(d) {
|
||||||
|
t.Fatalf("recorder deadline = %v, want %v", got, d)
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(80 * time.Millisecond)
|
||||||
|
if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 {
|
||||||
|
t.Fatalf("no warning events may fire for an already-past deadline, got %+v", r.snapshot())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateAncientPastRejected(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
good := time.Now().Add(time.Hour)
|
||||||
|
if err := w.Update(good); err != nil {
|
||||||
|
t.Fatalf("seed Update: %v", err)
|
||||||
|
}
|
||||||
|
// Drain the stateChange from the seed.
|
||||||
|
waitForEvents(t, r, 1)
|
||||||
|
|
||||||
|
err := w.Update(time.Now().Add(-31 * 24 * time.Hour))
|
||||||
|
if !errors.Is(err, ErrDeadlineInPast) {
|
||||||
|
t.Fatalf("want ErrDeadlineInPast, got %v", err)
|
||||||
|
}
|
||||||
|
if !w.Deadline().IsZero() {
|
||||||
|
t.Fatalf("rejected ancient-past update must clear the deadline, got %v", w.Deadline())
|
||||||
|
}
|
||||||
|
events := waitForEvents(t, r, 2)
|
||||||
|
if events[1].kind != stateChange {
|
||||||
|
t.Fatalf("expected stateChange on clear, got %+v", events[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCloseSilencesUpdates(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
|
w.Close()
|
||||||
|
|
||||||
|
if err := w.Update(time.Now().Add(time.Hour)); err != nil {
|
||||||
|
t.Fatalf("Update after Close: want nil, got %v", err)
|
||||||
|
}
|
||||||
|
if got := r.snapshot(); len(got) != 0 {
|
||||||
|
t.Fatalf("expected no events after Close, got %+v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCloseKeepsRecorderDeadline pins the reconnect-flap fix: the watcher
|
||||||
|
// closes on every engine restart (network change, sleep/wake) while the
|
||||||
|
// SSO deadline stays valid across those, so Close must leave the
|
||||||
|
// server-scoped recorder's value in place. The client run loop clears the
|
||||||
|
// recorder when it exits for real.
|
||||||
|
func TestCloseKeepsRecorderDeadline(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(time.Hour, r)
|
||||||
|
|
||||||
|
d := time.Now().Add(2 * time.Hour)
|
||||||
|
if err := w.Update(d); err != nil {
|
||||||
|
t.Fatalf("seed Update: %v", err)
|
||||||
|
}
|
||||||
|
if got := r.deadline(); !got.Equal(d) {
|
||||||
|
t.Fatalf("recorder deadline after Update = %v, want %v", got, d)
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Close()
|
||||||
|
|
||||||
|
if got := r.deadline(); !got.Equal(d) {
|
||||||
|
t.Fatalf("recorder deadline after Close = %v, want %v", got, d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCloseWithoutDeadlineLeavesRecorderUntouched guards the symmetric
|
||||||
|
// case: closing a watcher that never held a deadline must not emit a
|
||||||
|
// redundant clear (the recorder may legitimately hold a value written by
|
||||||
|
// some other path; the watcher only owns what it set).
|
||||||
|
func TestCloseWithoutDeadlineLeavesRecorderUntouched(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := newWatcher(time.Hour, r)
|
||||||
|
|
||||||
|
w.Close()
|
||||||
|
|
||||||
|
if got := r.snapshot(); len(got) != 0 {
|
||||||
|
t.Fatalf("expected no events from Close on an empty watcher, got %+v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFinalWarningFiresAfterRegularWarning(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
// Warning fires at deadline-80ms, final at deadline-30ms.
|
||||||
|
w := NewWithLeads(80*time.Millisecond, 30*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
d := time.Now().Add(100 * time.Millisecond)
|
||||||
|
_ = w.Update(d)
|
||||||
|
|
||||||
|
// Expect stateChange + warning + final-warning.
|
||||||
|
events := waitForEvents(t, r, 3)
|
||||||
|
|
||||||
|
if countWhere(events, func(e event) bool { return e.kind == stateChange }) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 stateChange, got %+v", events)
|
||||||
|
}
|
||||||
|
if countWhere(events, event.isWarning) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 warning publish, got %+v", events)
|
||||||
|
}
|
||||||
|
if countWhere(events, event.isFinalWarning) != 1 {
|
||||||
|
t.Fatalf("expected exactly 1 final-warning publish, got %+v", events)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warning must precede final (same deadline, longer lead fires first).
|
||||||
|
var wIdx, fIdx int
|
||||||
|
for i, e := range events {
|
||||||
|
switch {
|
||||||
|
case e.isWarning():
|
||||||
|
wIdx = i
|
||||||
|
case e.isFinalWarning():
|
||||||
|
fIdx = i
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if wIdx > fIdx {
|
||||||
|
t.Fatalf("warning must publish before final-warning, got order %+v", events)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDismissSuppressesFinalWarning(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := NewWithLeads(80*time.Millisecond, 30*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
d := time.Now().Add(100 * time.Millisecond)
|
||||||
|
_ = w.Update(d)
|
||||||
|
|
||||||
|
// Wait for the warning publish so we know we're inside the warning
|
||||||
|
// window, then dismiss before the final timer would fire.
|
||||||
|
deadline := time.Now().Add(500 * time.Millisecond)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if countWhere(r.snapshot(), event.isWarning) >= 1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
time.Sleep(2 * time.Millisecond)
|
||||||
|
}
|
||||||
|
if countWhere(r.snapshot(), event.isWarning) < 1 {
|
||||||
|
t.Fatalf("warning did not publish in time, events=%+v", r.snapshot())
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Dismiss()
|
||||||
|
|
||||||
|
// Now wait past when the final would have fired.
|
||||||
|
time.Sleep(120 * time.Millisecond)
|
||||||
|
|
||||||
|
if n := countWhere(r.snapshot(), event.isFinalWarning); n != 0 {
|
||||||
|
t.Fatalf("final-warning published after Dismiss(), events=%+v", r.snapshot())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDismissResetByNewDeadline(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := NewWithLeads(80*time.Millisecond, 30*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
first := time.Now().Add(100 * time.Millisecond)
|
||||||
|
_ = w.Update(first)
|
||||||
|
|
||||||
|
// Dismiss against the first deadline.
|
||||||
|
w.Dismiss()
|
||||||
|
|
||||||
|
// Replace with a fresh deadline before the first's timers complete.
|
||||||
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
second := time.Now().Add(100 * time.Millisecond)
|
||||||
|
_ = w.Update(second)
|
||||||
|
|
||||||
|
// The second cycle must publish a final-warning (the dismiss state
|
||||||
|
// did not carry over).
|
||||||
|
deadline := time.Now().Add(500 * time.Millisecond)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if countWhere(r.snapshot(), event.isFinalWarning) >= 1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
time.Sleep(5 * time.Millisecond)
|
||||||
|
}
|
||||||
|
if countWhere(r.snapshot(), event.isFinalWarning) < 1 {
|
||||||
|
t.Fatalf("final-warning did not publish on fresh deadline after Dismiss reset, events=%+v", r.snapshot())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDismissBeforeUpdateIsNoop(t *testing.T) {
|
||||||
|
r := &fakeRecorder{}
|
||||||
|
w := NewWithLeads(80*time.Millisecond, 30*time.Millisecond, r)
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
// No deadline tracked yet; Dismiss must be a no-op (no panic, no state).
|
||||||
|
w.Dismiss()
|
||||||
|
|
||||||
|
d := time.Now().Add(100 * time.Millisecond)
|
||||||
|
_ = w.Update(d)
|
||||||
|
|
||||||
|
// Final warning should still publish — Dismiss only acts on the current
|
||||||
|
// deadline, and there was none at the time of the call.
|
||||||
|
deadline := time.Now().Add(500 * time.Millisecond)
|
||||||
|
for time.Now().Before(deadline) {
|
||||||
|
if countWhere(r.snapshot(), event.isFinalWarning) >= 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(5 * time.Millisecond)
|
||||||
|
}
|
||||||
|
t.Fatalf("final-warning did not publish after no-op pre-Update Dismiss, events=%+v", r.snapshot())
|
||||||
|
}
|
||||||
@@ -20,14 +20,26 @@ func randomBytesInHex(count int) (string, error) {
|
|||||||
return hex.EncodeToString(buf), nil
|
return hex.EncodeToString(buf), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// isValidAccessToken is a simple validation of the access token
|
// validateTokenAudience checks that the token is a well-formed JWT whose
|
||||||
func isValidAccessToken(token string, audience string) error {
|
// audience claim matches the expected audience.
|
||||||
|
//
|
||||||
|
// It does NOT verify the token's cryptographic signature and therefore must not
|
||||||
|
// be treated as an authenticity check. The token is obtained by the client
|
||||||
|
// directly from the IdP token endpoint over TLS, and its signature is verified
|
||||||
|
// server-side by the management server against the IdP's JWKS
|
||||||
|
// (see shared/auth/jwt/validator.go). This function is only a client-side
|
||||||
|
// sanity check that the returned token targets the expected audience.
|
||||||
|
func validateTokenAudience(token string, audience string) error {
|
||||||
if token == "" {
|
if token == "" {
|
||||||
return fmt.Errorf("token received is empty")
|
return fmt.Errorf("token received is empty")
|
||||||
}
|
}
|
||||||
|
|
||||||
encodedClaims := strings.Split(token, ".")[1]
|
parts := strings.Split(token, ".")
|
||||||
claimsString, err := base64.RawURLEncoding.DecodeString(encodedClaims)
|
if len(parts) != 3 {
|
||||||
|
return fmt.Errorf("token is not a well-formed JWT")
|
||||||
|
}
|
||||||
|
|
||||||
|
claimsString, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
108
client/internal/auth/util_test.go
Normal file
108
client/internal/auth/util_test.go
Normal file
@@ -0,0 +1,108 @@
|
|||||||
|
package auth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// makeJWT builds an unsigned JWT-shaped string (header.payload.signature) with
|
||||||
|
// the given claims payload. The signature part is arbitrary because
|
||||||
|
// validateTokenAudience intentionally does not verify it.
|
||||||
|
func makeJWT(t *testing.T, claims map[string]interface{}) string {
|
||||||
|
t.Helper()
|
||||||
|
header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","typ":"JWT"}`))
|
||||||
|
payloadBytes, err := json.Marshal(claims)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marshal claims: %v", err)
|
||||||
|
}
|
||||||
|
payload := base64.RawURLEncoding.EncodeToString(payloadBytes)
|
||||||
|
return header + "." + payload + ".unverified-signature"
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateTokenAudience(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
token string
|
||||||
|
audience string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "empty token",
|
||||||
|
token: "",
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "not a JWT - no dots",
|
||||||
|
token: "notajwt",
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "not a JWT - two parts only",
|
||||||
|
token: "header.payload",
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "matching string audience",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"aud": "netbird"}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mismatching string audience",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"aud": "other"}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "matching audience in array",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"aud": []interface{}{"other", "netbird"}}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mismatching audience array",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"aud": []interface{}{"a", "b"}}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing audience claim",
|
||||||
|
token: makeJWT(t, map[string]interface{}{"sub": "user"}),
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "invalid base64 payload",
|
||||||
|
token: "header.!!!not-base64!!!.sig",
|
||||||
|
audience: "netbird",
|
||||||
|
wantErr: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
err := validateTokenAudience(tc.token, tc.audience)
|
||||||
|
if tc.wantErr && err == nil {
|
||||||
|
t.Fatalf("expected error, got nil")
|
||||||
|
}
|
||||||
|
if !tc.wantErr && err != nil {
|
||||||
|
t.Fatalf("expected no error, got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestValidateTokenAudienceNoPanic guards the regression where a non-empty
|
||||||
|
// token without the JWT dot structure caused an index-out-of-range panic.
|
||||||
|
func TestValidateTokenAudienceNoPanic(t *testing.T) {
|
||||||
|
inputs := []string{"a", ".", "a.", "aaaa", "no-dots-here"}
|
||||||
|
for _, in := range inputs {
|
||||||
|
if err := validateTokenAudience(in, "netbird"); err == nil {
|
||||||
|
t.Fatalf("expected error for malformed token %q, got nil", in)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -34,6 +34,8 @@ const (
|
|||||||
// - Handling connection establishment based on peer signaling
|
// - Handling connection establishment based on peer signaling
|
||||||
//
|
//
|
||||||
// The implementation is not thread-safe; it is protected by engine.syncMsgMux.
|
// The implementation is not thread-safe; it is protected by engine.syncMsgMux.
|
||||||
|
// The only exception is ActivatePeer, which is safe for concurrent use so the
|
||||||
|
// DNS warm-up path can call it without contending on the engine mutex.
|
||||||
type ConnMgr struct {
|
type ConnMgr struct {
|
||||||
peerStore *peerstore.Store
|
peerStore *peerstore.Store
|
||||||
statusRecorder *peer.Status
|
statusRecorder *peer.Status
|
||||||
@@ -42,12 +44,26 @@ type ConnMgr struct {
|
|||||||
rosenpassEnabled bool
|
rosenpassEnabled bool
|
||||||
|
|
||||||
lazyConnMgr *manager.Manager
|
lazyConnMgr *manager.Manager
|
||||||
|
// lazyConnMgrMu guards the lazyConnMgr pointer for readers outside the
|
||||||
|
// engine loop (ActivatePeer). Writers hold it in addition to
|
||||||
|
// engine.syncMsgMux; all other reads stay under engine.syncMsgMux only.
|
||||||
|
lazyConnMgrMu sync.RWMutex
|
||||||
|
|
||||||
|
// reconcileRoutedIPs re-applies a peer's routed allowed IPs after its lazy wake endpoint is
|
||||||
|
// (re)armed (Mode A at arm time). Injected by the engine; nil disables the reconcile.
|
||||||
|
reconcileRoutedIPs func(peerKey string) error
|
||||||
|
|
||||||
wg sync.WaitGroup
|
wg sync.WaitGroup
|
||||||
lazyCtx context.Context
|
lazyCtx context.Context
|
||||||
lazyCtxCancel context.CancelFunc
|
lazyCtxCancel context.CancelFunc
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetRoutedIPsReconciler injects the callback used to re-apply a peer's routed allowed IPs when
|
||||||
|
// its lazy wake endpoint is (re)armed. Must be called before the lazy manager starts.
|
||||||
|
func (e *ConnMgr) SetRoutedIPsReconciler(fn func(peerKey string) error) {
|
||||||
|
e.reconcileRoutedIPs = fn
|
||||||
|
}
|
||||||
|
|
||||||
func NewConnMgr(engineConfig *EngineConfig, statusRecorder *peer.Status, peerStore *peerstore.Store, iface lazyconn.WGIface) *ConnMgr {
|
func NewConnMgr(engineConfig *EngineConfig, statusRecorder *peer.Status, peerStore *peerstore.Store, iface lazyconn.WGIface) *ConnMgr {
|
||||||
e := &ConnMgr{
|
e := &ConnMgr{
|
||||||
peerStore: peerStore,
|
peerStore: peerStore,
|
||||||
@@ -109,7 +125,7 @@ func (e *ConnMgr) UpdatedRemoteFeatureFlag(ctx context.Context, enabled bool) er
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
log.Warnf("lazy connection manager is enabled by management feature flag")
|
log.Infof("lazy connection manager is enabled by the management feature flag")
|
||||||
e.initLazyManager(ctx)
|
e.initLazyManager(ctx)
|
||||||
e.statusRecorder.UpdateLazyConnection(true)
|
e.statusRecorder.UpdateLazyConnection(true)
|
||||||
return e.addPeersToLazyConnManager()
|
return e.addPeersToLazyConnManager()
|
||||||
@@ -238,12 +254,20 @@ func (e *ConnMgr) RemovePeerConn(peerKey string) {
|
|||||||
conn.Log.Infof("removed peer from lazy conn manager")
|
conn.Log.Infof("removed peer from lazy conn manager")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ActivatePeer wakes an idle lazy connection. Unlike the rest of ConnMgr it is
|
||||||
|
// safe for concurrent use: the lazy manager pointer is read under lazyConnMgrMu
|
||||||
|
// and the manager itself is internally synchronized, so callers outside the
|
||||||
|
// engine loop (DNS warm-up) do not need engine.syncMsgMux.
|
||||||
func (e *ConnMgr) ActivatePeer(ctx context.Context, conn *peer.Conn) {
|
func (e *ConnMgr) ActivatePeer(ctx context.Context, conn *peer.Conn) {
|
||||||
if !e.isStartedWithLazyMgr() {
|
e.lazyConnMgrMu.RLock()
|
||||||
|
lazyConnMgr := e.lazyConnMgr
|
||||||
|
started := lazyConnMgr != nil && e.lazyCtxCancel != nil
|
||||||
|
e.lazyConnMgrMu.RUnlock()
|
||||||
|
if !started {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if found := e.lazyConnMgr.ActivatePeer(conn.GetKey()); found {
|
if found := lazyConnMgr.ActivatePeer(conn.GetKey()); found {
|
||||||
if err := conn.Open(ctx); err != nil {
|
if err := conn.Open(ctx); err != nil {
|
||||||
conn.Log.Errorf("failed to open connection: %v", err)
|
conn.Log.Errorf("failed to open connection: %v", err)
|
||||||
}
|
}
|
||||||
@@ -268,16 +292,22 @@ func (e *ConnMgr) Close() {
|
|||||||
|
|
||||||
e.lazyCtxCancel()
|
e.lazyCtxCancel()
|
||||||
e.wg.Wait()
|
e.wg.Wait()
|
||||||
|
|
||||||
|
e.lazyConnMgrMu.Lock()
|
||||||
e.lazyConnMgr = nil
|
e.lazyConnMgr = nil
|
||||||
|
e.lazyConnMgrMu.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (e *ConnMgr) initLazyManager(engineCtx context.Context) {
|
func (e *ConnMgr) initLazyManager(engineCtx context.Context) {
|
||||||
cfg := manager.Config{
|
cfg := manager.Config{
|
||||||
InactivityThreshold: inactivityThresholdEnv(),
|
InactivityThreshold: inactivityThresholdEnv(),
|
||||||
|
ReconcileAllowedIPs: e.reconcileRoutedIPs,
|
||||||
}
|
}
|
||||||
e.lazyConnMgr = manager.NewManager(cfg, engineCtx, e.peerStore, e.iface)
|
|
||||||
|
|
||||||
|
e.lazyConnMgrMu.Lock()
|
||||||
|
e.lazyConnMgr = manager.NewManager(cfg, engineCtx, e.peerStore, e.iface)
|
||||||
e.lazyCtx, e.lazyCtxCancel = context.WithCancel(engineCtx)
|
e.lazyCtx, e.lazyCtxCancel = context.WithCancel(engineCtx)
|
||||||
|
e.lazyConnMgrMu.Unlock()
|
||||||
|
|
||||||
e.wg.Add(1)
|
e.wg.Add(1)
|
||||||
go func() {
|
go func() {
|
||||||
@@ -316,7 +346,10 @@ func (e *ConnMgr) closeManager(ctx context.Context) {
|
|||||||
|
|
||||||
e.lazyCtxCancel()
|
e.lazyCtxCancel()
|
||||||
e.wg.Wait()
|
e.wg.Wait()
|
||||||
|
|
||||||
|
e.lazyConnMgrMu.Lock()
|
||||||
e.lazyConnMgr = nil
|
e.lazyConnMgr = nil
|
||||||
|
e.lazyConnMgrMu.Unlock()
|
||||||
|
|
||||||
for _, peerID := range e.peerStore.PeersPubKey() {
|
for _, peerID := range e.peerStore.PeersPubKey() {
|
||||||
e.peerStore.PeerConnOpen(ctx, peerID)
|
e.peerStore.PeerConnOpen(ctx, peerID)
|
||||||
@@ -352,11 +385,20 @@ func inactivityThresholdEnv() *time.Duration {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
parsedMinutes, err := strconv.Atoi(envValue)
|
// Documented format: a Go duration such as "30m" or "1h".
|
||||||
if err != nil || parsedMinutes <= 0 {
|
if d, err := time.ParseDuration(envValue); err == nil {
|
||||||
return nil
|
if d <= 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return &d
|
||||||
}
|
}
|
||||||
|
|
||||||
d := time.Duration(parsedMinutes) * time.Minute
|
// Backwards compatibility: a bare integer used to be interpreted as minutes.
|
||||||
return &d
|
if parsedMinutes, err := strconv.Atoi(envValue); err == nil && parsedMinutes > 0 {
|
||||||
|
d := time.Duration(parsedMinutes) * time.Minute
|
||||||
|
return &d
|
||||||
|
}
|
||||||
|
|
||||||
|
log.Warnf("invalid %s value %q: expected a Go duration such as 30m or 1h", lazyconn.EnvInactivityThreshold, envValue)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,10 +1,21 @@
|
|||||||
package internal
|
package internal
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||||
"github.com/netbirdio/netbird/client/internal/lazyconn"
|
"github.com/netbirdio/netbird/client/internal/lazyconn"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peerstore"
|
||||||
|
"github.com/netbirdio/netbird/monotime"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestResolveLazyForce(t *testing.T) {
|
func TestResolveLazyForce(t *testing.T) {
|
||||||
@@ -38,3 +49,93 @@ func TestResolveLazyForce(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type mockLazyWGIface struct{}
|
||||||
|
|
||||||
|
func (mockLazyWGIface) RemovePeer(string) error { return nil }
|
||||||
|
func (mockLazyWGIface) UpdatePeer(string, []netip.Prefix, time.Duration, *net.UDPAddr, *wgtypes.Key) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (mockLazyWGIface) IsUserspaceBind() bool { return false }
|
||||||
|
func (mockLazyWGIface) Address() wgaddr.Address { return wgaddr.Address{} }
|
||||||
|
func (mockLazyWGIface) LastActivities() map[string]monotime.Time { return nil }
|
||||||
|
func (mockLazyWGIface) MTU() uint16 { return 1280 }
|
||||||
|
|
||||||
|
// TestConnMgr_ActivatePeerConcurrentWithLifecycle exercises ActivatePeer from
|
||||||
|
// non-engine goroutines (the DNS warm-up path) racing the manager lifecycle,
|
||||||
|
// which stays on the engine loop. Run with -race: it fails if ActivatePeer
|
||||||
|
// still requires engine.syncMsgMux for safety.
|
||||||
|
func TestConnMgr_ActivatePeerConcurrentWithLifecycle(t *testing.T) {
|
||||||
|
t.Setenv(lazyconn.EnvLazyConn, "on")
|
||||||
|
|
||||||
|
status := peer.NewRecorder("https://mgm")
|
||||||
|
store := peerstore.NewConnStore()
|
||||||
|
connMgr := NewConnMgr(&EngineConfig{}, status, store, mockLazyWGIface{})
|
||||||
|
|
||||||
|
conn := newTestPeerConn(t, "peerA")
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
connMgr.Start(ctx)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for range 4 {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
connMgr.ActivatePeer(ctx, conn)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Let the activators spin against the started manager, then tear it down
|
||||||
|
// underneath them and let them spin against the stopped manager.
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
connMgr.Close()
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
close(done)
|
||||||
|
wg.Wait()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInactivityThresholdEnv(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
val string
|
||||||
|
want *time.Duration
|
||||||
|
}{
|
||||||
|
{name: "unset", val: "", want: nil},
|
||||||
|
{name: "go duration minutes", val: "30m", want: durPtr(30 * time.Minute)},
|
||||||
|
{name: "go duration hours", val: "1h", want: durPtr(time.Hour)},
|
||||||
|
{name: "go duration seconds", val: "90s", want: durPtr(90 * time.Second)},
|
||||||
|
{name: "bare integer is minutes (backwards compat)", val: "5", want: durPtr(5 * time.Minute)},
|
||||||
|
{name: "zero duration", val: "0s", want: nil},
|
||||||
|
{name: "zero integer", val: "0", want: nil},
|
||||||
|
{name: "negative duration", val: "-5m", want: nil},
|
||||||
|
{name: "garbage", val: "abc", want: nil},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range tests {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
t.Setenv(lazyconn.EnvInactivityThreshold, tc.val)
|
||||||
|
got := inactivityThresholdEnv()
|
||||||
|
switch {
|
||||||
|
case tc.want == nil && got != nil:
|
||||||
|
t.Fatalf("want nil, got %v", *got)
|
||||||
|
case tc.want != nil && got == nil:
|
||||||
|
t.Fatalf("want %v, got nil", *tc.want)
|
||||||
|
case tc.want != nil && *got != *tc.want:
|
||||||
|
t.Fatalf("want %v, got %v", *tc.want, *got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func durPtr(d time.Duration) *time.Duration { return &d }
|
||||||
|
|||||||
@@ -34,6 +34,7 @@ import (
|
|||||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||||
"github.com/netbirdio/netbird/client/internal/statemanager"
|
"github.com/netbirdio/netbird/client/internal/statemanager"
|
||||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/tunnelnotifier"
|
||||||
"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"
|
||||||
@@ -112,11 +113,14 @@ func (c *ConnectClient) RunOnAndroid(
|
|||||||
stateFilePath string,
|
stateFilePath string,
|
||||||
cacheDir string,
|
cacheDir string,
|
||||||
) error {
|
) error {
|
||||||
|
notifier := tunnelnotifier.New(networkChangeListener, nil)
|
||||||
|
defer notifier.Close()
|
||||||
|
|
||||||
// in case of non Android os these variables will be nil
|
// in case of non Android os these variables will be nil
|
||||||
mobileDependency := MobileDependency{
|
mobileDependency := MobileDependency{
|
||||||
TunAdapter: tunAdapter,
|
TunAdapter: tunAdapter,
|
||||||
IFaceDiscover: iFaceDiscover,
|
IFaceDiscover: iFaceDiscover,
|
||||||
NetworkChangeListener: networkChangeListener,
|
NetworkChangeListener: notifier,
|
||||||
HostDNSAddresses: dnsAddresses,
|
HostDNSAddresses: dnsAddresses,
|
||||||
DnsReadyListener: dnsReadyListener,
|
DnsReadyListener: dnsReadyListener,
|
||||||
StateFilePath: stateFilePath,
|
StateFilePath: stateFilePath,
|
||||||
@@ -136,10 +140,13 @@ func (c *ConnectClient) RunOniOS(
|
|||||||
// Set GC percent to 5% to reduce memory usage as iOS only allows 50MB of memory for the extension.
|
// Set GC percent to 5% to reduce memory usage as iOS only allows 50MB of memory for the extension.
|
||||||
debug.SetGCPercent(5)
|
debug.SetGCPercent(5)
|
||||||
|
|
||||||
|
notifier := tunnelnotifier.New(networkChangeListener, dnsManager)
|
||||||
|
defer notifier.Close()
|
||||||
|
|
||||||
mobileDependency := MobileDependency{
|
mobileDependency := MobileDependency{
|
||||||
FileDescriptor: fileDescriptor,
|
FileDescriptor: fileDescriptor,
|
||||||
NetworkChangeListener: networkChangeListener,
|
NetworkChangeListener: notifier,
|
||||||
DnsManager: dnsManager,
|
DnsManager: notifier,
|
||||||
StateFilePath: stateFilePath,
|
StateFilePath: stateFilePath,
|
||||||
TempDir: cacheDir,
|
TempDir: cacheDir,
|
||||||
}
|
}
|
||||||
@@ -257,7 +264,10 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
|
|||||||
log.Errorf("failed to clean up temporary installer file: %v", err)
|
log.Errorf("failed to clean up temporary installer file: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
defer c.statusRecorder.ClientStop()
|
defer func() {
|
||||||
|
c.statusRecorder.SetSessionExpiresAt(time.Time{})
|
||||||
|
c.statusRecorder.ClientStop()
|
||||||
|
}()
|
||||||
operation := func() error {
|
operation := func() error {
|
||||||
// if context cancelled we not start new backoff cycle
|
// if context cancelled we not start new backoff cycle
|
||||||
if c.ctx.Err() != nil {
|
if c.ctx.Err() != nil {
|
||||||
@@ -277,6 +287,15 @@ 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)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
// On daemon shutdown / Down() the parent context is cancelled
|
||||||
|
// and the dial fails with "context canceled". Wrapping that
|
||||||
|
// into state would leave the snapshot stuck at Connecting+err
|
||||||
|
// until the backoff loop wakes up — instead let the operation
|
||||||
|
// return cleanly so the deferred state.Set(StatusIdle) takes
|
||||||
|
// effect on the next iteration.
|
||||||
|
if c.ctx.Err() != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
return wrapErr(gstatus.Errorf(codes.FailedPrecondition, "failed connecting to Management Service : %s", err))
|
return wrapErr(gstatus.Errorf(codes.FailedPrecondition, "failed connecting to Management Service : %s", err))
|
||||||
}
|
}
|
||||||
mgmNotifier := statusRecorderToMgmConnStateNotifier(c.statusRecorder)
|
mgmNotifier := statusRecorderToMgmConnStateNotifier(c.statusRecorder)
|
||||||
@@ -415,6 +434,10 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
|
|||||||
return wrapErr(err)
|
return wrapErr(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Seed the session-expiry deadline from the LoginResponse. Subsequent
|
||||||
|
// changes flow in through SyncResponse and are applied in handleSync.
|
||||||
|
engine.ApplySessionDeadline(loginResp.GetSessionExpiresAt())
|
||||||
|
|
||||||
log.Infof("Netbird engine started, the IP is: %s", peerConfig.GetAddress())
|
log.Infof("Netbird engine started, the IP is: %s", peerConfig.GetAddress())
|
||||||
state.Set(StatusConnected)
|
state.Set(StatusConnected)
|
||||||
|
|
||||||
@@ -451,6 +474,10 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
|
|||||||
}
|
}
|
||||||
|
|
||||||
c.statusRecorder.ClientStart()
|
c.statusRecorder.ClientStart()
|
||||||
|
// Wrap the backoff with c.ctx so Down()/actCancel propagates into the
|
||||||
|
// inter-attempt sleep — otherwise a 15s MaxInterval can keep the retry
|
||||||
|
// loop alive long after the caller asked to give up, leaving the
|
||||||
|
// 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 {
|
||||||
log.Debugf("exiting client retry loop due to unrecoverable error: %s", err)
|
log.Debugf("exiting client retry loop due to unrecoverable error: %s", err)
|
||||||
@@ -601,6 +628,7 @@ func createEngineConfig(key wgtypes.Key, config *profilemanager.Config, peerConf
|
|||||||
BlockLANAccess: config.BlockLANAccess,
|
BlockLANAccess: config.BlockLANAccess,
|
||||||
BlockInbound: config.BlockInbound,
|
BlockInbound: config.BlockInbound,
|
||||||
DisableIPv6: config.DisableIPv6,
|
DisableIPv6: config.DisableIPv6,
|
||||||
|
SyncMessageVersion: config.SyncMessageVersion,
|
||||||
|
|
||||||
LazyConnection: lazyconn.ParseState(config.LazyConnection),
|
LazyConnection: lazyconn.ParseState(config.LazyConnection),
|
||||||
|
|
||||||
@@ -676,6 +704,7 @@ func loginToManagement(ctx context.Context, client mgm.Client, pubSSHKey []byte,
|
|||||||
config.BlockLANAccess,
|
config.BlockLANAccess,
|
||||||
config.BlockInbound,
|
config.BlockInbound,
|
||||||
config.DisableIPv6,
|
config.DisableIPv6,
|
||||||
|
config.SyncMessageVersion,
|
||||||
config.EnableSSHRoot,
|
config.EnableSSHRoot,
|
||||||
config.EnableSSHSFTP,
|
config.EnableSSHSFTP,
|
||||||
config.EnableSSHLocalPortForwarding,
|
config.EnableSSHLocalPortForwarding,
|
||||||
|
|||||||
15
client/internal/daemonaddr/owner.go
Normal file
15
client/internal/daemonaddr/owner.go
Normal file
@@ -0,0 +1,15 @@
|
|||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
// DaemonRunsAsSelf reports whether the daemon listening at addr runs as this very
|
||||||
|
// user. That is what makes an unprivileged daemon authorize this process for the
|
||||||
|
// changes it otherwise restricts to root or an administrator, so a client can tell
|
||||||
|
// up front whether those controls are usable instead of letting a save fail.
|
||||||
|
//
|
||||||
|
// It is answered from the ownership of the socket or pipe the daemon created, so it
|
||||||
|
// costs no round trip and needs no cooperation from the daemon. Ownership that
|
||||||
|
// cannot be read is reported as false, including for a TCP address, so a caller
|
||||||
|
// reading this as "the daemon would allow it" fails closed. The daemon remains the
|
||||||
|
// only thing that authorizes anything: this only decides what a client offers.
|
||||||
|
func DaemonRunsAsSelf(addr string) bool {
|
||||||
|
return daemonRunsAsSelf(addr)
|
||||||
|
}
|
||||||
40
client/internal/daemonaddr/owner_unix.go
Normal file
40
client/internal/daemonaddr/owner_unix.go
Normal file
@@ -0,0 +1,40 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
)
|
||||||
|
|
||||||
|
// daemonRunsAsSelf compares the owner of the daemon's Unix socket with this
|
||||||
|
// process's uid. Root is not treated specially here: a root caller is privileged
|
||||||
|
// on its own merits, and a root-owned socket says nothing about the caller.
|
||||||
|
func daemonRunsAsSelf(addr string) bool {
|
||||||
|
path, ok := strings.CutPrefix(addr, "unix://")
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
info, err := os.Stat(path)
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("stat daemon socket %s: %v", path, err)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only a socket says anything about a daemon. A directory or a leftover
|
||||||
|
// regular file at that path is not one, and reading it as "the daemon runs as
|
||||||
|
// us" would offer controls the daemon then refuses.
|
||||||
|
if info.Mode()&os.ModeSocket == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
stat, ok := info.Sys().(*syscall.Stat_t)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return stat.Uid == uint32(os.Getuid())
|
||||||
|
}
|
||||||
62
client/internal/daemonaddr/owner_unix_test.go
Normal file
62
client/internal/daemonaddr/owner_unix_test.go
Normal file
@@ -0,0 +1,62 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// A socket this user created means the daemon runs as this user, which is the
|
||||||
|
// rootless case where the daemon delegates its authority to its own identity.
|
||||||
|
func TestDaemonRunsAsSelf_OwnSocket(t *testing.T) {
|
||||||
|
path := filepath.Join(t.TempDir(), "netbird.sock")
|
||||||
|
ln, err := net.Listen("unix", path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := ln.Close(); err != nil {
|
||||||
|
t.Logf("close listener: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
if !DaemonRunsAsSelf("unix://" + path) {
|
||||||
|
t.Error("a socket owned by this user must count as the daemon running as us")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Everything that is not a readable socket of ours has to answer false, because
|
||||||
|
// the caller reads a true as "the daemon would authorize me".
|
||||||
|
func TestDaemonRunsAsSelf_FailsClosed(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
// A socket owned by another user, which is what a root-run daemon looks like
|
||||||
|
// to an unprivileged client. Only assertable when we are not root ourselves.
|
||||||
|
rootOwned := "unix:///var/run/netbird.sock"
|
||||||
|
if _, err := os.Stat("/var/run/netbird.sock"); err == nil && os.Getuid() != 0 {
|
||||||
|
if DaemonRunsAsSelf(rootOwned) {
|
||||||
|
t.Error("a socket owned by another user must not count as ours")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for name, addr := range map[string]string{
|
||||||
|
"missing socket": "unix://" + filepath.Join(dir, "absent.sock"),
|
||||||
|
"tcp address": "tcp://127.0.0.1:41731",
|
||||||
|
"named pipe": "npipe://netbird",
|
||||||
|
"empty": "",
|
||||||
|
"no scheme": filepath.Join(dir, "absent.sock"),
|
||||||
|
"directory": "unix://" + dir,
|
||||||
|
"unknown scheme": "http://localhost:8080",
|
||||||
|
"scheme only": "unix://",
|
||||||
|
"relative socket": "unix://netbird.sock",
|
||||||
|
} {
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
if DaemonRunsAsSelf(addr) {
|
||||||
|
t.Errorf("%q must not count as a daemon running as us", addr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
42
client/internal/daemonaddr/owner_windows.go
Normal file
42
client/internal/daemonaddr/owner_windows.go
Normal file
@@ -0,0 +1,42 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||||
|
)
|
||||||
|
|
||||||
|
// daemonRunsAsSelf reads the owner of the daemon's pipe. A daemon running as the
|
||||||
|
// service account owns its pipe as LocalSystem, and an elevated one as
|
||||||
|
// BUILTIN\Administrators, so only a daemon the user started themselves matches.
|
||||||
|
func daemonRunsAsSelf(addr string) bool {
|
||||||
|
name, ok := strings.CutPrefix(addr, pipeScheme)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, path := range PipePaths(name) {
|
||||||
|
// Bounded: this runs on the UI's path for deciding which controls to
|
||||||
|
// offer, so a pipe that does not answer promptly must not stall it. A
|
||||||
|
// timeout leaves the caller unprivileged, which only disables controls.
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), probeTimeout)
|
||||||
|
conn, err := dialPipe(ctx, path)
|
||||||
|
cancel()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
owned := ipcauth.PipeOwnedBySelf(conn)
|
||||||
|
if cerr := conn.Close(); cerr != nil {
|
||||||
|
log.Debugf("close daemon pipe %s after ownership check: %v", path, cerr)
|
||||||
|
}
|
||||||
|
return owned
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
103
client/internal/daemonaddr/pipe.go
Normal file
103
client/internal/daemonaddr/pipe.go
Normal file
@@ -0,0 +1,103 @@
|
|||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"google.golang.org/grpc"
|
||||||
|
"google.golang.org/grpc/credentials/insecure"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// WindowsPipeAddr is the default daemon address on Windows. A named pipe
|
||||||
|
// carries the connecting process's token, which loopback TCP does not, so
|
||||||
|
// it is the only Windows transport on which the daemon can tell who is
|
||||||
|
// calling it.
|
||||||
|
WindowsPipeAddr = "npipe://netbird"
|
||||||
|
|
||||||
|
// legacyWindowsAddr is the loopback-TCP address the Windows daemon used
|
||||||
|
// before named-pipe support.
|
||||||
|
legacyWindowsAddr = "tcp://127.0.0.1:41731"
|
||||||
|
|
||||||
|
pipeScheme = "npipe://"
|
||||||
|
|
||||||
|
// protectedPrefix is the NPFS namespace in which only LocalSystem and
|
||||||
|
// members of BUILTIN\Administrators may create a pipe. A daemon running as
|
||||||
|
// the service account creates its pipe there so that an unprivileged process
|
||||||
|
// cannot pre-create the name, which would keep the daemon from starting and
|
||||||
|
// leave callers talking to the squatter. Opening such a pipe needs no
|
||||||
|
// privilege, so unprivileged clients still reach the daemon.
|
||||||
|
protectedPrefix = `ProtectedPrefix\Administrators\`
|
||||||
|
)
|
||||||
|
|
||||||
|
// DialTarget returns the gRPC dial target and transport options for a daemon
|
||||||
|
// address. The npipe scheme needs a context dialer because gRPC has no
|
||||||
|
// named-pipe resolver; unix and tcp are handled by gRPC itself.
|
||||||
|
func DialTarget(addr string) (string, []grpc.DialOption) {
|
||||||
|
opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())}
|
||||||
|
|
||||||
|
if name, ok := strings.CutPrefix(addr, pipeScheme); ok {
|
||||||
|
paths := PipePaths(name)
|
||||||
|
opts = append(opts, grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) {
|
||||||
|
return dialPipePaths(ctx, paths)
|
||||||
|
}))
|
||||||
|
return "passthrough:///netbird-daemon-pipe", opts
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.TrimPrefix(addr, "tcp://"), opts
|
||||||
|
}
|
||||||
|
|
||||||
|
// PipePath maps an npipe address name ("netbird", from "npipe://netbird") to a
|
||||||
|
// Windows named-pipe path (\\.\pipe\netbird). A fully qualified path is left as
|
||||||
|
// is.
|
||||||
|
func PipePath(name string) string {
|
||||||
|
if strings.HasPrefix(name, `\\`) {
|
||||||
|
return name
|
||||||
|
}
|
||||||
|
return `\\.\pipe\` + name
|
||||||
|
}
|
||||||
|
|
||||||
|
// PipePaths returns the paths a daemon control pipe may live at for an npipe
|
||||||
|
// address name, in the order both sides must try them: the protected name first,
|
||||||
|
// then the plain one.
|
||||||
|
//
|
||||||
|
// The daemon serves the first it can create, which is the protected name when it
|
||||||
|
// runs as the service account and the plain one when it runs as an ordinary user,
|
||||||
|
// as it does in netstack mode. Clients therefore have to try both, and because a
|
||||||
|
// client cannot tell from the name alone who created the pipe, the plain name is
|
||||||
|
// only usable once the server's identity has been checked: see
|
||||||
|
// verifyPipeServer.
|
||||||
|
//
|
||||||
|
// A fully qualified path is what the operator asked for and is used as is.
|
||||||
|
func PipePaths(name string) []string {
|
||||||
|
if strings.HasPrefix(name, `\\`) {
|
||||||
|
return []string{name}
|
||||||
|
}
|
||||||
|
return []string{PipePath(protectedPrefix + name), PipePath(name)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsProtectedPipePath reports whether a pipe path is in the namespace only an
|
||||||
|
// administrator or LocalSystem can create in, which is what lets a client trust
|
||||||
|
// such a pipe from its name alone.
|
||||||
|
func IsProtectedPipePath(path string) bool {
|
||||||
|
return strings.HasPrefix(path, `\\.\pipe\`+protectedPrefix)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MigrateLegacy upgrades the pre-named-pipe Windows daemon address to the named
|
||||||
|
// pipe, reporting whether it rewrote the address. Existing installs persist the
|
||||||
|
// daemon address, so without this an upgraded daemon would keep listening on
|
||||||
|
// loopback TCP, where callers carry no identity and privileged operations would
|
||||||
|
// have to be refused for everyone. Only the exact legacy default is rewritten:
|
||||||
|
// a deliberately chosen custom address is left alone.
|
||||||
|
func MigrateLegacy(addr string) (string, bool) {
|
||||||
|
return migrateLegacyForOS(runtime.GOOS, addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
func migrateLegacyForOS(goos, addr string) (string, bool) {
|
||||||
|
if goos == "windows" && addr == legacyWindowsAddr {
|
||||||
|
return WindowsPipeAddr, true
|
||||||
|
}
|
||||||
|
return addr, false
|
||||||
|
}
|
||||||
15
client/internal/daemonaddr/pipe_other.go
Normal file
15
client/internal/daemonaddr/pipe_other.go
Normal file
@@ -0,0 +1,15 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
)
|
||||||
|
|
||||||
|
// dialPipePaths is Windows-only: no other platform serves the daemon on a named
|
||||||
|
// pipe.
|
||||||
|
func dialPipePaths(context.Context, []string) (net.Conn, error) {
|
||||||
|
return nil, fmt.Errorf("named pipes are only supported on Windows")
|
||||||
|
}
|
||||||
30
client/internal/daemonaddr/pipe_test.go
Normal file
30
client/internal/daemonaddr/pipe_test.go
Normal file
@@ -0,0 +1,30 @@
|
|||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The protected name must be tried before the plain one on both sides: it is the
|
||||||
|
// one an unprivileged process cannot create, so preferring it is what keeps a
|
||||||
|
// squatter from owning the name the service daemon would otherwise use.
|
||||||
|
func TestPipePaths_PrefersTheProtectedName(t *testing.T) {
|
||||||
|
got := PipePaths("netbird")
|
||||||
|
want := []string{
|
||||||
|
`\\.\pipe\ProtectedPrefix\Administrators\netbird`,
|
||||||
|
`\\.\pipe\netbird`,
|
||||||
|
}
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("PipePaths = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// An operator who passes a full path chose exactly one pipe, so neither side may
|
||||||
|
// look anywhere else.
|
||||||
|
func TestPipePaths_QualifiedPathIsUsedAsIs(t *testing.T) {
|
||||||
|
path := `\\.\pipe\custom-netbird`
|
||||||
|
got := PipePaths(path)
|
||||||
|
if !slices.Equal(got, []string{path}) {
|
||||||
|
t.Errorf("PipePaths = %q, want just %q", got, path)
|
||||||
|
}
|
||||||
|
}
|
||||||
59
client/internal/daemonaddr/pipe_windows.go
Normal file
59
client/internal/daemonaddr/pipe_windows.go
Normal file
@@ -0,0 +1,59 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"github.com/Microsoft/go-winio"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||||
|
)
|
||||||
|
|
||||||
|
// dialPipePaths connects to the first path that answers with a pipe server this
|
||||||
|
// client may trust, and returns the last error when none does.
|
||||||
|
func dialPipePaths(ctx context.Context, paths []string) (net.Conn, error) {
|
||||||
|
var lastErr error
|
||||||
|
for _, path := range paths {
|
||||||
|
conn, err := dialPipe(ctx, path)
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("dial daemon pipe %s: %v", path, err)
|
||||||
|
lastErr = err
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// A pipe in the protected namespace could only have been created by an
|
||||||
|
// administrator or LocalSystem, so its name is the guarantee. Any other
|
||||||
|
// name has to be checked, because any local user can create one.
|
||||||
|
if !IsProtectedPipePath(path) {
|
||||||
|
if err := ipcauth.PipeServerTrusted(conn); err != nil {
|
||||||
|
if closeErr := conn.Close(); closeErr != nil {
|
||||||
|
log.Debugf("close untrusted pipe %s: %v", path, closeErr)
|
||||||
|
}
|
||||||
|
lastErr = fmt.Errorf("%s: %w", path, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return conn, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if lastErr == nil {
|
||||||
|
lastErr = errors.New("no daemon pipe to connect to")
|
||||||
|
}
|
||||||
|
return nil, lastErr
|
||||||
|
}
|
||||||
|
|
||||||
|
// dialPipe connects to the daemon control pipe at SECURITY_IDENTIFICATION.
|
||||||
|
// winio's plain DialPipe connects at SECURITY_ANONYMOUS, under which the daemon
|
||||||
|
// cannot read the caller's token at all. Identification lets the daemon read the
|
||||||
|
// caller's SID and groups without granting it the ability to act as the caller.
|
||||||
|
func dialPipe(ctx context.Context, path string) (net.Conn, error) {
|
||||||
|
access := uint32(windows.GENERIC_READ | windows.GENERIC_WRITE)
|
||||||
|
return winio.DialPipeAccessImpLevel(ctx, path, access, winio.PipeImpLevelIdentification)
|
||||||
|
}
|
||||||
9
client/internal/daemonaddr/resolve_pipe_other.go
Normal file
9
client/internal/daemonaddr/resolve_pipe_other.go
Normal file
@@ -0,0 +1,9 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
// ResolveDaemonAddr is a no-op off Windows, where there is no named-pipe
|
||||||
|
// default to fall back from.
|
||||||
|
func ResolveDaemonAddr(addr string) string {
|
||||||
|
return addr
|
||||||
|
}
|
||||||
82
client/internal/daemonaddr/resolve_pipe_windows.go
Normal file
82
client/internal/daemonaddr/resolve_pipe_windows.go
Normal file
@@ -0,0 +1,82 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package daemonaddr
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/Microsoft/go-winio"
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
)
|
||||||
|
|
||||||
|
// probeTimeout bounds each transport probe. Both are local, so a daemon that is
|
||||||
|
// listening answers immediately and one that is not fails immediately.
|
||||||
|
const probeTimeout = 300 * time.Millisecond
|
||||||
|
|
||||||
|
// ResolveDaemonAddr keeps a client on the named pipe and never silently moves it
|
||||||
|
// off. When the pipe does not answer it checks the legacy loopback TCP address, so
|
||||||
|
// a client meeting a daemon that has not restarted since the upgrade can say what
|
||||||
|
// is wrong, but it does not connect there.
|
||||||
|
//
|
||||||
|
// Using that address automatically would be a downgrade the user never asked for:
|
||||||
|
// any local process can bind 127.0.0.1 while the daemon is not listening, and the
|
||||||
|
// transport carries no caller identity, so a client that accepted whatever answered
|
||||||
|
// would hand a setup key, a pre-shared key or an SSO prompt to a local impostor. An
|
||||||
|
// operator who needs the legacy address during the upgrade window can still pass
|
||||||
|
// --daemon-addr explicitly, which is a deliberate choice and still refuses the
|
||||||
|
// privileged operations.
|
||||||
|
//
|
||||||
|
// Only the pipe address is resolved. A custom address is left alone, though passing
|
||||||
|
// --daemon-addr npipe://netbird explicitly is indistinguishable from the default
|
||||||
|
// here, so it is treated the same way.
|
||||||
|
func ResolveDaemonAddr(addr string) string {
|
||||||
|
if addr != WindowsPipeAddr {
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, path := range PipePaths("netbird") {
|
||||||
|
if pipeAvailable(path) {
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if tcpAvailable(legacyWindowsAddr) {
|
||||||
|
log.Warnf("the daemon is not serving %s, but something is listening on the legacy %s. "+
|
||||||
|
"Restart the NetBird service so it serves the pipe. That address is not used automatically: "+
|
||||||
|
"any local user can bind it and it carries no caller identity, so pass --daemon-addr %s "+
|
||||||
|
"explicitly if you accept that",
|
||||||
|
WindowsPipeAddr, legacyWindowsAddr, legacyWindowsAddr)
|
||||||
|
}
|
||||||
|
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func pipeAvailable(path string) bool {
|
||||||
|
timeout := probeTimeout
|
||||||
|
conn, err := winio.DialPipe(path, &timeout)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if err := conn.Close(); err != nil {
|
||||||
|
log.Debugf("close daemon pipe probe: %v", err)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func tcpAvailable(addr string) bool {
|
||||||
|
host := addr
|
||||||
|
if _, after, ok := strings.Cut(addr, "://"); ok {
|
||||||
|
host = after
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := net.DialTimeout("tcp", host, probeTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if err := conn.Close(); err != nil {
|
||||||
|
log.Debugf("close daemon TCP probe: %v", err)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
@@ -229,9 +229,16 @@ scutil_dns.txt (macOS only):
|
|||||||
|
|
||||||
const (
|
const (
|
||||||
clientLogFile = "client.log"
|
clientLogFile = "client.log"
|
||||||
|
uiLogFile = "gui-client.log"
|
||||||
errorLogFile = "netbird.err"
|
errorLogFile = "netbird.err"
|
||||||
stdoutLogFile = "netbird.out"
|
stdoutLogFile = "netbird.out"
|
||||||
|
|
||||||
|
// Rotated-log glob prefixes (base log name without extension) passed to
|
||||||
|
// addRotatedLogFiles. The daemon's own log and the GUI log live in the same
|
||||||
|
// dir, so the prefixes must be disjoint to keep their rotated siblings apart.
|
||||||
|
clientLogPrefix = "client"
|
||||||
|
uiLogPrefix = "gui-client"
|
||||||
|
|
||||||
darwinErrorLogPath = "/var/log/netbird.out.log"
|
darwinErrorLogPath = "/var/log/netbird.out.log"
|
||||||
darwinStdoutLogPath = "/var/log/netbird.err.log"
|
darwinStdoutLogPath = "/var/log/netbird.err.log"
|
||||||
)
|
)
|
||||||
@@ -249,6 +256,7 @@ type BundleGenerator struct {
|
|||||||
statusRecorder *peer.Status
|
statusRecorder *peer.Status
|
||||||
syncResponse *mgmProto.SyncResponse
|
syncResponse *mgmProto.SyncResponse
|
||||||
logPath string
|
logPath string
|
||||||
|
uiLogPath string
|
||||||
tempDir string
|
tempDir string
|
||||||
statePath string
|
statePath string
|
||||||
cpuProfile []byte
|
cpuProfile []byte
|
||||||
@@ -276,6 +284,7 @@ type GeneratorDependencies struct {
|
|||||||
StatusRecorder *peer.Status
|
StatusRecorder *peer.Status
|
||||||
SyncResponse *mgmProto.SyncResponse
|
SyncResponse *mgmProto.SyncResponse
|
||||||
LogPath string
|
LogPath string
|
||||||
|
UILogPath string // Absolute path to the desktop UI's gui-client.log, reported via RegisterUILog. Empty if no UI registered one.
|
||||||
TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used.
|
TempDir string // Directory for temporary bundle zip files. If empty, os.TempDir() is used.
|
||||||
StatePath string // Path to the state file. If empty, the ServiceManager default path is used.
|
StatePath string // Path to the state file. If empty, the ServiceManager default path is used.
|
||||||
CPUProfile []byte
|
CPUProfile []byte
|
||||||
@@ -300,6 +309,7 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
|
|||||||
statusRecorder: deps.StatusRecorder,
|
statusRecorder: deps.StatusRecorder,
|
||||||
syncResponse: deps.SyncResponse,
|
syncResponse: deps.SyncResponse,
|
||||||
logPath: deps.LogPath,
|
logPath: deps.LogPath,
|
||||||
|
uiLogPath: deps.UILogPath,
|
||||||
tempDir: deps.TempDir,
|
tempDir: deps.TempDir,
|
||||||
statePath: deps.StatePath,
|
statePath: deps.StatePath,
|
||||||
cpuProfile: deps.CPUProfile,
|
cpuProfile: deps.CPUProfile,
|
||||||
@@ -411,6 +421,10 @@ func (g *BundleGenerator) createArchive() error {
|
|||||||
log.Errorf("failed to add logs to debug bundle: %v", err)
|
log.Errorf("failed to add logs to debug bundle: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if err := g.addUILog(); err != nil {
|
||||||
|
log.Errorf("failed to add UI log to debug bundle: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
if err := g.addUpdateLogs(); err != nil {
|
if err := g.addUpdateLogs(); err != nil {
|
||||||
log.Errorf("failed to add updater logs: %v", err)
|
log.Errorf("failed to add updater logs: %v", err)
|
||||||
}
|
}
|
||||||
@@ -466,7 +480,6 @@ func (g *BundleGenerator) addStatus() error {
|
|||||||
|
|
||||||
fullStatus := g.statusRecorder.GetFullStatus()
|
fullStatus := g.statusRecorder.GetFullStatus()
|
||||||
protoFullStatus := nbstatus.ToProtoFullStatus(fullStatus)
|
protoFullStatus := nbstatus.ToProtoFullStatus(fullStatus)
|
||||||
protoFullStatus.Events = g.statusRecorder.GetEventHistory()
|
|
||||||
overview := nbstatus.ConvertToStatusOutputOverview(protoFullStatus, nbstatus.ConvertOptions{
|
overview := nbstatus.ConvertToStatusOutputOverview(protoFullStatus, nbstatus.ConvertOptions{
|
||||||
Anonymize: g.anonymize,
|
Anonymize: g.anonymize,
|
||||||
ProfileName: profName,
|
ProfileName: profName,
|
||||||
@@ -663,6 +676,7 @@ func (g *BundleGenerator) addCommonConfigFields(configContent *strings.Builder)
|
|||||||
configContent.WriteString(fmt.Sprintf("BlockLANAccess: %v\n", g.internalConfig.BlockLANAccess))
|
configContent.WriteString(fmt.Sprintf("BlockLANAccess: %v\n", g.internalConfig.BlockLANAccess))
|
||||||
configContent.WriteString(fmt.Sprintf("BlockInbound: %v\n", g.internalConfig.BlockInbound))
|
configContent.WriteString(fmt.Sprintf("BlockInbound: %v\n", g.internalConfig.BlockInbound))
|
||||||
configContent.WriteString(fmt.Sprintf("DisableIPv6: %v\n", g.internalConfig.DisableIPv6))
|
configContent.WriteString(fmt.Sprintf("DisableIPv6: %v\n", g.internalConfig.DisableIPv6))
|
||||||
|
configContent.WriteString(fmt.Sprintf("SyncMessageVersion: %v\n", g.internalConfig.SyncMessageVersion))
|
||||||
|
|
||||||
if g.internalConfig.DisableNotifications != nil {
|
if g.internalConfig.DisableNotifications != nil {
|
||||||
configContent.WriteString(fmt.Sprintf("DisableNotifications: %v\n", *g.internalConfig.DisableNotifications))
|
configContent.WriteString(fmt.Sprintf("DisableNotifications: %v\n", *g.internalConfig.DisableNotifications))
|
||||||
@@ -986,7 +1000,7 @@ func (g *BundleGenerator) addLogfile() error {
|
|||||||
return fmt.Errorf("add client log file to zip: %w", err)
|
return fmt.Errorf("add client log file to zip: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
g.addRotatedLogFiles(logDir)
|
g.addRotatedLogFiles(logDir, clientLogPrefix)
|
||||||
|
|
||||||
stdErrLogPath := filepath.Join(logDir, errorLogFile)
|
stdErrLogPath := filepath.Join(logDir, errorLogFile)
|
||||||
stdoutLogPath := filepath.Join(logDir, stdoutLogFile)
|
stdoutLogPath := filepath.Join(logDir, stdoutLogFile)
|
||||||
@@ -1006,6 +1020,25 @@ func (g *BundleGenerator) addLogfile() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// addUILog adds the desktop UI's gui-client.log (and its rotated siblings) to
|
||||||
|
// the bundle. The path is reported by the UI via RegisterUILog; empty when no
|
||||||
|
// UI registered one (e.g. headless / server). Missing file is non-fatal — the
|
||||||
|
// UI only writes it while the daemon is in debug, so it's often absent.
|
||||||
|
func (g *BundleGenerator) addUILog() error {
|
||||||
|
if g.uiLogPath == "" {
|
||||||
|
log.Debugf("no UI log path registered, skipping in debug bundle")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := g.addSingleLogfile(g.uiLogPath, uiLogFile); err != nil {
|
||||||
|
return fmt.Errorf("add UI log file to zip: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
g.addRotatedLogFiles(filepath.Dir(g.uiLogPath), uiLogPrefix)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// addSingleLogfile adds a single log file to the archive
|
// addSingleLogfile adds a single log file to the archive
|
||||||
func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error {
|
func (g *BundleGenerator) addSingleLogfile(logPath, targetName string) error {
|
||||||
logFile, err := os.Open(logPath)
|
logFile, err := os.Open(logPath)
|
||||||
@@ -1078,14 +1111,16 @@ func (g *BundleGenerator) addSingleLogFileGz(logPath, targetName string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// addRotatedLogFiles adds rotated log files to the bundle based on logFileCount
|
// addRotatedLogFiles adds rotated log files to the bundle based on logFileCount.
|
||||||
func (g *BundleGenerator) addRotatedLogFiles(logDir string) {
|
// prefix is the base log name without extension (e.g. "client", "gui-client");
|
||||||
|
// the glob matches both files rotated by us and by logrotate on linux.
|
||||||
|
func (g *BundleGenerator) addRotatedLogFiles(logDir, prefix string) {
|
||||||
if g.logFileCount == 0 {
|
if g.logFileCount == 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// This regex will match both logs rotated by us and logrotate on linux
|
// This pattern matches both logs rotated by us and logrotate on linux
|
||||||
pattern := filepath.Join(logDir, "client*.log.*")
|
pattern := filepath.Join(logDir, prefix+"*.log.*")
|
||||||
files, err := filepath.Glob(pattern)
|
files, err := filepath.Glob(pattern)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Warnf("failed to glob rotated logs: %v", err)
|
log.Warnf("failed to glob rotated logs: %v", err)
|
||||||
|
|||||||
@@ -40,6 +40,25 @@ func TestAddRotatedLogFiles_PicksUpAllVariants(t *testing.T) {
|
|||||||
require.NotContains(t, names, "other.log", "unrelated files should not be in bundle")
|
require.NotContains(t, names, "other.log", "unrelated files should not be in bundle")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestAddRotatedLogFiles_GUIPrefix asserts the prefix parameter scopes the glob
|
||||||
|
// to the GUI log: gui-client.log.* rotated siblings are picked up and the
|
||||||
|
// daemon's own client.log.* are not (and vice versa, covered above). This is
|
||||||
|
// the load-bearing check for the gui-client.log bundle collection — the old
|
||||||
|
// "client*.log.*" glob would have missed gui-client rotations.
|
||||||
|
func TestAddRotatedLogFiles_GUIPrefix(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
|
||||||
|
writeFile(t, filepath.Join(dir, "gui-client.log.1"), "gui rotated\n")
|
||||||
|
writeGzFile(t, filepath.Join(dir, "gui-client.log.2.gz"), "gui rotated gz\n")
|
||||||
|
writeFile(t, filepath.Join(dir, "client.log.1"), "daemon rotated\n")
|
||||||
|
|
||||||
|
names := runAddRotatedLogFilesPrefix(t, dir, "gui-client", 10)
|
||||||
|
|
||||||
|
require.Contains(t, names, "gui-client.log.1", "gui-client rotated file should be in bundle")
|
||||||
|
require.Contains(t, names, "gui-client.log.2.gz", "gui-client gz rotated file should be in bundle")
|
||||||
|
require.NotContains(t, names, "client.log.1", "daemon rotated file must not match the gui-client prefix")
|
||||||
|
}
|
||||||
|
|
||||||
// TestAddRotatedLogFiles_RespectsLogFileCount asserts that only the newest
|
// TestAddRotatedLogFiles_RespectsLogFileCount asserts that only the newest
|
||||||
// logFileCount rotated files are bundled, ordered by mtime.
|
// logFileCount rotated files are bundled, ordered by mtime.
|
||||||
func TestAddRotatedLogFiles_RespectsLogFileCount(t *testing.T) {
|
func TestAddRotatedLogFiles_RespectsLogFileCount(t *testing.T) {
|
||||||
@@ -67,6 +86,10 @@ func TestAddRotatedLogFiles_RespectsLogFileCount(t *testing.T) {
|
|||||||
// runAddRotatedLogFiles calls addRotatedLogFiles against a fresh in-memory
|
// runAddRotatedLogFiles calls addRotatedLogFiles against a fresh in-memory
|
||||||
// zip writer and returns the set of entry names that ended up in the archive.
|
// zip writer and returns the set of entry names that ended up in the archive.
|
||||||
func runAddRotatedLogFiles(t *testing.T, dir string, logFileCount uint32) map[string]struct{} {
|
func runAddRotatedLogFiles(t *testing.T, dir string, logFileCount uint32) map[string]struct{} {
|
||||||
|
return runAddRotatedLogFilesPrefix(t, dir, "client", logFileCount)
|
||||||
|
}
|
||||||
|
|
||||||
|
func runAddRotatedLogFilesPrefix(t *testing.T, dir, prefix string, logFileCount uint32) map[string]struct{} {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
@@ -74,7 +97,7 @@ func runAddRotatedLogFiles(t *testing.T, dir string, logFileCount uint32) map[st
|
|||||||
archive: zip.NewWriter(&buf),
|
archive: zip.NewWriter(&buf),
|
||||||
logFileCount: logFileCount,
|
logFileCount: logFileCount,
|
||||||
}
|
}
|
||||||
g.addRotatedLogFiles(dir)
|
g.addRotatedLogFiles(dir, prefix)
|
||||||
require.NoError(t, g.archive.Close())
|
require.NoError(t, g.archive.Close())
|
||||||
|
|
||||||
zr, err := zip.NewReader(bytes.NewReader(buf.Bytes()), int64(buf.Len()))
|
zr, err := zip.NewReader(bytes.NewReader(buf.Bytes()), int64(buf.Len()))
|
||||||
|
|||||||
@@ -887,6 +887,8 @@ func TestAddConfig_AllFieldsCovered(t *testing.T) {
|
|||||||
ClientCertKeyPath: "/tmp/key",
|
ClientCertKeyPath: "/tmp/key",
|
||||||
LazyConnection: "on",
|
LazyConnection: "on",
|
||||||
MTU: 1280,
|
MTU: 1280,
|
||||||
|
DisableIPv6: true,
|
||||||
|
SyncMessageVersion: func(v int) *int { return &v }(1),
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, anonymize := range []bool{false, true} {
|
for _, anonymize := range []bool{false, true} {
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -36,7 +37,43 @@ type resolver interface {
|
|||||||
// record is left alone (it points at something outside our mesh, e.g.
|
// record is left alone (it points at something outside our mesh, e.g.
|
||||||
// a non-peer upstream).
|
// a non-peer upstream).
|
||||||
type PeerConnectivity interface {
|
type PeerConnectivity interface {
|
||||||
IsConnectedByIP(ip string) (known, connected bool)
|
IsConnectedByIP(ip netip.Addr) (known, connected bool)
|
||||||
|
}
|
||||||
|
|
||||||
|
// PeerActivator wakes lazy-connection peers on demand. The local resolver calls
|
||||||
|
// it with the tunnel IPs an answer points at, so a peer that is idle (lazily
|
||||||
|
// disconnected) starts connecting at DNS-resolution time rather than racing the
|
||||||
|
// client's first request packet. nil disables warm-up.
|
||||||
|
type PeerActivator interface {
|
||||||
|
// ActivatePeersByIP triggers wake-up for the peer(s) owning addrs and blocks
|
||||||
|
// until one is connected or ctx (a short per-query budget) expires. It is a
|
||||||
|
// fast no-op for unknown or already-connected addresses.
|
||||||
|
ActivatePeersByIP(ctx context.Context, addrs []netip.Addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultLazyWarmupTimeout = 2 * time.Second
|
||||||
|
envLazyWarmupTimeout = "NB_DNS_LAZY_WARMUP_TIMEOUT"
|
||||||
|
)
|
||||||
|
|
||||||
|
// lazyWarmupTimeoutFromEnv returns the per-query budget for waking a
|
||||||
|
// lazy-connection peer a DNS answer points at. Tunable via
|
||||||
|
// NB_DNS_LAZY_WARMUP_TIMEOUT (a Go duration). Parsed once at construction time.
|
||||||
|
func lazyWarmupTimeoutFromEnv() time.Duration {
|
||||||
|
v := os.Getenv(envLazyWarmupTimeout)
|
||||||
|
if v == "" {
|
||||||
|
return defaultLazyWarmupTimeout
|
||||||
|
}
|
||||||
|
d, err := time.ParseDuration(v)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("invalid %s value %q, using default %s: %v", envLazyWarmupTimeout, v, defaultLazyWarmupTimeout, err)
|
||||||
|
return defaultLazyWarmupTimeout
|
||||||
|
}
|
||||||
|
if d <= 0 {
|
||||||
|
log.Warnf("non-positive %s value %q, using default %s", envLazyWarmupTimeout, v, defaultLazyWarmupTimeout)
|
||||||
|
return defaultLazyWarmupTimeout
|
||||||
|
}
|
||||||
|
return d
|
||||||
}
|
}
|
||||||
|
|
||||||
type Resolver struct {
|
type Resolver struct {
|
||||||
@@ -51,6 +88,12 @@ type Resolver struct {
|
|||||||
// filter and preserves the legacy "return whatever is registered"
|
// filter and preserves the legacy "return whatever is registered"
|
||||||
// behaviour for callers that never wire a status source.
|
// behaviour for callers that never wire a status source.
|
||||||
peerConn PeerConnectivity
|
peerConn PeerConnectivity
|
||||||
|
// peerActivator, when non-nil, is called at resolution time to warm the
|
||||||
|
// lazy connection to the peer(s) an answer points at. nil disables warm-up.
|
||||||
|
peerActivator PeerActivator
|
||||||
|
// warmupTimeout is the per-query budget for the lazy-connection warm-up
|
||||||
|
// wait, resolved from the environment once at construction time.
|
||||||
|
warmupTimeout time.Duration
|
||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
@@ -59,11 +102,12 @@ type Resolver struct {
|
|||||||
func NewResolver() *Resolver {
|
func NewResolver() *Resolver {
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
return &Resolver{
|
return &Resolver{
|
||||||
records: make(map[dns.Question][]dns.RR),
|
records: make(map[dns.Question][]dns.RR),
|
||||||
domains: make(map[domain.Domain]struct{}),
|
domains: make(map[domain.Domain]struct{}),
|
||||||
zones: make(map[domain.Domain]bool),
|
zones: make(map[domain.Domain]bool),
|
||||||
ctx: ctx,
|
warmupTimeout: lazyWarmupTimeoutFromEnv(),
|
||||||
cancel: cancel,
|
ctx: ctx,
|
||||||
|
cancel: cancel,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -76,6 +120,14 @@ func (d *Resolver) SetPeerConnectivity(p PeerConnectivity) {
|
|||||||
d.peerConn = p
|
d.peerConn = p
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetPeerActivator wires the DNS-time lazy-connection warm-up. Pass nil to
|
||||||
|
// disable. Safe to call multiple times; the latest value wins.
|
||||||
|
func (d *Resolver) SetPeerActivator(a PeerActivator) {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
d.peerActivator = a
|
||||||
|
}
|
||||||
|
|
||||||
func (d *Resolver) MatchSubdomains() bool {
|
func (d *Resolver) MatchSubdomains() bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
@@ -122,6 +174,9 @@ func (d *Resolver) ServeDNS(w dns.ResponseWriter, r *dns.Msg) {
|
|||||||
replyMessage.RecursionAvailable = true
|
replyMessage.RecursionAvailable = true
|
||||||
|
|
||||||
result := d.lookupRecords(logger, question)
|
result := d.lookupRecords(logger, question)
|
||||||
|
// Warm before filtering: activation flips a lazily-idle target to connected,
|
||||||
|
// which then lets it survive the disconnected-peer filter below.
|
||||||
|
d.warmLazyPeers(question, result.records)
|
||||||
result.records = d.filterDisconnectedPeerAnswers(logger, question, result.records)
|
result.records = d.filterDisconnectedPeerAnswers(logger, question, result.records)
|
||||||
replyMessage.Authoritative = !result.hasExternalData
|
replyMessage.Authoritative = !result.hasExternalData
|
||||||
replyMessage.Answer = result.records
|
replyMessage.Answer = result.records
|
||||||
@@ -495,8 +550,8 @@ func (d *Resolver) filterDisconnectedPeerAnswers(logger *log.Entry, question dns
|
|||||||
kept := make([]dns.RR, 0, len(records))
|
kept := make([]dns.RR, 0, len(records))
|
||||||
var dropped int
|
var dropped int
|
||||||
for _, rr := range records {
|
for _, rr := range records {
|
||||||
ip := extractRecordIP(rr)
|
ip, ok := extractRecordAddr(rr)
|
||||||
if ip == "" {
|
if !ok {
|
||||||
kept = append(kept, rr)
|
kept = append(kept, rr)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -518,22 +573,57 @@ func (d *Resolver) filterDisconnectedPeerAnswers(logger *log.Entry, question dns
|
|||||||
return kept
|
return kept
|
||||||
}
|
}
|
||||||
|
|
||||||
// extractRecordIP returns the dotted-decimal / colon-hex IP carried by
|
// warmLazyPeers triggers lazy-connection wake-up for the peers a resolved
|
||||||
// an A or AAAA record, or "" for any other record type.
|
// answer points at and waits briefly for one to connect, so the caller's first
|
||||||
func extractRecordIP(rr dns.RR) string {
|
// request doesn't race the connection establishment. Warm-up is scoped to
|
||||||
|
// match-only (non-authoritative) zones — the synthesized private-service zones
|
||||||
|
// and user-created zones whose records point at specific peers. The account's
|
||||||
|
// peer zone is authoritative, so plain peer-name lookups never trigger warm-up;
|
||||||
|
// otherwise resolving any peer's name would wake its idle connection, defeating
|
||||||
|
// laziness mesh-wide. No-op when no activator is wired (lazy connections
|
||||||
|
// disabled) or the answer carries no peer IPs.
|
||||||
|
func (d *Resolver) warmLazyPeers(question dns.Question, records []dns.RR) {
|
||||||
|
if len(records) < 2 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
d.mu.RLock()
|
||||||
|
activator := d.peerActivator
|
||||||
|
var nonAuth, found bool
|
||||||
|
if activator != nil {
|
||||||
|
nonAuth, found = d.findZone(question.Name)
|
||||||
|
}
|
||||||
|
d.mu.RUnlock()
|
||||||
|
if activator == nil || !found || !nonAuth {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var addrs []netip.Addr
|
||||||
|
for _, rr := range records {
|
||||||
|
if addr, ok := extractRecordAddr(rr); ok {
|
||||||
|
addrs = append(addrs, addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(addrs) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(d.ctx, d.warmupTimeout)
|
||||||
|
defer cancel()
|
||||||
|
activator.ActivatePeersByIP(ctx, addrs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractRecordAddr returns the IP address carried by an A or AAAA record.
|
||||||
|
// ok is false for any other record type or a record with no address.
|
||||||
|
func extractRecordAddr(rr dns.RR) (netip.Addr, bool) {
|
||||||
switch r := rr.(type) {
|
switch r := rr.(type) {
|
||||||
case *dns.A:
|
case *dns.A:
|
||||||
if r.A == nil {
|
addr, ok := netip.AddrFromSlice(r.A)
|
||||||
return ""
|
return addr.Unmap(), ok
|
||||||
}
|
|
||||||
return r.A.String()
|
|
||||||
case *dns.AAAA:
|
case *dns.AAAA:
|
||||||
if r.AAAA == nil {
|
addr, ok := netip.AddrFromSlice(r.AAAA)
|
||||||
return ""
|
return addr.Unmap(), ok
|
||||||
}
|
|
||||||
return r.AAAA.String()
|
|
||||||
}
|
}
|
||||||
return ""
|
return netip.Addr{}, false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Update replaces all zones and their records
|
// Update replaces all zones and their records
|
||||||
|
|||||||
@@ -37,8 +37,8 @@ type mockPeerConnectivity struct {
|
|||||||
byIP map[string]struct{ known, connected bool }
|
byIP map[string]struct{ known, connected bool }
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m mockPeerConnectivity) IsConnectedByIP(ip string) (known, connected bool) {
|
func (m mockPeerConnectivity) IsConnectedByIP(ip netip.Addr) (known, connected bool) {
|
||||||
v, ok := m.byIP[ip]
|
v, ok := m.byIP[ip.String()]
|
||||||
if !ok {
|
if !ok {
|
||||||
return false, false
|
return false, false
|
||||||
}
|
}
|
||||||
|
|||||||
204
client/internal/dns/local/warmup_test.go
Normal file
204
client/internal/dns/local/warmup_test.go
Normal file
@@ -0,0 +1,204 @@
|
|||||||
|
package local
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/dns/test"
|
||||||
|
nbdns "github.com/netbirdio/netbird/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
// recordingActivator records the addresses it was asked to warm and returns
|
||||||
|
// immediately, so ServeDNS is not blocked by the test.
|
||||||
|
type recordingActivator struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
called bool
|
||||||
|
addrs []netip.Addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *recordingActivator) ActivatePeersByIP(_ context.Context, addrs []netip.Addr) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.called = true
|
||||||
|
r.addrs = append(r.addrs, addrs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func serveA(t *testing.T, resolver *Resolver, name string) *dns.Msg {
|
||||||
|
t.Helper()
|
||||||
|
var resp *dns.Msg
|
||||||
|
w := &test.MockResponseWriter{WriteMsgFunc: func(m *dns.Msg) error { resp = m; return nil }}
|
||||||
|
resolver.ServeDNS(w, new(dns.Msg).SetQuestion(name, dns.TypeA))
|
||||||
|
return resp
|
||||||
|
}
|
||||||
|
|
||||||
|
// serviceZone registers rec in a match-only (non-authoritative) zone, the shape
|
||||||
|
// the synthesized private-service zones arrive in.
|
||||||
|
func serviceZone(t *testing.T, resolver *Resolver, zone string, records ...nbdns.SimpleRecord) {
|
||||||
|
t.Helper()
|
||||||
|
resolver.Update([]nbdns.CustomZone{{
|
||||||
|
Domain: zone,
|
||||||
|
Records: records,
|
||||||
|
NonAuthoritative: true,
|
||||||
|
}})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLocalResolver_WarmsLazyPeerOnResolve(t *testing.T) {
|
||||||
|
// Warm-up fires only for multi-record answers (the HA / round-robin shape of
|
||||||
|
// the synthesized private-service zones), so register two peer targets.
|
||||||
|
const name = "svc.proxy.netbird.cloud."
|
||||||
|
recs := []nbdns.SimpleRecord{
|
||||||
|
{Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"},
|
||||||
|
{Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.8"},
|
||||||
|
}
|
||||||
|
resolver := NewResolver()
|
||||||
|
serviceZone(t, resolver, "proxy.netbird.cloud", recs...)
|
||||||
|
|
||||||
|
act := &recordingActivator{}
|
||||||
|
resolver.SetPeerActivator(act)
|
||||||
|
|
||||||
|
resp := serveA(t, resolver, name)
|
||||||
|
require.NotNil(t, resp, "resolver must answer")
|
||||||
|
require.NotEmpty(t, resp.Answer, "answer must carry the A records")
|
||||||
|
|
||||||
|
act.mu.Lock()
|
||||||
|
defer act.mu.Unlock()
|
||||||
|
assert.True(t, act.called, "activator must be invoked for a multi-record service-zone answer")
|
||||||
|
assert.Contains(t, act.addrs, netip.MustParseAddr("100.64.0.7"), "activator must receive the first peer IP")
|
||||||
|
assert.Contains(t, act.addrs, netip.MustParseAddr("100.64.0.8"), "activator must receive the second peer IP")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLocalResolver_NoWarmupForSingleRecord(t *testing.T) {
|
||||||
|
// A single-record answer does not trigger warm-up; the resolver only warms
|
||||||
|
// multi-record answers.
|
||||||
|
rec := nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"}
|
||||||
|
resolver := NewResolver()
|
||||||
|
serviceZone(t, resolver, "proxy.netbird.cloud", rec)
|
||||||
|
|
||||||
|
act := &recordingActivator{}
|
||||||
|
resolver.SetPeerActivator(act)
|
||||||
|
|
||||||
|
resp := serveA(t, resolver, rec.Name)
|
||||||
|
require.NotNil(t, resp, "resolver must answer")
|
||||||
|
require.NotEmpty(t, resp.Answer, "answer must carry the A record")
|
||||||
|
|
||||||
|
act.mu.Lock()
|
||||||
|
defer act.mu.Unlock()
|
||||||
|
assert.False(t, act.called, "activator must not be invoked for a single-record answer")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLocalResolver_NoActivatorNoWarmup(t *testing.T) {
|
||||||
|
// With no activator wired the resolver behaves exactly as before.
|
||||||
|
rec := nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"}
|
||||||
|
resolver := NewResolver()
|
||||||
|
serviceZone(t, resolver, "proxy.netbird.cloud", rec)
|
||||||
|
|
||||||
|
resp := serveA(t, resolver, rec.Name)
|
||||||
|
require.NotNil(t, resp, "resolver must still answer without an activator")
|
||||||
|
require.NotEmpty(t, resp.Answer, "answer must carry the A record")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLocalResolver_NoWarmupForMissingRecord(t *testing.T) {
|
||||||
|
// A query that resolves to nothing must not invoke the activator (no IPs).
|
||||||
|
resolver := NewResolver()
|
||||||
|
serviceZone(t, resolver, "proxy.netbird.cloud",
|
||||||
|
nbdns.SimpleRecord{Name: "svc.proxy.netbird.cloud.", Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.7"})
|
||||||
|
|
||||||
|
act := &recordingActivator{}
|
||||||
|
resolver.SetPeerActivator(act)
|
||||||
|
|
||||||
|
serveA(t, resolver, "absent.proxy.netbird.cloud.")
|
||||||
|
|
||||||
|
act.mu.Lock()
|
||||||
|
defer act.mu.Unlock()
|
||||||
|
assert.False(t, act.called, "activator must not be invoked when there is no answer")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLocalResolver_NoWarmupInAuthoritativeZone(t *testing.T) {
|
||||||
|
// The account's peer zone is authoritative; resolving a peer's name there
|
||||||
|
// must not wake its lazy connection — warm-up is scoped to match-only
|
||||||
|
// (non-authoritative) zones such as the synthesized private-service zones.
|
||||||
|
// Use a multi-record answer so the authoritative-zone scoping is the only
|
||||||
|
// reason warm-up is skipped, not the single-record guard.
|
||||||
|
const name = "peer.netbird.cloud."
|
||||||
|
recs := []nbdns.SimpleRecord{
|
||||||
|
{Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.9"},
|
||||||
|
{Name: name, Type: 1, Class: nbdns.DefaultClass, TTL: 300, RData: "100.64.0.10"},
|
||||||
|
}
|
||||||
|
resolver := NewResolver()
|
||||||
|
resolver.Update([]nbdns.CustomZone{{
|
||||||
|
Domain: "netbird.cloud",
|
||||||
|
Records: recs,
|
||||||
|
}})
|
||||||
|
|
||||||
|
act := &recordingActivator{}
|
||||||
|
resolver.SetPeerActivator(act)
|
||||||
|
|
||||||
|
resp := serveA(t, resolver, name)
|
||||||
|
require.NotNil(t, resp, "resolver must answer")
|
||||||
|
require.NotEmpty(t, resp.Answer, "answer must carry the A records")
|
||||||
|
|
||||||
|
act.mu.Lock()
|
||||||
|
defer act.mu.Unlock()
|
||||||
|
assert.False(t, act.called, "activator must not be invoked for authoritative-zone answers")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLazyWarmupTimeoutFromEnv(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
value string
|
||||||
|
envSet bool
|
||||||
|
want time.Duration
|
||||||
|
}{
|
||||||
|
{name: "unset uses default", want: defaultLazyWarmupTimeout},
|
||||||
|
{name: "valid overrides", value: "5s", envSet: true, want: 5 * time.Second},
|
||||||
|
{name: "invalid falls back", value: "not-a-duration", envSet: true, want: defaultLazyWarmupTimeout},
|
||||||
|
{name: "negative falls back", value: "-1s", envSet: true, want: defaultLazyWarmupTimeout},
|
||||||
|
{name: "zero falls back", value: "0s", envSet: true, want: defaultLazyWarmupTimeout},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
if tt.envSet {
|
||||||
|
t.Setenv(envLazyWarmupTimeout, tt.value)
|
||||||
|
}
|
||||||
|
assert.Equal(t, tt.want, lazyWarmupTimeoutFromEnv())
|
||||||
|
assert.Equal(t, tt.want, NewResolver().warmupTimeout, "constructor must resolve the timeout once")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExtractRecordAddr(t *testing.T) {
|
||||||
|
t.Run("A record yields unmapped v4", func(t *testing.T) {
|
||||||
|
// net.ParseIP returns the 16-byte v4-in-v6 form, the same shape
|
||||||
|
// miekg/dns stores after parsing an A record; the extracted address
|
||||||
|
// must compare equal to a plain v4 netip.Addr.
|
||||||
|
addr, ok := extractRecordAddr(&dns.A{A: net.ParseIP("100.64.0.7")})
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.True(t, addr.Is4())
|
||||||
|
assert.Equal(t, netip.MustParseAddr("100.64.0.7"), addr)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("AAAA record yields v6", func(t *testing.T) {
|
||||||
|
addr, ok := extractRecordAddr(&dns.AAAA{AAAA: net.ParseIP("fd00::1")})
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Equal(t, netip.MustParseAddr("fd00::1"), addr)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("A record without address", func(t *testing.T) {
|
||||||
|
_, ok := extractRecordAddr(&dns.A{})
|
||||||
|
assert.False(t, ok)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("non-address record", func(t *testing.T) {
|
||||||
|
_, ok := extractRecordAddr(&dns.CNAME{Target: "target.netbird.cloud."})
|
||||||
|
assert.False(t, ok)
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
|
|
||||||
dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config"
|
dnsconfig "github.com/netbirdio/netbird/client/internal/dns/config"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/dns/local"
|
||||||
nbdns "github.com/netbirdio/netbird/dns"
|
nbdns "github.com/netbirdio/netbird/dns"
|
||||||
"github.com/netbirdio/netbird/route"
|
"github.com/netbirdio/netbird/route"
|
||||||
"github.com/netbirdio/netbird/shared/management/domain"
|
"github.com/netbirdio/netbird/shared/management/domain"
|
||||||
@@ -92,6 +93,11 @@ func (m *MockServer) SetFirewall(Firewall) {
|
|||||||
// Mock implementation - no-op
|
// Mock implementation - no-op
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetPeerActivator mock implementation of SetPeerActivator from Server interface
|
||||||
|
func (m *MockServer) SetPeerActivator(local.PeerActivator) {
|
||||||
|
// Mock implementation - no-op
|
||||||
|
}
|
||||||
|
|
||||||
// BeginBatch mock implementation of BeginBatch from Server interface
|
// BeginBatch mock implementation of BeginBatch from Server interface
|
||||||
func (m *MockServer) BeginBatch() {
|
func (m *MockServer) BeginBatch() {
|
||||||
// Mock implementation - no-op
|
// Mock implementation - no-op
|
||||||
|
|||||||
@@ -51,7 +51,5 @@ func (n *notifier) notify() {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
go func(l listener.NetworkChangeListener) {
|
n.listener.OnNetworkChanged("")
|
||||||
l.OnNetworkChanged("")
|
|
||||||
}(n.listener)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -82,6 +82,7 @@ type Server interface {
|
|||||||
PopulateManagementDomain(mgmtURL *url.URL) error
|
PopulateManagementDomain(mgmtURL *url.URL) error
|
||||||
SetRouteSources(selected, active func() route.HAMap)
|
SetRouteSources(selected, active func() route.HAMap)
|
||||||
SetFirewall(Firewall)
|
SetFirewall(Firewall)
|
||||||
|
SetPeerActivator(local.PeerActivator)
|
||||||
}
|
}
|
||||||
|
|
||||||
type nsGroupsByDomain struct {
|
type nsGroupsByDomain struct {
|
||||||
@@ -251,7 +252,7 @@ func NewDefaultServerPermanentUpstream(
|
|||||||
ds.hostsDNSHolder.set(hostsDnsList)
|
ds.hostsDNSHolder.set(hostsDnsList)
|
||||||
ds.permanent = true
|
ds.permanent = true
|
||||||
ds.currentConfig = dnsConfigToHostDNSConfig(config, ds.service.RuntimeIP(), ds.service.RuntimePort())
|
ds.currentConfig = dnsConfigToHostDNSConfig(config, ds.service.RuntimeIP(), ds.service.RuntimePort())
|
||||||
ds.searchDomainNotifier = newNotifier(ds.SearchDomains())
|
ds.searchDomainNotifier = newNotifier(ds.searchDomains())
|
||||||
ds.searchDomainNotifier.setListener(listener)
|
ds.searchDomainNotifier.setListener(listener)
|
||||||
setServerDns(ds)
|
setServerDns(ds)
|
||||||
return ds
|
return ds
|
||||||
@@ -491,6 +492,13 @@ func (s *DefaultServer) SetFirewall(fw Firewall) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetPeerActivator wires the DNS-time lazy-connection warm-up on the local
|
||||||
|
// resolver. Injected after the connection manager exists (it does not at
|
||||||
|
// DNS-server construction time). Pass nil to disable.
|
||||||
|
func (s *DefaultServer) SetPeerActivator(a local.PeerActivator) {
|
||||||
|
s.localResolver.SetPeerActivator(a)
|
||||||
|
}
|
||||||
|
|
||||||
// Stop stops the server
|
// Stop stops the server
|
||||||
func (s *DefaultServer) Stop() {
|
func (s *DefaultServer) Stop() {
|
||||||
s.ctxCancel()
|
s.ctxCancel()
|
||||||
@@ -594,6 +602,12 @@ func (s *DefaultServer) UpdateDNSServer(serial uint64, update nbdns.Config) erro
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *DefaultServer) SearchDomains() []string {
|
func (s *DefaultServer) SearchDomains() []string {
|
||||||
|
s.mux.Lock()
|
||||||
|
defer s.mux.Unlock()
|
||||||
|
return s.searchDomains()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *DefaultServer) searchDomains() []string {
|
||||||
var searchDomains []string
|
var searchDomains []string
|
||||||
|
|
||||||
for _, dConf := range s.currentConfig.Domains {
|
for _, dConf := range s.currentConfig.Domains {
|
||||||
@@ -678,7 +692,7 @@ func (s *DefaultServer) applyConfiguration(update nbdns.Config) error {
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
if s.searchDomainNotifier != nil {
|
if s.searchDomainNotifier != nil {
|
||||||
s.searchDomainNotifier.onNewSearchDomains(s.SearchDomains())
|
s.searchDomainNotifier.onNewSearchDomains(s.searchDomains())
|
||||||
}
|
}
|
||||||
|
|
||||||
s.updateNSGroupStates(update.NameServerGroups)
|
s.updateNSGroupStates(update.NameServerGroups)
|
||||||
@@ -1435,11 +1449,11 @@ type localPeerConnectivity struct {
|
|||||||
|
|
||||||
// IsConnectedByIP looks the IP up in the peerstore and surfaces both
|
// IsConnectedByIP looks the IP up in the peerstore and surfaces both
|
||||||
// the known and connected bits. Used by Resolver.filterDisconnectedPeerAnswers.
|
// the known and connected bits. Used by Resolver.filterDisconnectedPeerAnswers.
|
||||||
func (l localPeerConnectivity) IsConnectedByIP(ip string) (known, connected bool) {
|
func (l localPeerConnectivity) IsConnectedByIP(ip netip.Addr) (known, connected bool) {
|
||||||
if l.status == nil {
|
if l.status == nil {
|
||||||
return false, false
|
return false, false
|
||||||
}
|
}
|
||||||
state, ok := l.status.PeerStateByIP(ip)
|
state, ok := l.status.PeerStateByIP(ip.String())
|
||||||
if !ok {
|
if !ok {
|
||||||
return false, false
|
return false, false
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -292,18 +292,16 @@ func (s *serviceViaListener) generateFreePort() (uint16, error) {
|
|||||||
return customPort, nil
|
return customPort, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
udpAddr := net.UDPAddrFromAddrPort(netip.MustParseAddrPort("0.0.0.0:0"))
|
probeListener, err := net.ListenUDP("udp4", &net.UDPAddr{})
|
||||||
probeListener, err := net.ListenUDP("udp", udpAddr)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Debugf("failed to bind random port for DNS: %s", err)
|
log.Debugf("failed to bind random port for DNS: %s", err)
|
||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
addrPort := netip.MustParseAddrPort(probeListener.LocalAddr().String()) // might panic if address is incorrect
|
port := uint16(probeListener.LocalAddr().(*net.UDPAddr).Port)
|
||||||
err = probeListener.Close()
|
if err = probeListener.Close(); err != nil {
|
||||||
if err != nil {
|
|
||||||
log.Debugf("failed to free up DNS port: %s", err)
|
log.Debugf("failed to free up DNS port: %s", err)
|
||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
return addrPort.Port(), nil
|
return port, nil
|
||||||
}
|
}
|
||||||
|
|||||||
76
client/internal/dns_peer_activator.go
Normal file
76
client/internal/dns_peer_activator.go
Normal file
@@ -0,0 +1,76 @@
|
|||||||
|
package internal
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/netip"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peerstore"
|
||||||
|
)
|
||||||
|
|
||||||
|
const dnsActivationPollInterval = 50 * time.Millisecond
|
||||||
|
|
||||||
|
// dnsPeerActivator wakes lazy-connection peers from the DNS resolution path. It
|
||||||
|
// implements dns/local.PeerActivator. DNS queries run on their own goroutines,
|
||||||
|
// so it only touches state that is safe for concurrent use — ConnMgr.ActivatePeer,
|
||||||
|
// peerstore.Store and peer.Status — and never takes the engine's syncMsgMux,
|
||||||
|
// keeping DNS resolution from contending with network-map processing.
|
||||||
|
type dnsPeerActivator struct {
|
||||||
|
connMgr *ConnMgr
|
||||||
|
peerStore *peerstore.Store
|
||||||
|
status *peer.Status
|
||||||
|
// ctx is the engine's long-lived context. The connection dial is tied to it
|
||||||
|
// (not the per-query DNS wait budget) so a handshake that outlasts the wait
|
||||||
|
// still completes in the background rather than being cancelled at the deadline.
|
||||||
|
ctx context.Context
|
||||||
|
}
|
||||||
|
|
||||||
|
// ActivatePeersByIP triggers wake-up for the peer(s) owning addrs and waits
|
||||||
|
// until one is connected or ctx (the per-query DNS wait budget) expires.
|
||||||
|
// Activation itself is tied to the engine's long-lived context so the dial
|
||||||
|
// survives a wait that times out. Unknown or already-connected addresses are
|
||||||
|
// skipped, so the steady-state (warm) path adds no latency.
|
||||||
|
func (a *dnsPeerActivator) ActivatePeersByIP(ctx context.Context, addrs []netip.Addr) {
|
||||||
|
if a == nil || a.connMgr == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var pending []string
|
||||||
|
for _, addr := range addrs {
|
||||||
|
ip := addr.String()
|
||||||
|
st, ok := a.status.PeerStateByIP(ip)
|
||||||
|
if !ok || st.ConnStatus == peer.StatusConnected {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
conn, ok := a.peerStore.PeerConn(st.PubKey)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
a.connMgr.ActivatePeer(a.ctx, conn)
|
||||||
|
pending = append(pending, ip)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(pending) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
a.waitConnected(ctx, pending)
|
||||||
|
}
|
||||||
|
|
||||||
|
// waitConnected blocks until any of ips reports a connected peer or ctx expires.
|
||||||
|
func (a *dnsPeerActivator) waitConnected(ctx context.Context, ips []string) {
|
||||||
|
ticker := time.NewTicker(dnsActivationPollInterval)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
for _, ip := range ips {
|
||||||
|
if st, ok := a.status.PeerStateByIP(ip); ok && st.ConnStatus == peer.StatusConnected {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
129
client/internal/dns_peer_activator_test.go
Normal file
129
client/internal/dns_peer_activator_test.go
Normal file
@@ -0,0 +1,129 @@
|
|||||||
|
package internal
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peerstore"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestPeerConn(t *testing.T, key string) *peer.Conn {
|
||||||
|
t.Helper()
|
||||||
|
conn, err := peer.NewConn(peer.ConnConfig{
|
||||||
|
Key: key,
|
||||||
|
LocalKey: "local",
|
||||||
|
WgConfig: peer.WgConfig{
|
||||||
|
AllowedIps: []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")},
|
||||||
|
},
|
||||||
|
}, peer.ServiceDependencies{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestDNSPeerActivator(t *testing.T) (*dnsPeerActivator, *peer.Status, *peerstore.Store) {
|
||||||
|
t.Helper()
|
||||||
|
status := peer.NewRecorder("https://mgm")
|
||||||
|
store := peerstore.NewConnStore()
|
||||||
|
// ConnMgr without Start: the lazy manager is nil, so ActivatePeer is a
|
||||||
|
// no-op — these tests exercise the activator's skip/wait logic.
|
||||||
|
connMgr := NewConnMgr(&EngineConfig{}, status, store, nil)
|
||||||
|
return &dnsPeerActivator{
|
||||||
|
connMgr: connMgr,
|
||||||
|
peerStore: store,
|
||||||
|
status: status,
|
||||||
|
ctx: context.Background(),
|
||||||
|
}, status, store
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDNSPeerActivator_NilSafe(t *testing.T) {
|
||||||
|
var a *dnsPeerActivator
|
||||||
|
a.ActivatePeersByIP(context.Background(), []netip.Addr{netip.MustParseAddr("100.64.0.1")})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDNSPeerActivator_SkipsUnknownAndConnectedPeers verifies the steady-state
|
||||||
|
// (warm) path adds no latency: already-connected and unknown addresses never
|
||||||
|
// enter the wait loop.
|
||||||
|
func TestDNSPeerActivator_SkipsUnknownAndConnectedPeers(t *testing.T) {
|
||||||
|
a, status, store := newTestDNSPeerActivator(t)
|
||||||
|
|
||||||
|
require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", "fd00::1"))
|
||||||
|
require.NoError(t, status.UpdatePeerState(peer.State{PubKey: "peerA", ConnStatus: peer.StatusConnected}))
|
||||||
|
store.AddPeerConn("peerA", newTestPeerConn(t, "peerA"))
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
a.ActivatePeersByIP(ctx, []netip.Addr{
|
||||||
|
netip.MustParseAddr("100.64.0.1"), // known, connected -> skipped
|
||||||
|
netip.MustParseAddr("fd00::1"), // known via IPv6, connected -> skipped
|
||||||
|
netip.MustParseAddr("100.64.0.99"), // unknown -> skipped
|
||||||
|
})
|
||||||
|
require.Less(t, time.Since(start), time.Second, "no pending peer must mean no wait")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDNSPeerActivator_WaitsForPendingPeerToConnect verifies the wait loop
|
||||||
|
// returns as soon as a pending peer reports connected, well before the
|
||||||
|
// per-query budget expires.
|
||||||
|
func TestDNSPeerActivator_WaitsForPendingPeerToConnect(t *testing.T) {
|
||||||
|
a, status, store := newTestDNSPeerActivator(t)
|
||||||
|
|
||||||
|
require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", ""))
|
||||||
|
store.AddPeerConn("peerA", newTestPeerConn(t, "peerA"))
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
time.Sleep(150 * time.Millisecond)
|
||||||
|
_ = status.UpdatePeerState(peer.State{PubKey: "peerA", ConnStatus: peer.StatusConnected})
|
||||||
|
}()
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")})
|
||||||
|
elapsed := time.Since(start)
|
||||||
|
|
||||||
|
require.GreaterOrEqual(t, elapsed, 100*time.Millisecond, "must wait for the pending peer")
|
||||||
|
require.Less(t, elapsed, 5*time.Second, "must return on connect, not at the deadline")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDNSPeerActivator_ReturnsAtBudgetWhenPeerStaysIdle verifies a peer that
|
||||||
|
// never connects releases the DNS response at the per-query budget instead of
|
||||||
|
// blocking it indefinitely.
|
||||||
|
func TestDNSPeerActivator_ReturnsAtBudgetWhenPeerStaysIdle(t *testing.T) {
|
||||||
|
a, status, store := newTestDNSPeerActivator(t)
|
||||||
|
|
||||||
|
require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", ""))
|
||||||
|
store.AddPeerConn("peerA", newTestPeerConn(t, "peerA"))
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")})
|
||||||
|
elapsed := time.Since(start)
|
||||||
|
|
||||||
|
require.GreaterOrEqual(t, elapsed, 250*time.Millisecond, "must wait out the budget for a pending peer")
|
||||||
|
require.Less(t, elapsed, 5*time.Second, "must not block past the budget")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDNSPeerActivator_NoWaitWithoutPeerConn verifies a known-but-idle peer
|
||||||
|
// with no connection object in the store is not waited on: there is nothing to
|
||||||
|
// activate, so waiting could only ever time out.
|
||||||
|
func TestDNSPeerActivator_NoWaitWithoutPeerConn(t *testing.T) {
|
||||||
|
a, status, _ := newTestDNSPeerActivator(t)
|
||||||
|
|
||||||
|
require.NoError(t, status.AddPeer("peerA", "a.netbird.cloud", "100.64.0.1", ""))
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
a.ActivatePeersByIP(ctx, []netip.Addr{netip.MustParseAddr("100.64.0.1")})
|
||||||
|
require.Less(t, time.Since(start), time.Second, "peer without a conn must not be waited on")
|
||||||
|
}
|
||||||
Binary file not shown.
Binary file not shown.
@@ -52,11 +52,14 @@ int xdp_dns_fwd(struct iphdr *ip, struct udphdr *udp) {
|
|||||||
|
|
||||||
if (udp->dest == GENERAL_DNS_PORT && ip->daddr == dns_ip) {
|
if (udp->dest == GENERAL_DNS_PORT && ip->daddr == dns_ip) {
|
||||||
udp->dest = dns_port;
|
udp->dest = dns_port;
|
||||||
|
// Clear the now-stale checksum; zero means "not computed" for IPv4.
|
||||||
|
udp->check = 0;
|
||||||
return XDP_PASS;
|
return XDP_PASS;
|
||||||
}
|
}
|
||||||
|
|
||||||
if (udp->source == dns_port && ip->saddr == dns_ip) {
|
if (udp->source == dns_port && ip->saddr == dns_ip) {
|
||||||
udp->source = GENERAL_DNS_PORT;
|
udp->source = GENERAL_DNS_PORT;
|
||||||
|
udp->check = 0;
|
||||||
return XDP_PASS;
|
return XDP_PASS;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -50,5 +50,11 @@ int xdp_wg_proxy(struct iphdr *ip, struct udphdr *udp) {
|
|||||||
__be16 new_dst_port = htons(proxy_port);
|
__be16 new_dst_port = htons(proxy_port);
|
||||||
udp->dest = new_dst_port;
|
udp->dest = new_dst_port;
|
||||||
udp->source = new_src_port;
|
udp->source = new_src_port;
|
||||||
|
|
||||||
|
// The ports are covered by the UDP checksum. This is an IPv4 loopback hop
|
||||||
|
// and the payload is already integrity-protected, so clear the checksum (a
|
||||||
|
// zero UDP checksum means "not computed" for IPv4) rather than leave a
|
||||||
|
// stale value the kernel would drop as UDP_CSUM.
|
||||||
|
udp->check = 0;
|
||||||
return XDP_PASS;
|
return XDP_PASS;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -64,7 +64,10 @@ import (
|
|||||||
"github.com/netbirdio/netbird/route"
|
"github.com/netbirdio/netbird/route"
|
||||||
mgm "github.com/netbirdio/netbird/shared/management/client"
|
mgm "github.com/netbirdio/netbird/shared/management/client"
|
||||||
"github.com/netbirdio/netbird/shared/management/domain"
|
"github.com/netbirdio/netbird/shared/management/domain"
|
||||||
|
sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc"
|
||||||
|
nbnetworkmap "github.com/netbirdio/netbird/shared/management/networkmap"
|
||||||
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
||||||
|
types "github.com/netbirdio/netbird/shared/management/types"
|
||||||
"github.com/netbirdio/netbird/shared/netiputil"
|
"github.com/netbirdio/netbird/shared/netiputil"
|
||||||
auth "github.com/netbirdio/netbird/shared/relay/auth/hmac"
|
auth "github.com/netbirdio/netbird/shared/relay/auth/hmac"
|
||||||
relayClient "github.com/netbirdio/netbird/shared/relay/client"
|
relayClient "github.com/netbirdio/netbird/shared/relay/client"
|
||||||
@@ -147,6 +150,7 @@ type EngineConfig struct {
|
|||||||
BlockLANAccess bool
|
BlockLANAccess bool
|
||||||
BlockInbound bool
|
BlockInbound bool
|
||||||
DisableIPv6 bool
|
DisableIPv6 bool
|
||||||
|
SyncMessageVersion *int
|
||||||
|
|
||||||
// LazyConnection is the MDM-sourced lazy-connection override; StateUnset defers to
|
// LazyConnection is the MDM-sourced lazy-connection override; StateUnset defers to
|
||||||
// the env var and management feature flag.
|
// the env var and management feature flag.
|
||||||
@@ -220,11 +224,12 @@ type Engine struct {
|
|||||||
// networkSerial is the latest CurrentSerial (state ID) of the network sent by the Management service
|
// networkSerial is the latest CurrentSerial (state ID) of the network sent by the Management service
|
||||||
networkSerial uint64
|
networkSerial uint64
|
||||||
|
|
||||||
// forwardingRules holds the ingress forward rules applied for the current target.
|
// latestComponents is the most-recent NetworkMapComponents decoded from
|
||||||
// Wholesale sections (incl. forward rules) run only on the first pass of a target;
|
// a NetworkMapEnvelope (capability=3 peers only). Held alongside the
|
||||||
// it is stashed here so the final, peer-converged pass can build the lazy-connection
|
// NetworkMap that Calculate() produced from it so future incremental
|
||||||
// exclude list without recomputing them on every bounded peer pass.
|
// updates have a base to apply changes against. nil for legacy-format
|
||||||
forwardingRules []firewallManager.ForwardRule
|
// peers. Guarded by syncMsgMux.
|
||||||
|
latestComponents *types.NetworkMapComponents
|
||||||
|
|
||||||
networkMonitor *networkmonitor.NetworkMonitor
|
networkMonitor *networkmonitor.NetworkMonitor
|
||||||
|
|
||||||
@@ -280,6 +285,20 @@ type Engine struct {
|
|||||||
jobExecutorWG sync.WaitGroup
|
jobExecutorWG sync.WaitGroup
|
||||||
|
|
||||||
exposeManager *expose.Manager
|
exposeManager *expose.Manager
|
||||||
|
|
||||||
|
sessionWatcher sessionDeadlineWatcher
|
||||||
|
}
|
||||||
|
|
||||||
|
// sessionDeadlineWatcher is the engine-facing surface of the SSO session
|
||||||
|
// expiry watcher. The concrete implementation (sessionwatch.Watcher) is wired
|
||||||
|
// in via newSessionWatcher, which is build-tagged so the js/wasm build links a
|
||||||
|
// no-op stub instead of pulling the full sessionwatch package (and its timer
|
||||||
|
// machinery) into the binary — the wasm client never runs the engine's
|
||||||
|
// session-warning flow.
|
||||||
|
type sessionDeadlineWatcher interface {
|
||||||
|
Update(deadline time.Time) error
|
||||||
|
Dismiss()
|
||||||
|
Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Peer is an instance of the Connection Peer
|
// Peer is an instance of the Connection Peer
|
||||||
@@ -331,6 +350,17 @@ func NewEngine(
|
|||||||
updateManager: services.UpdateManager,
|
updateManager: services.UpdateManager,
|
||||||
syncStoreDir: config.StateDir,
|
syncStoreDir: config.StateDir,
|
||||||
}
|
}
|
||||||
|
// sessionWatcher keeps the SubscribeStatus consumers in sync with the
|
||||||
|
// session expiry deadline. Deadline-change ticks come for free via
|
||||||
|
// Status.SetSessionExpiresAt; the watcher exists to push a wake-up at
|
||||||
|
// T-WarningLead and T-FinalWarningLead so the UI repaints the remaining
|
||||||
|
// time / warning state even when nothing else changed, and to publish
|
||||||
|
// two SystemEvents (the warning composition lives in sessionwatch so
|
||||||
|
// the wire format stays owned by one package):
|
||||||
|
// - T-WarningLead → interactive "Extend now / Dismiss" notification
|
||||||
|
// - T-FinalWarningLead → auto-opened SessionAboutToExpire dialog,
|
||||||
|
// suppressed when the user dismissed the earlier warning
|
||||||
|
engine.sessionWatcher = newSessionWatcher(engine.statusRecorder)
|
||||||
|
|
||||||
log.Infof("I am: %s", config.WgPrivateKey.PublicKey().String())
|
log.Infof("I am: %s", config.WgPrivateKey.PublicKey().String())
|
||||||
return engine
|
return engine
|
||||||
@@ -397,6 +427,10 @@ func (e *Engine) stopLocked() {
|
|||||||
e.srWatcher.Close()
|
e.srWatcher.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if e.sessionWatcher != nil {
|
||||||
|
e.sessionWatcher.Close()
|
||||||
|
}
|
||||||
|
|
||||||
if e.updateManager != nil {
|
if e.updateManager != nil {
|
||||||
e.updateManager.SetDownloadOnly()
|
e.updateManager.SetDownloadOnly()
|
||||||
}
|
}
|
||||||
@@ -528,7 +562,7 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
|
|||||||
} else {
|
} else {
|
||||||
log.Infof("running rosenpass in strict mode")
|
log.Infof("running rosenpass in strict mode")
|
||||||
}
|
}
|
||||||
e.rpManager, err = rosenpass.NewManager(e.config.PreSharedKey, e.config.WgIfaceName)
|
e.rpManager, err = rosenpass.NewManager(e.config.PreSharedKey, e.config.WgIfaceName, publicKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("create rosenpass manager: %w", err)
|
return fmt.Errorf("create rosenpass manager: %w", err)
|
||||||
}
|
}
|
||||||
@@ -538,12 +572,7 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
|
|||||||
}
|
}
|
||||||
e.stateManager.Start()
|
e.stateManager.Start()
|
||||||
|
|
||||||
initialRoutes, dnsConfig, dnsFeatureFlag, err := e.readInitialSettings()
|
dnsServer, err := e.newDnsServer()
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("read initial settings: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
dnsServer, err := e.newDnsServer(dnsConfig)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("create dns server: %w", err)
|
return fmt.Errorf("create dns server: %w", err)
|
||||||
}
|
}
|
||||||
@@ -561,10 +590,8 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
|
|||||||
WGInterface: e.wgInterface,
|
WGInterface: e.wgInterface,
|
||||||
StatusRecorder: e.statusRecorder,
|
StatusRecorder: e.statusRecorder,
|
||||||
RelayManager: e.relayManager,
|
RelayManager: e.relayManager,
|
||||||
InitialRoutes: initialRoutes,
|
|
||||||
StateManager: e.stateManager,
|
StateManager: e.stateManager,
|
||||||
DNSServer: dnsServer,
|
DNSServer: dnsServer,
|
||||||
DNSFeatureFlag: dnsFeatureFlag,
|
|
||||||
PeerStore: e.peerStore,
|
PeerStore: e.peerStore,
|
||||||
DisableClientRoutes: e.config.DisableClientRoutes,
|
DisableClientRoutes: e.config.DisableClientRoutes,
|
||||||
DisableServerRoutes: e.config.DisableServerRoutes,
|
DisableServerRoutes: e.config.DisableServerRoutes,
|
||||||
@@ -629,8 +656,24 @@ func (e *Engine) Start(netbirdConfig *mgmProto.NetbirdConfig, mgmtURL *url.URL)
|
|||||||
iceCfg := e.createICEConfig()
|
iceCfg := e.createICEConfig()
|
||||||
|
|
||||||
e.connMgr = NewConnMgr(e.config, e.statusRecorder, e.peerStore, wgIface)
|
e.connMgr = NewConnMgr(e.config, e.statusRecorder, e.peerStore, wgIface)
|
||||||
|
e.connMgr.SetRoutedIPsReconciler(func(peerKey string) error {
|
||||||
|
if e.routeManager == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return e.routeManager.ReconcilePeerAllowedIPs(peerKey)
|
||||||
|
})
|
||||||
e.connMgr.Start(e.ctx)
|
e.connMgr.Start(e.ctx)
|
||||||
|
|
||||||
|
// Wire DNS-time lazy-connection warm-up now that the connection manager
|
||||||
|
// exists (it does not at DNS-server construction time). A DNS answer that
|
||||||
|
// points at an idle peer then wakes it before the client's first request.
|
||||||
|
e.dnsServer.SetPeerActivator(&dnsPeerActivator{
|
||||||
|
connMgr: e.connMgr,
|
||||||
|
peerStore: e.peerStore,
|
||||||
|
status: e.statusRecorder,
|
||||||
|
ctx: e.ctx,
|
||||||
|
})
|
||||||
|
|
||||||
e.srWatcher = guard.NewSRWatcher(e.signal, e.relayManager, e.mobileDep.IFaceDiscover, iceCfg)
|
e.srWatcher = guard.NewSRWatcher(e.signal, e.relayManager, e.mobileDep.IFaceDiscover, iceCfg)
|
||||||
e.srWatcher.Start(peer.IsForceRelayed())
|
e.srWatcher.Start(peer.IsForceRelayed())
|
||||||
|
|
||||||
@@ -780,15 +823,7 @@ func (e *Engine) blockLanAccess() {
|
|||||||
|
|
||||||
// modifyPeers updates peers that have been modified (e.g. IP address has been changed).
|
// modifyPeers updates peers that have been modified (e.g. IP address has been changed).
|
||||||
// It closes the existing connection, removes it from the peerConns map, and creates a new one.
|
// It closes the existing connection, removes it from the peerConns map, and creates a new one.
|
||||||
// maxPeersPerSyncPass is the default per-pass cap on how many peers each of
|
func (e *Engine) modifyPeers(peersUpdate []*mgmProto.RemotePeerConfig) error {
|
||||||
// removePeers/modifyPeers/addNewPeers applies, so syncMsgMux is held only for a
|
|
||||||
// batch at a time and other subsystems can interleave between passes. It is
|
|
||||||
// passed in (not read globally) so tests can exercise the multi-pass path.
|
|
||||||
const maxPeersPerSyncPass = 300
|
|
||||||
|
|
||||||
// modifyPeers re-applies up to maxBatch changed peers per call. It returns true
|
|
||||||
// when more changed peers remained than the cap, so the caller re-runs.
|
|
||||||
func (e *Engine) modifyPeers(peersUpdate []*mgmProto.RemotePeerConfig, maxBatch int) (bool, error) {
|
|
||||||
|
|
||||||
// first, check if peers have been modified
|
// first, check if peers have been modified
|
||||||
var modified []*mgmProto.RemotePeerConfig
|
var modified []*mgmProto.RemotePeerConfig
|
||||||
@@ -818,32 +853,26 @@ func (e *Engine) modifyPeers(peersUpdate []*mgmProto.RemotePeerConfig, maxBatch
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
more := false
|
|
||||||
if len(modified) > maxBatch {
|
|
||||||
modified = modified[:maxBatch]
|
|
||||||
more = true
|
|
||||||
}
|
|
||||||
|
|
||||||
// second, close all modified connections and remove them from the state map
|
// second, close all modified connections and remove them from the state map
|
||||||
for _, p := range modified {
|
for _, p := range modified {
|
||||||
if err := e.removePeer(p.GetWgPubKey()); err != nil {
|
err := e.removePeer(p.GetWgPubKey())
|
||||||
return false, err
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// third, add the peer connections again
|
// third, add the peer connections again
|
||||||
for _, p := range modified {
|
for _, p := range modified {
|
||||||
if err := e.addNewPeer(p); err != nil {
|
err := e.addNewPeer(p)
|
||||||
return false, err
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return more, nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// removePeers finds and removes peers that do not exist anymore in the network map received from the Management Service.
|
// removePeers finds and removes peers that do not exist anymore in the network map received from the Management Service.
|
||||||
// It also removes peers that have been modified (e.g. change of IP address). They will be added again in addPeers method.
|
// It also removes peers that have been modified (e.g. change of IP address). They will be added again in addPeers method.
|
||||||
// removePeers removes up to maxBatch peers per call. It returns true when more
|
func (e *Engine) removePeers(peersUpdate []*mgmProto.RemotePeerConfig) error {
|
||||||
// peers remained to remove than the cap, so the caller re-runs.
|
|
||||||
func (e *Engine) removePeers(peersUpdate []*mgmProto.RemotePeerConfig, maxBatch int) (bool, error) {
|
|
||||||
newPeers := make([]string, 0, len(peersUpdate))
|
newPeers := make([]string, 0, len(peersUpdate))
|
||||||
for _, p := range peersUpdate {
|
for _, p := range peersUpdate {
|
||||||
newPeers = append(newPeers, p.GetWgPubKey())
|
newPeers = append(newPeers, p.GetWgPubKey())
|
||||||
@@ -851,19 +880,14 @@ func (e *Engine) removePeers(peersUpdate []*mgmProto.RemotePeerConfig, maxBatch
|
|||||||
|
|
||||||
toRemove := util.SliceDiff(e.peerStore.PeersPubKey(), newPeers)
|
toRemove := util.SliceDiff(e.peerStore.PeersPubKey(), newPeers)
|
||||||
|
|
||||||
more := false
|
|
||||||
if len(toRemove) > maxBatch {
|
|
||||||
toRemove = toRemove[:maxBatch]
|
|
||||||
more = true
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, p := range toRemove {
|
for _, p := range toRemove {
|
||||||
if err := e.removePeer(p); err != nil {
|
err := e.removePeer(p)
|
||||||
return false, err
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
log.Infof("removed peer %s", p)
|
log.Infof("removed peer %s", p)
|
||||||
}
|
}
|
||||||
return more, nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (e *Engine) removeAllPeers() error {
|
func (e *Engine) removeAllPeers() error {
|
||||||
@@ -942,28 +966,72 @@ func (e *Engine) phase(name string) func() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// applySyncPass applies one bounded pass of the sync update under syncMsgMux and
|
func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
|
||||||
// returns true if more peers remained than the per-pass cap. It is driven by the
|
started := time.Now()
|
||||||
// mapStateManager, which re-invokes it (releasing the lock between passes) until
|
defer func() {
|
||||||
// the update is fully applied.
|
duration := time.Since(started)
|
||||||
func (e *Engine) applySyncPass(update *mgmProto.SyncResponse, firstPass bool) (bool, error) {
|
log.Infof("sync finished in %s", duration)
|
||||||
|
e.clientMetrics.RecordSyncDuration(e.ctx, duration)
|
||||||
|
}()
|
||||||
e.syncMsgMux.Lock()
|
e.syncMsgMux.Lock()
|
||||||
defer e.syncMsgMux.Unlock()
|
defer e.syncMsgMux.Unlock()
|
||||||
|
|
||||||
// Check context INSIDE lock to ensure atomicity with shutdown
|
// Check context INSIDE lock to ensure atomicity with shutdown
|
||||||
if e.ctx.Err() != nil {
|
if e.ctx.Err() != nil {
|
||||||
return false, e.ctx.Err()
|
return e.ctx.Err()
|
||||||
}
|
}
|
||||||
|
|
||||||
if update.NetworkMap != nil && update.NetworkMap.PeerConfig != nil {
|
e.ApplySessionDeadline(update.GetSessionExpiresAt())
|
||||||
e.handleAutoUpdateVersion(update.NetworkMap.PeerConfig.AutoUpdate)
|
|
||||||
|
// Envelope sync responses carry PeerConfig at the top level; legacy
|
||||||
|
// NetworkMap syncs carry it under NetworkMap.PeerConfig.
|
||||||
|
if pc := update.GetPeerConfig(); pc != nil {
|
||||||
|
e.handleAutoUpdateVersion(pc.GetAutoUpdate())
|
||||||
|
} else if nm := update.GetNetworkMap(); nm != nil && nm.GetPeerConfig() != nil {
|
||||||
|
e.handleAutoUpdateVersion(nm.GetPeerConfig().GetAutoUpdate())
|
||||||
}
|
}
|
||||||
|
|
||||||
done := e.phase("netbird_config")
|
done := e.phase("netbird_config")
|
||||||
err := e.updateNetbirdConfig(update.GetNetbirdConfig())
|
err := e.updateNetbirdConfig(update.GetNetbirdConfig())
|
||||||
done()
|
done()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Decode the network map from either the components envelope or the
|
||||||
|
// legacy proto.NetworkMap before the posture-check gating below, so the
|
||||||
|
// "is there a network map" decision covers both wire shapes.
|
||||||
|
var (
|
||||||
|
nm *mgmProto.NetworkMap
|
||||||
|
components *types.NetworkMapComponents
|
||||||
|
)
|
||||||
|
if version := update.GetVersion(); version == int32(sharedgrpc.ComponentNetworkMap) {
|
||||||
|
// Components-format peer: decode the envelope back to typed
|
||||||
|
// components, run Calculate() locally, and convert to the wire
|
||||||
|
// NetworkMap shape the rest of the engine consumes. Components are
|
||||||
|
// retained so future incremental updates can apply deltas instead
|
||||||
|
// of doing a full reconstruction.
|
||||||
|
envelope := update.GetNetworkMapEnvelope()
|
||||||
|
if envelope == nil {
|
||||||
|
return fmt.Errorf("received a SyncReponse indicating use of components network map, but components are missing")
|
||||||
|
}
|
||||||
|
|
||||||
|
localKey := e.config.WgPrivateKey.PublicKey().String()
|
||||||
|
dnsName := ""
|
||||||
|
if pc := update.GetPeerConfig(); pc != nil {
|
||||||
|
// PeerConfig.Fqdn = "<dns_label>.<dns_domain>" — extract the
|
||||||
|
// shared domain by stripping the peer's own label prefix. Falls
|
||||||
|
// back to empty if the FQDN doesn't have the expected shape.
|
||||||
|
dnsName = extractDNSDomainFromFQDN(pc.GetFqdn())
|
||||||
|
}
|
||||||
|
result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("decode network map envelope: %w", err)
|
||||||
|
}
|
||||||
|
nm = result.NetworkMap
|
||||||
|
components = result.Components
|
||||||
|
} else {
|
||||||
|
nm = update.GetNetworkMap()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Posture checks are bound to the network map presence:
|
// Posture checks are bound to the network map presence:
|
||||||
@@ -971,27 +1039,50 @@ func (e *Engine) applySyncPass(update *mgmProto.SyncResponse, firstPass bool) (b
|
|||||||
// NetworkMap != nil, checks nil -> posture checks were removed, clear them
|
// NetworkMap != nil, checks nil -> posture checks were removed, clear them
|
||||||
// NetworkMap == nil -> config-only update (e.g. relay token rotation),
|
// NetworkMap == nil -> config-only update (e.g. relay token rotation),
|
||||||
// leave the previously applied checks untouched
|
// leave the previously applied checks untouched
|
||||||
nm := update.GetNetworkMap()
|
|
||||||
if nm == nil {
|
if nm == nil {
|
||||||
return false, nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
done = e.phase("checks")
|
done = e.phase("checks")
|
||||||
err = e.updateChecksIfNew(update.Checks)
|
err = e.updateChecksIfNew(update.Checks)
|
||||||
done()
|
done()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
done = e.phase("persist")
|
||||||
|
// Only retain the components view when the server sent the envelope
|
||||||
|
// path. A legacy proto.NetworkMap means components == nil; writing it
|
||||||
|
// here would clobber a previously-cached snapshot, breaking the
|
||||||
|
// incremental-delta base on a future envelope sync.
|
||||||
|
if components != nil {
|
||||||
|
e.latestComponents = components
|
||||||
|
}
|
||||||
|
|
||||||
|
e.persistSyncResponse(update)
|
||||||
|
done()
|
||||||
|
|
||||||
// only apply new changes and ignore old ones
|
// only apply new changes and ignore old ones
|
||||||
more, err := e.updateNetworkMap(nm, maxPeersPerSyncPass, firstPass)
|
if err := e.updateNetworkMap(nm); err != nil {
|
||||||
if err != nil {
|
return err
|
||||||
return false, err
|
|
||||||
}
|
}
|
||||||
|
|
||||||
e.statusRecorder.PublishEvent(cProto.SystemEvent_INFO, cProto.SystemEvent_SYSTEM, "Network map updated", "", nil)
|
e.statusRecorder.PublishEvent(cProto.SystemEvent_INFO, cProto.SystemEvent_SYSTEM, "Network map updated", "", nil)
|
||||||
|
|
||||||
return more, nil
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractDNSDomainFromFQDN returns the trailing dotted domain part of the
|
||||||
|
// receiving peer's FQDN — the same value the management server fills as
|
||||||
|
// dnsName when it builds the legacy NetworkMap. "peer42.netbird.cloud" →
|
||||||
|
// "netbird.cloud". An empty string is returned for unrecognized formats.
|
||||||
|
func extractDNSDomainFromFQDN(fqdn string) string {
|
||||||
|
for i := 0; i < len(fqdn); i++ {
|
||||||
|
if fqdn[i] == '.' && i+1 < len(fqdn) {
|
||||||
|
return fqdn[i+1:]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
// updateNetbirdConfig applies the management-provided NetBird configuration:
|
// updateNetbirdConfig applies the management-provided NetBird configuration:
|
||||||
@@ -1039,13 +1130,6 @@ func (e *Engine) updateNetbirdConfig(wCfg *mgmProto.NetbirdConfig) error {
|
|||||||
// (not syncMsgMux) is held for the whole Set so the store cannot be cleared (disabled /
|
// (not syncMsgMux) is held for the whole Set so the store cannot be cleared (disabled /
|
||||||
// engine close) mid-call and have this write resurrect a file that was just removed.
|
// engine close) mid-call and have this write resurrect a file that was just removed.
|
||||||
func (e *Engine) persistSyncResponse(update *mgmProto.SyncResponse) {
|
func (e *Engine) persistSyncResponse(update *mgmProto.SyncResponse) {
|
||||||
// Only persist updates that carry a network map. Config-only updates (e.g. relay
|
|
||||||
// token rotation, STUN/TURN) have a nil NetworkMap; persisting them would overwrite
|
|
||||||
// the last full map on disk and break restore-on-restart.
|
|
||||||
if update.GetNetworkMap() == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
e.syncRespMux.RLock()
|
e.syncRespMux.RLock()
|
||||||
defer e.syncRespMux.RUnlock()
|
defer e.syncRespMux.RUnlock()
|
||||||
|
|
||||||
@@ -1160,6 +1244,7 @@ func (e *Engine) applyInfoFlags(info *system.Info) {
|
|||||||
e.config.BlockLANAccess,
|
e.config.BlockLANAccess,
|
||||||
e.config.BlockInbound,
|
e.config.BlockInbound,
|
||||||
e.config.DisableIPv6,
|
e.config.DisableIPv6,
|
||||||
|
e.config.SyncMessageVersion,
|
||||||
e.config.EnableSSHRoot,
|
e.config.EnableSSHRoot,
|
||||||
e.config.EnableSSHSFTP,
|
e.config.EnableSSHSFTP,
|
||||||
e.config.EnableSSHLocalPortForwarding,
|
e.config.EnableSSHLocalPortForwarding,
|
||||||
@@ -1294,7 +1379,7 @@ func (e *Engine) handleBundle(params *mgmProto.BundleParameters) (*mgmProto.JobR
|
|||||||
ClientMetrics: e.clientMetrics,
|
ClientMetrics: e.clientMetrics,
|
||||||
DaemonVersion: version.NetbirdVersion(),
|
DaemonVersion: version.NetbirdVersion(),
|
||||||
RefreshStatus: func() {
|
RefreshStatus: func() {
|
||||||
e.RunHealthProbes(true)
|
e.RunHealthProbes(e.ctx, true)
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1333,24 +1418,7 @@ func (e *Engine) receiveManagementEvents() {
|
|||||||
}
|
}
|
||||||
e.applyInfoFlags(info)
|
e.applyInfoFlags(info)
|
||||||
|
|
||||||
// The map-state manager converges the latest update in the background in
|
err := e.mgmClient.Sync(e.ctx, info, e.handleSync)
|
||||||
// bounded passes; the stream callback only hands it the newest target.
|
|
||||||
persist := func(u *mgmProto.SyncResponse) {
|
|
||||||
done := e.phase("persist")
|
|
||||||
e.persistSyncResponse(u)
|
|
||||||
done()
|
|
||||||
}
|
|
||||||
manager := newMapStateManager(e.applySyncPass, persist, func(d time.Duration) {
|
|
||||||
log.Infof("sync finished in %s", d)
|
|
||||||
e.clientMetrics.RecordSyncDuration(e.ctx, d)
|
|
||||||
})
|
|
||||||
e.shutdownWg.Add(1)
|
|
||||||
go func() {
|
|
||||||
defer e.shutdownWg.Done()
|
|
||||||
manager.run(e.ctx)
|
|
||||||
}()
|
|
||||||
|
|
||||||
err := e.mgmClient.Sync(e.ctx, info, manager.SetTarget)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// happens if management is unavailable for a long time.
|
// happens if management is unavailable for a long time.
|
||||||
// We want to cancel the operation of the whole client
|
// We want to cancel the operation of the whole client
|
||||||
@@ -1401,107 +1469,21 @@ func (e *Engine) updateTURNs(turns []*mgmProto.ProtectedHostConfig) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// updateNetworkMap applies the wholesale parts (config, routes, ACL, DNS) in full
|
func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap) error {
|
||||||
// and up to maxBatch peers per phase. It returns true when more peers remained
|
|
||||||
// than the cap, so the caller re-runs until convergence.
|
|
||||||
func (e *Engine) updateNetworkMap(networkMap *mgmProto.NetworkMap, maxBatch int, firstPass bool) (bool, error) {
|
|
||||||
// intentionally leave it before checking serial because for now it can happen that peer IP changed but serial didn't
|
// intentionally leave it before checking serial because for now it can happen that peer IP changed but serial didn't
|
||||||
if networkMap.GetPeerConfig() != nil {
|
if networkMap.GetPeerConfig() != nil {
|
||||||
err := e.updateConfig(networkMap.GetPeerConfig())
|
err := e.updateConfig(networkMap.GetPeerConfig())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
serial := networkMap.GetSerial()
|
serial := networkMap.GetSerial()
|
||||||
if e.networkSerial > serial {
|
if e.networkSerial > serial {
|
||||||
log.Debugf("received outdated NetworkMap with serial %d, ignoring", serial)
|
log.Debugf("received outdated NetworkMap with serial %d, ignoring", serial)
|
||||||
return false, nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Wholesale sections (firewall/ACL, DNS, routes, forward rules) are applied
|
|
||||||
// up-front and only once per target: they are cheap, local, idempotent and must
|
|
||||||
// be in place before peers come up (fail-closed). On the bounded re-runs that only
|
|
||||||
// drain the remaining peer batches they are skipped — the applied forward rules are
|
|
||||||
// reused from e.forwardingRules for the lazy-exclude finalize.
|
|
||||||
if firstPass {
|
|
||||||
e.applyWholesale(networkMap, serial)
|
|
||||||
}
|
|
||||||
|
|
||||||
log.Debugf("got peers update from Management Service, total peers to connect to = %d", len(networkMap.GetRemotePeers()))
|
|
||||||
|
|
||||||
doneOffline := e.phase("offline_peers")
|
|
||||||
e.updateOfflinePeers(networkMap.GetOfflinePeers())
|
|
||||||
doneOffline()
|
|
||||||
|
|
||||||
// Filter out own peer from the remote peers list
|
|
||||||
localPubKey := e.config.WgPrivateKey.PublicKey().String()
|
|
||||||
remotePeers := make([]*mgmProto.RemotePeerConfig, 0, len(networkMap.GetRemotePeers()))
|
|
||||||
for _, p := range networkMap.GetRemotePeers() {
|
|
||||||
if p.GetWgPubKey() != localPubKey {
|
|
||||||
remotePeers = append(remotePeers, p)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// No special case for cleanup: when management signals RemotePeersIsEmpty (e.g. our
|
|
||||||
// peer was deleted), remotePeers is already empty, so the bounded diff below removes
|
|
||||||
// every peer in batches — same path as a normal update, no unbounded removeAllPeers
|
|
||||||
// held under syncMsgMux in one shot.
|
|
||||||
doneRemoved := e.phase("removed_peers")
|
|
||||||
removeMore, err := e.removePeers(remotePeers, maxBatch)
|
|
||||||
doneRemoved()
|
|
||||||
if err != nil {
|
|
||||||
return false, err
|
|
||||||
}
|
|
||||||
|
|
||||||
doneModified := e.phase("modified_peers")
|
|
||||||
modifyMore, err := e.modifyPeers(remotePeers, maxBatch)
|
|
||||||
doneModified()
|
|
||||||
if err != nil {
|
|
||||||
return false, err
|
|
||||||
}
|
|
||||||
|
|
||||||
doneAdded := e.phase("added_peers")
|
|
||||||
addMore, err := e.addNewPeers(remotePeers, maxBatch)
|
|
||||||
doneAdded()
|
|
||||||
if err != nil {
|
|
||||||
return false, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// needMore signals the caller to re-run when a peer phase hit its per-pass cap.
|
|
||||||
needMore := removeMore || modifyMore || addMore
|
|
||||||
|
|
||||||
e.statusRecorder.FinishPeerListModifications()
|
|
||||||
|
|
||||||
e.updatePeerSSHHostKeys(remotePeers)
|
|
||||||
|
|
||||||
if err := e.updateSSHClientConfig(remotePeers); err != nil {
|
|
||||||
log.Warnf("failed to update SSH client config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
e.updateSSHServerAuth(networkMap.GetSshAuth())
|
|
||||||
|
|
||||||
// Set the exclude list only once peers have fully converged (this pass added
|
|
||||||
// the last batch). It needs all target peers present in the store, and
|
|
||||||
// ExcludePeer has replace-semantics — a partial set mid-convergence would be wrong.
|
|
||||||
if !needMore {
|
|
||||||
doneLazy := e.phase("lazy_exclude")
|
|
||||||
excludedLazyPeers := e.toExcludedLazyPeers(e.forwardingRules, remotePeers)
|
|
||||||
e.connMgr.SetExcludeList(e.ctx, excludedLazyPeers)
|
|
||||||
doneLazy()
|
|
||||||
}
|
|
||||||
|
|
||||||
e.networkSerial = serial
|
|
||||||
|
|
||||||
return needMore, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// applyWholesale applies the cheap, local, idempotent map sections — lazy feature
|
|
||||||
// flag, firewall/legacy management, DNS, routes, ACL filtering, DNS forwarder and
|
|
||||||
// ingress forward rules — that must be in place before peers come up. It runs once
|
|
||||||
// per target (first pass only); the resulting forward rules are stashed in
|
|
||||||
// e.forwardingRules for the lazy-exclude finalize on the peer-converged pass.
|
|
||||||
func (e *Engine) applyWholesale(networkMap *mgmProto.NetworkMap, serial uint64) {
|
|
||||||
if err := e.connMgr.UpdatedRemoteFeatureFlag(e.ctx, networkMap.GetPeerConfig().GetLazyConnectionEnabled()); err != nil {
|
if err := e.connMgr.UpdatedRemoteFeatureFlag(e.ctx, networkMap.GetPeerConfig().GetLazyConnectionEnabled()); err != nil {
|
||||||
log.Errorf("failed to update lazy connection feature flag: %v", err)
|
log.Errorf("failed to update lazy connection feature flag: %v", err)
|
||||||
}
|
}
|
||||||
@@ -1574,7 +1556,84 @@ func (e *Engine) applyWholesale(networkMap *mgmProto.NetworkMap, serial uint64)
|
|||||||
log.Errorf("failed to update forward rules, err: %v", err)
|
log.Errorf("failed to update forward rules, err: %v", err)
|
||||||
}
|
}
|
||||||
done()
|
done()
|
||||||
e.forwardingRules = forwardingRules
|
|
||||||
|
log.Debugf("got peers update from Management Service, total peers to connect to = %d", len(networkMap.GetRemotePeers()))
|
||||||
|
|
||||||
|
done = e.phase("offline_peers")
|
||||||
|
e.updateOfflinePeers(networkMap.GetOfflinePeers())
|
||||||
|
done()
|
||||||
|
|
||||||
|
remotePeers, err := e.reconcilePeers(networkMap)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// must set the exclude list after the peers are added. Without it the manager can not figure out the peers parameters from the store
|
||||||
|
done = e.phase("lazy_exclude")
|
||||||
|
excludedLazyPeers := e.toExcludedLazyPeers(forwardingRules, remotePeers)
|
||||||
|
e.connMgr.SetExcludeList(e.ctx, excludedLazyPeers)
|
||||||
|
done()
|
||||||
|
|
||||||
|
e.networkSerial = serial
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// reconcilePeers applies the remote peer list from the network map (removing,
|
||||||
|
// modifying and adding peers, then updating SSH config) and returns the remote
|
||||||
|
// peers with our own peer filtered out, for use by later sync steps.
|
||||||
|
func (e *Engine) reconcilePeers(networkMap *mgmProto.NetworkMap) ([]*mgmProto.RemotePeerConfig, error) {
|
||||||
|
// Filter out own peer from the remote peers list
|
||||||
|
localPubKey := e.config.WgPrivateKey.PublicKey().String()
|
||||||
|
remotePeers := make([]*mgmProto.RemotePeerConfig, 0, len(networkMap.GetRemotePeers()))
|
||||||
|
for _, p := range networkMap.GetRemotePeers() {
|
||||||
|
if p.GetWgPubKey() != localPubKey {
|
||||||
|
remotePeers = append(remotePeers, p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// cleanup request, most likely our peer has been deleted
|
||||||
|
if networkMap.GetRemotePeersIsEmpty() {
|
||||||
|
err := e.removeAllPeers()
|
||||||
|
e.statusRecorder.FinishPeerListModifications()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return remotePeers, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
done := e.phase("removed_peers")
|
||||||
|
err := e.removePeers(remotePeers)
|
||||||
|
done()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
done = e.phase("modified_peers")
|
||||||
|
err = e.modifyPeers(remotePeers)
|
||||||
|
done()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
done = e.phase("added_peers")
|
||||||
|
err = e.addNewPeers(remotePeers)
|
||||||
|
done()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
e.statusRecorder.FinishPeerListModifications()
|
||||||
|
|
||||||
|
e.updatePeerSSHHostKeys(remotePeers)
|
||||||
|
|
||||||
|
if err := e.updateSSHClientConfig(remotePeers); err != nil {
|
||||||
|
log.Warnf("failed to update SSH client config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
e.updateSSHServerAuth(networkMap.GetSshAuth())
|
||||||
|
|
||||||
|
return remotePeers, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func toDNSFeatureFlag(networkMap *mgmProto.NetworkMap) bool {
|
func toDNSFeatureFlag(networkMap *mgmProto.NetworkMap) bool {
|
||||||
@@ -1754,23 +1813,14 @@ func addrToString(addr netip.Addr) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// addNewPeers adds peers that were not know before but arrived from the Management service with the update
|
// addNewPeers adds peers that were not know before but arrived from the Management service with the update
|
||||||
// addNewPeers adds up to maxBatch not-yet-present peers per call. It returns true
|
func (e *Engine) addNewPeers(peersUpdate []*mgmProto.RemotePeerConfig) error {
|
||||||
// when more new peers remained than the cap, so the caller re-runs.
|
|
||||||
func (e *Engine) addNewPeers(peersUpdate []*mgmProto.RemotePeerConfig, maxBatch int) (bool, error) {
|
|
||||||
added := 0
|
|
||||||
for _, p := range peersUpdate {
|
for _, p := range peersUpdate {
|
||||||
if _, ok := e.peerStore.PeerConn(p.GetWgPubKey()); ok {
|
err := e.addNewPeer(p)
|
||||||
continue // already present (cheap skip), does not count toward the cap
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
if added >= maxBatch {
|
|
||||||
return true, nil // at least one more new peer remains
|
|
||||||
}
|
|
||||||
if err := e.addNewPeer(p); err != nil {
|
|
||||||
return false, err
|
|
||||||
}
|
|
||||||
added++
|
|
||||||
}
|
}
|
||||||
return false, nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// addNewPeer add peer if connection doesn't exist
|
// addNewPeer add peer if connection doesn't exist
|
||||||
@@ -2045,41 +2095,6 @@ func (e *Engine) close() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (e *Engine) readInitialSettings() ([]*route.Route, *nbdns.Config, bool, error) {
|
|
||||||
if runtime.GOOS != "android" {
|
|
||||||
// nolint:nilnil
|
|
||||||
return nil, nil, false, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
info := system.GetInfo(e.ctx)
|
|
||||||
info.SetFlags(
|
|
||||||
e.config.RosenpassEnabled,
|
|
||||||
e.config.RosenpassPermissive,
|
|
||||||
&e.config.ServerSSHAllowed,
|
|
||||||
e.config.DisableClientRoutes,
|
|
||||||
e.config.DisableServerRoutes,
|
|
||||||
e.config.DisableDNS,
|
|
||||||
e.config.DisableFirewall,
|
|
||||||
e.config.BlockLANAccess,
|
|
||||||
e.config.BlockInbound,
|
|
||||||
e.config.DisableIPv6,
|
|
||||||
e.config.EnableSSHRoot,
|
|
||||||
e.config.EnableSSHSFTP,
|
|
||||||
e.config.EnableSSHLocalPortForwarding,
|
|
||||||
e.config.EnableSSHRemotePortForwarding,
|
|
||||||
e.config.DisableSSHAuth,
|
|
||||||
)
|
|
||||||
|
|
||||||
netMap, err := e.mgmClient.GetNetworkMap(info)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, false, err
|
|
||||||
}
|
|
||||||
routes := toRoutes(netMap.GetRoutes())
|
|
||||||
dnsCfg := toDNSConfig(netMap.GetDNSConfig(), e.wgInterface.Address())
|
|
||||||
dnsFeatureFlag := toDNSFeatureFlag(netMap)
|
|
||||||
return routes, &dnsCfg, dnsFeatureFlag, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *Engine) newWgIface() (*iface.WGIface, error) {
|
func (e *Engine) newWgIface() (*iface.WGIface, error) {
|
||||||
transportNet, err := e.newStdNet()
|
transportNet, err := e.newStdNet()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -2114,7 +2129,7 @@ func (e *Engine) newWgIface() (*iface.WGIface, error) {
|
|||||||
func (e *Engine) wgInterfaceCreate() (err error) {
|
func (e *Engine) wgInterfaceCreate() (err error) {
|
||||||
switch runtime.GOOS {
|
switch runtime.GOOS {
|
||||||
case "android":
|
case "android":
|
||||||
err = e.wgInterface.CreateOnAndroid(e.routeManager.InitialRouteRange(), e.dnsServer.DnsIP().String(), e.dnsServer.SearchDomains())
|
err = e.wgInterface.CreateOnAndroid(e.routeManager.CurrentRouteRange(), e.dnsServer.DnsIP().String(), e.dnsServer.SearchDomains())
|
||||||
case "ios":
|
case "ios":
|
||||||
e.mobileDep.NetworkChangeListener.SetInterfaceIP(e.config.WgAddr.String())
|
e.mobileDep.NetworkChangeListener.SetInterfaceIP(e.config.WgAddr.String())
|
||||||
if e.config.WgAddr.HasIPv6() {
|
if e.config.WgAddr.HasIPv6() {
|
||||||
@@ -2127,7 +2142,7 @@ func (e *Engine) wgInterfaceCreate() (err error) {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (e *Engine) newDnsServer(dnsConfig *nbdns.Config) (dns.Server, error) {
|
func (e *Engine) newDnsServer() (dns.Server, error) {
|
||||||
// due to tests where we are using a mocked version of the DNS server
|
// due to tests where we are using a mocked version of the DNS server
|
||||||
if e.dnsServer != nil {
|
if e.dnsServer != nil {
|
||||||
return e.dnsServer, nil
|
return e.dnsServer, nil
|
||||||
@@ -2139,7 +2154,7 @@ func (e *Engine) newDnsServer(dnsConfig *nbdns.Config) (dns.Server, error) {
|
|||||||
e.ctx,
|
e.ctx,
|
||||||
e.wgInterface,
|
e.wgInterface,
|
||||||
e.mobileDep.HostDNSAddresses,
|
e.mobileDep.HostDNSAddresses,
|
||||||
*dnsConfig,
|
nbdns.Config{},
|
||||||
e.mobileDep.NetworkChangeListener,
|
e.mobileDep.NetworkChangeListener,
|
||||||
e.statusRecorder,
|
e.statusRecorder,
|
||||||
e.config.DisableDNS,
|
e.config.DisableDNS,
|
||||||
@@ -2255,7 +2270,20 @@ func (e *Engine) getRosenpassAddr() string {
|
|||||||
|
|
||||||
// RunHealthProbes executes health checks for Signal, Management, Relay, and WireGuard services
|
// RunHealthProbes executes health checks for Signal, Management, Relay, and WireGuard services
|
||||||
// and updates the status recorder with the latest states.
|
// and updates the status recorder with the latest states.
|
||||||
func (e *Engine) RunHealthProbes(waitForResult bool) bool {
|
//
|
||||||
|
// ctx scopes the (potentially slow) STUN/TURN probing: a caller that gives up —
|
||||||
|
// e.g. a Status RPC whose client disconnected — cancels its ctx and the probe
|
||||||
|
// returns instead of running to its per-component timeout. The engine's own
|
||||||
|
// lifetime ctx still applies independently, so an engine shutdown aborts the
|
||||||
|
// probe even if the caller's ctx is context.Background().
|
||||||
|
func (e *Engine) RunHealthProbes(ctx context.Context, waitForResult bool) bool {
|
||||||
|
// Tie the caller's ctx to the engine lifetime: either cancelling aborts
|
||||||
|
// the probe below.
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
defer cancel()
|
||||||
|
stop := context.AfterFunc(e.ctx, cancel)
|
||||||
|
defer stop()
|
||||||
|
|
||||||
e.syncMsgMux.Lock()
|
e.syncMsgMux.Lock()
|
||||||
|
|
||||||
signalHealthy := e.signal.IsHealthy()
|
signalHealthy := e.signal.IsHealthy()
|
||||||
@@ -2278,9 +2306,9 @@ func (e *Engine) RunHealthProbes(waitForResult bool) bool {
|
|||||||
if runtime.GOOS != "js" {
|
if runtime.GOOS != "js" {
|
||||||
var results []relay.ProbeResult
|
var results []relay.ProbeResult
|
||||||
if waitForResult {
|
if waitForResult {
|
||||||
results = e.probeStunTurn.ProbeAllWaitResult(e.ctx, stuns, turns)
|
results = e.probeStunTurn.ProbeAllWaitResult(ctx, stuns, turns)
|
||||||
} else {
|
} else {
|
||||||
results = e.probeStunTurn.ProbeAll(e.ctx, stuns, turns)
|
results = e.probeStunTurn.ProbeAll(ctx, stuns, turns)
|
||||||
}
|
}
|
||||||
e.statusRecorder.UpdateRelayStates(results)
|
e.statusRecorder.UpdateRelayStates(results)
|
||||||
|
|
||||||
@@ -2623,13 +2651,14 @@ func (e *Engine) updateForwardRules(rules []*mgmProto.ForwardingRule) ([]firewal
|
|||||||
|
|
||||||
func (e *Engine) toExcludedLazyPeers(rules []firewallManager.ForwardRule, peers []*mgmProto.RemotePeerConfig) map[string]bool {
|
func (e *Engine) toExcludedLazyPeers(rules []firewallManager.ForwardRule, peers []*mgmProto.RemotePeerConfig) map[string]bool {
|
||||||
excludedPeers := make(map[string]bool)
|
excludedPeers := make(map[string]bool)
|
||||||
|
|
||||||
|
// Ingress forward targets: inbound forwarded traffic is initiated remotely and
|
||||||
|
// cannot wake a lazy connection, so the peer routing the target must stay
|
||||||
|
// permanently connected. AllowedIPs are already parsed on the peer conn, so
|
||||||
|
// reuse those typed prefixes instead of re-parsing the network map strings.
|
||||||
for _, r := range rules {
|
for _, r := range rules {
|
||||||
ip := r.TranslatedAddress
|
|
||||||
for _, p := range peers {
|
for _, p := range peers {
|
||||||
for _, allowedIP := range p.GetAllowedIps() {
|
if e.peerRoutesAddr(p, r.TranslatedAddress) {
|
||||||
if allowedIP != ip.String() {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
log.Infof("exclude forwarder peer from lazy connection: %s", p.GetWgPubKey())
|
log.Infof("exclude forwarder peer from lazy connection: %s", p.GetWgPubKey())
|
||||||
excludedPeers[p.GetWgPubKey()] = true
|
excludedPeers[p.GetWgPubKey()] = true
|
||||||
}
|
}
|
||||||
@@ -2639,6 +2668,27 @@ func (e *Engine) toExcludedLazyPeers(rules []firewallManager.ForwardRule, peers
|
|||||||
return excludedPeers
|
return excludedPeers
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// peerRoutesAddr reports whether the peer is a router for addr, matched against
|
||||||
|
// the peer's already-parsed AllowedIPs from the store (the same typed value the
|
||||||
|
// lazy manager consumes) rather than re-parsing the network map strings.
|
||||||
|
func (e *Engine) peerRoutesAddr(p *mgmProto.RemotePeerConfig, addr netip.Addr) bool {
|
||||||
|
prefixes, ok := e.peerStore.AllowedIPs(p.GetWgPubKey())
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return prefixesContain(prefixes, addr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// prefixesContain reports whether addr falls within any of the prefixes.
|
||||||
|
func prefixesContain(prefixes []netip.Prefix, addr netip.Addr) bool {
|
||||||
|
for _, prefix := range prefixes {
|
||||||
|
if prefix.Contains(addr) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
// isChecksEqual checks if two slices of checks are equal.
|
// isChecksEqual checks if two slices of checks are equal.
|
||||||
func isChecksEqual(checks1, checks2 []*mgmProto.Checks) bool {
|
func isChecksEqual(checks1, checks2 []*mgmProto.Checks) bool {
|
||||||
normalize := func(checks []*mgmProto.Checks) []string {
|
normalize := func(checks []*mgmProto.Checks) []string {
|
||||||
|
|||||||
108
client/internal/engine_authsession.go
Normal file
108
client/internal/engine_authsession.go
Normal file
@@ -0,0 +1,108 @@
|
|||||||
|
package internal
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"google.golang.org/protobuf/types/known/timestamppb"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
|
||||||
|
cProto "github.com/netbirdio/netbird/client/proto"
|
||||||
|
"github.com/netbirdio/netbird/client/system"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ApplySessionDeadline propagates the absolute SSO session deadline carried on
|
||||||
|
// LoginResponse / SyncResponse to both the watcher (for the edge-triggered
|
||||||
|
// warning) and the status recorder (for the SubscribeStatus / Status RPC
|
||||||
|
// snapshot the UI consumes).
|
||||||
|
//
|
||||||
|
// The wire field is 3-state:
|
||||||
|
// - nil → snapshot carries no info; keep the
|
||||||
|
// previously-anchored deadline (no-op)
|
||||||
|
// - explicit zero (s=0, n=0) → peer is not SSO-registered or expiry is
|
||||||
|
// disabled; clear both sinks
|
||||||
|
// - valid timestamp → new deadline; arm watcher, expose on
|
||||||
|
// status recorder
|
||||||
|
//
|
||||||
|
// Deadline sanity-checks live in sessionwatch.Watcher.Update. Any rejected
|
||||||
|
// value is treated as a clear on both sinks: the alternative — leaving the
|
||||||
|
// previously-known deadline in place — risks the UI confidently displaying
|
||||||
|
// a stale "expires in X" while the server has actually invalidated it.
|
||||||
|
func (e *Engine) ApplySessionDeadline(ts *timestamppb.Timestamp) {
|
||||||
|
if ts == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var deadline time.Time
|
||||||
|
// Explicit zero (seconds=0 AND nanos=0) is the sentinel for "disabled".
|
||||||
|
// Everything else flows through Watcher.Update, whose sanity-checks
|
||||||
|
// reject out-of-range / pre-epoch / far-future / too-stale values and
|
||||||
|
// clear on rejection.
|
||||||
|
if ts.GetSeconds() != 0 || ts.GetNanos() != 0 {
|
||||||
|
deadline = ts.AsTime().UTC()
|
||||||
|
}
|
||||||
|
if e.sessionWatcher == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// Watcher.Update owns the propagation to the status recorder (the
|
||||||
|
// SubscribeStatus / Status snapshot the UI reads): a set writes the
|
||||||
|
// deadline, a clear or a sanity-check rejection writes the zero value.
|
||||||
|
// Keeping a single writer is what stops the recorder from drifting out
|
||||||
|
// of sync with the warning timers.
|
||||||
|
if err := e.sessionWatcher.Update(deadline); err != nil {
|
||||||
|
log.Errorf("auth session deadline rejected: %v, clearing", err)
|
||||||
|
e.statusRecorder.PublishEvent(
|
||||||
|
cProto.SystemEvent_ERROR,
|
||||||
|
cProto.SystemEvent_AUTHENTICATION,
|
||||||
|
"session deadline rejected",
|
||||||
|
"",
|
||||||
|
map[string]string{sessionwatch.MetaSessionDeadlineRejected: err.Error()},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// DismissSessionWarning records the user's "Dismiss" click on the
|
||||||
|
// T-WarningLead interactive notification and suppresses the upcoming
|
||||||
|
// T-FinalWarningLead fallback for the current deadline. No-op when the
|
||||||
|
// watcher is not running or holds no deadline.
|
||||||
|
func (e *Engine) DismissSessionWarning() {
|
||||||
|
if e.sessionWatcher == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
e.sessionWatcher.Dismiss()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExtendAuthSession asks the management server to refresh the SSO session
|
||||||
|
// expiry deadline using the supplied JWT, then mirrors the new deadline into
|
||||||
|
// the daemon's state. The tunnel is untouched; no resync, no reconnect.
|
||||||
|
//
|
||||||
|
// Returns the new absolute UTC deadline (or zero time when the server
|
||||||
|
// reports the peer is not eligible for extension).
|
||||||
|
func (e *Engine) ExtendAuthSession(ctx context.Context, jwtToken string) (time.Time, error) {
|
||||||
|
if jwtToken == "" {
|
||||||
|
return time.Time{}, errors.New("jwt token is required")
|
||||||
|
}
|
||||||
|
if e.mgmClient == nil {
|
||||||
|
return time.Time{}, errors.New("management client is not initialised")
|
||||||
|
}
|
||||||
|
|
||||||
|
info, err := system.GetInfoWithChecks(ctx, e.checks)
|
||||||
|
if err != nil {
|
||||||
|
log.Warnf("failed to collect system info for session extend: %v", err)
|
||||||
|
info = system.GetInfo(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := e.mgmClient.ExtendAuthSession(info, jwtToken)
|
||||||
|
if err != nil {
|
||||||
|
return time.Time{}, fmt.Errorf("extend auth session on management: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
e.ApplySessionDeadline(resp.GetSessionExpiresAt())
|
||||||
|
|
||||||
|
if resp.GetSessionExpiresAt().IsValid() {
|
||||||
|
return resp.GetSessionExpiresAt().AsTime().UTC(), nil
|
||||||
|
}
|
||||||
|
return time.Time{}, nil
|
||||||
|
}
|
||||||
87
client/internal/engine_lazy_exclude_test.go
Normal file
87
client/internal/engine_lazy_exclude_test.go
Normal file
@@ -0,0 +1,87 @@
|
|||||||
|
package internal
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
firewallManager "github.com/netbirdio/netbird/client/firewall/manager"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peerstore"
|
||||||
|
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPrefixesContain(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
prefixes []string
|
||||||
|
addr string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{name: "own overlay /32 matches", prefixes: []string{"100.110.8.145/32"}, addr: "100.110.8.145", want: true},
|
||||||
|
{name: "addr inside routed subnet", prefixes: []string{"10.121.0.0/16"}, addr: "10.121.208.4", want: true},
|
||||||
|
{name: "addr outside subnet", prefixes: []string{"10.121.0.0/16"}, addr: "10.122.0.1", want: false},
|
||||||
|
{name: "different /32", prefixes: []string{"100.110.8.145/32"}, addr: "100.110.8.146", want: false},
|
||||||
|
{name: "ipv6 /128 matches", prefixes: []string{"fd00::1/128"}, addr: "fd00::1", want: true},
|
||||||
|
{name: "no prefixes", prefixes: nil, addr: "10.121.208.4", want: false},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
prefixes := make([]netip.Prefix, 0, len(tt.prefixes))
|
||||||
|
for _, p := range tt.prefixes {
|
||||||
|
prefixes = append(prefixes, netip.MustParsePrefix(p))
|
||||||
|
}
|
||||||
|
require.Equal(t, tt.want, prefixesContain(prefixes, netip.MustParseAddr(tt.addr)))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestToExcludedLazyPeers_ForwardTarget guards a regression: the forward-target
|
||||||
|
// peer (the peer routing a ForwardRule.TranslatedAddress) must be excluded from
|
||||||
|
// lazy connections, matched via the peer's already-parsed AllowedIPs.
|
||||||
|
func TestToExcludedLazyPeers_ForwardTarget(t *testing.T) {
|
||||||
|
const targetPeerKey = "cccccccccccccccccccccccccccccccccccccccccc0="
|
||||||
|
const otherPeerKey = "dddddddddddddddddddddddddddddddddddddddddd0="
|
||||||
|
|
||||||
|
store := peerstore.NewConnStore()
|
||||||
|
store.AddPeerConn(targetPeerKey, newTestConn(t, targetPeerKey, "100.110.8.145/32"))
|
||||||
|
store.AddPeerConn(otherPeerKey, newTestConn(t, otherPeerKey, "100.110.9.10/32"))
|
||||||
|
|
||||||
|
e := &Engine{peerStore: store}
|
||||||
|
|
||||||
|
peers := []*mgmProto.RemotePeerConfig{
|
||||||
|
{WgPubKey: targetPeerKey, AllowedIps: []string{"100.110.8.145/32"}},
|
||||||
|
{WgPubKey: otherPeerKey, AllowedIps: []string{"100.110.9.10/32"}},
|
||||||
|
}
|
||||||
|
rules := []firewallManager.ForwardRule{
|
||||||
|
{TranslatedAddress: netip.MustParseAddr("100.110.8.145")},
|
||||||
|
}
|
||||||
|
|
||||||
|
excluded := e.toExcludedLazyPeers(rules, peers)
|
||||||
|
|
||||||
|
require.True(t, excluded[targetPeerKey], "forward-target peer must be excluded from lazy connections")
|
||||||
|
require.False(t, excluded[otherPeerKey], "non-target peer must not be excluded")
|
||||||
|
require.Len(t, excluded, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToExcludedLazyPeers_NoRules(t *testing.T) {
|
||||||
|
e := &Engine{peerStore: peerstore.NewConnStore()}
|
||||||
|
|
||||||
|
peers := []*mgmProto.RemotePeerConfig{
|
||||||
|
{WgPubKey: "peer-a", AllowedIps: []string{"100.110.8.145/32"}},
|
||||||
|
}
|
||||||
|
|
||||||
|
require.Empty(t, e.toExcludedLazyPeers(nil, peers))
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestConn(t *testing.T, key, allowedIP string) *peer.Conn {
|
||||||
|
t.Helper()
|
||||||
|
conn, err := peer.NewConn(peer.ConnConfig{
|
||||||
|
Key: key,
|
||||||
|
WgConfig: peer.WgConfig{AllowedIps: []netip.Prefix{netip.MustParsePrefix(allowedIP)}},
|
||||||
|
}, peer.ServiceDependencies{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
return conn
|
||||||
|
}
|
||||||
@@ -124,7 +124,7 @@ func TestEngine_SSH(t *testing.T) {
|
|||||||
RemotePeersIsEmpty: false,
|
RemotePeersIsEmpty: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = engine.updateNetworkMap(networkMap, maxPeersPerSyncPass, true)
|
err = engine.updateNetworkMap(networkMap)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
assert.Nil(t, engine.sshServer)
|
assert.Nil(t, engine.sshServer)
|
||||||
@@ -146,7 +146,7 @@ func TestEngine_SSH(t *testing.T) {
|
|||||||
RemotePeersIsEmpty: false,
|
RemotePeersIsEmpty: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = engine.updateNetworkMap(networkMap, maxPeersPerSyncPass, true)
|
err = engine.updateNetworkMap(networkMap)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
time.Sleep(250 * time.Millisecond)
|
time.Sleep(250 * time.Millisecond)
|
||||||
@@ -159,7 +159,7 @@ func TestEngine_SSH(t *testing.T) {
|
|||||||
RemotePeersIsEmpty: false,
|
RemotePeersIsEmpty: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = engine.updateNetworkMap(networkMap, maxPeersPerSyncPass, true)
|
err = engine.updateNetworkMap(networkMap)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
// time.Sleep(250 * time.Millisecond)
|
// time.Sleep(250 * time.Millisecond)
|
||||||
@@ -174,7 +174,7 @@ func TestEngine_SSH(t *testing.T) {
|
|||||||
RemotePeersIsEmpty: false,
|
RemotePeersIsEmpty: false,
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = engine.updateNetworkMap(networkMap, maxPeersPerSyncPass, true)
|
err = engine.updateNetworkMap(networkMap)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
assert.Nil(t, engine.sshServer)
|
assert.Nil(t, engine.sshServer)
|
||||||
|
|||||||
88
client/internal/engine_session_deadline_test.go
Normal file
88
client/internal/engine_session_deadline_test.go
Normal file
@@ -0,0 +1,88 @@
|
|||||||
|
package internal
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"google.golang.org/protobuf/types/known/timestamppb"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestApplySessionDeadline_ThreeState pins down the 3-state semantics of the
|
||||||
|
// wire field carried on LoginResponse / SyncResponse:
|
||||||
|
//
|
||||||
|
// - nil pointer → no info; previously-anchored deadline survives
|
||||||
|
// - explicit zero value → "expiry disabled" sentinel; both sinks cleared
|
||||||
|
// - valid future timestamp → new deadline propagated to both sinks
|
||||||
|
func TestApplySessionDeadline_ThreeState(t *testing.T) {
|
||||||
|
newEngine := func() *Engine {
|
||||||
|
recorder := peer.NewRecorder("")
|
||||||
|
return &Engine{
|
||||||
|
statusRecorder: recorder,
|
||||||
|
sessionWatcher: sessionwatch.New(recorder),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("valid timestamp sets deadline on both sinks", func(t *testing.T) {
|
||||||
|
e := newEngine()
|
||||||
|
deadline := time.Now().Add(time.Hour).UTC().Truncate(time.Second)
|
||||||
|
|
||||||
|
e.ApplySessionDeadline(timestamppb.New(deadline))
|
||||||
|
|
||||||
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(deadline),
|
||||||
|
"status recorder should hold the new deadline")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nil is a no-op and preserves previous deadline", func(t *testing.T) {
|
||||||
|
e := newEngine()
|
||||||
|
seeded := time.Now().Add(time.Hour).UTC().Truncate(time.Second)
|
||||||
|
e.ApplySessionDeadline(timestamppb.New(seeded))
|
||||||
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(seeded))
|
||||||
|
|
||||||
|
e.ApplySessionDeadline(nil)
|
||||||
|
|
||||||
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(seeded),
|
||||||
|
"nil snapshot must not disturb the existing deadline")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("explicit zero clears a previously-anchored deadline", func(t *testing.T) {
|
||||||
|
e := newEngine()
|
||||||
|
seeded := time.Now().Add(time.Hour).UTC().Truncate(time.Second)
|
||||||
|
e.ApplySessionDeadline(timestamppb.New(seeded))
|
||||||
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(seeded))
|
||||||
|
|
||||||
|
// Explicit zero Timestamp{} (seconds=0, nanos=0) is the
|
||||||
|
// "expiry disabled / not SSO" sentinel.
|
||||||
|
e.ApplySessionDeadline(×tamppb.Timestamp{})
|
||||||
|
|
||||||
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().IsZero(),
|
||||||
|
"explicit zero sentinel must clear the deadline")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("invalid timestamp clears the deadline", func(t *testing.T) {
|
||||||
|
e := newEngine()
|
||||||
|
seeded := time.Now().Add(time.Hour).UTC().Truncate(time.Second)
|
||||||
|
e.ApplySessionDeadline(timestamppb.New(seeded))
|
||||||
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(seeded))
|
||||||
|
|
||||||
|
// Out-of-range nanos → IsValid()==false; same-meaning as the
|
||||||
|
// disabled sentinel for downstream sinks.
|
||||||
|
e.ApplySessionDeadline(×tamppb.Timestamp{Seconds: 1, Nanos: -1})
|
||||||
|
|
||||||
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().IsZero(),
|
||||||
|
"invalid timestamp must clear the deadline")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("recently expired timestamp stays visible as expired", func(t *testing.T) {
|
||||||
|
e := newEngine()
|
||||||
|
expired := time.Now().Add(-5 * time.Minute).UTC().Truncate(time.Second)
|
||||||
|
|
||||||
|
e.ApplySessionDeadline(timestamppb.New(expired))
|
||||||
|
|
||||||
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(expired),
|
||||||
|
"recently-expired deadline must stay on the recorder so consumers render it as expired")
|
||||||
|
})
|
||||||
|
}
|
||||||
16
client/internal/engine_sessionwatch.go
Normal file
16
client/internal/engine_sessionwatch.go
Normal file
@@ -0,0 +1,16 @@
|
|||||||
|
//go:build !js
|
||||||
|
|
||||||
|
package internal
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/netbirdio/netbird/client/internal/auth/sessionwatch"
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
)
|
||||||
|
|
||||||
|
// newSessionWatcher returns the real SSO session expiry watcher for every
|
||||||
|
// non-wasm build. The js/wasm build gets a no-op stub from
|
||||||
|
// engine_sessionwatch_js.go so the sessionwatch package (and its timer
|
||||||
|
// machinery) never links into the wasm binary.
|
||||||
|
func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
|
||||||
|
return sessionwatch.New(recorder)
|
||||||
|
}
|
||||||
44
client/internal/engine_sessionwatch_js.go
Normal file
44
client/internal/engine_sessionwatch_js.go
Normal file
@@ -0,0 +1,44 @@
|
|||||||
|
//go:build js
|
||||||
|
|
||||||
|
package internal
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/netbirdio/netbird/client/internal/peer"
|
||||||
|
)
|
||||||
|
|
||||||
|
// noopSessionWatcher is the js/wasm stand-in for sessionwatch.Watcher. The
|
||||||
|
// wasm client never runs the engine's session-warning flow (the interactive
|
||||||
|
// T-WarningLead notification and the T-FinalWarningLead fallback dialog live
|
||||||
|
// in the desktop UI), so linking the full sessionwatch package (timers, event
|
||||||
|
// composition) would only bloat the binary.
|
||||||
|
//
|
||||||
|
// It still mirrors the deadline into the status recorder so the SubscribeStatus
|
||||||
|
// / Status snapshot the UI consumes stays correct — only the timer-driven
|
||||||
|
// warnings are dropped.
|
||||||
|
type noopSessionWatcher struct {
|
||||||
|
recorder *peer.Status
|
||||||
|
}
|
||||||
|
|
||||||
|
func newSessionWatcher(recorder *peer.Status) sessionDeadlineWatcher {
|
||||||
|
return noopSessionWatcher{recorder: recorder}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update mirrors the real watcher's recorder propagation without the timers or
|
||||||
|
// sanity-check sentinels: a valid deadline is exposed on the status snapshot,
|
||||||
|
// the zero time clears it.
|
||||||
|
func (w noopSessionWatcher) Update(deadline time.Time) error {
|
||||||
|
if w.recorder != nil {
|
||||||
|
w.recorder.SetSessionExpiresAt(deadline)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (noopSessionWatcher) Dismiss() {
|
||||||
|
// No-op: only suppresses the timer-driven final-warning, which this stub never arms.
|
||||||
|
}
|
||||||
|
|
||||||
|
func (noopSessionWatcher) Close() {
|
||||||
|
// No-op: no timers to stop and no state to unwind; the recorder is cleared via Update(zero).
|
||||||
|
}
|
||||||
@@ -437,7 +437,7 @@ func TestEngine_UpdateNetworkMap(t *testing.T) {
|
|||||||
|
|
||||||
for _, c := range []testCase{case1, case2, case3, case4, case5, case6} {
|
for _, c := range []testCase{case1, case2, case3, case4, case5, case6} {
|
||||||
t.Run(c.name, func(t *testing.T) {
|
t.Run(c.name, func(t *testing.T) {
|
||||||
_, err = engine.updateNetworkMap(c.networkMap, maxPeersPerSyncPass, true)
|
err = engine.updateNetworkMap(c.networkMap)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
return
|
return
|
||||||
@@ -464,47 +464,6 @@ func TestEngine_UpdateNetworkMap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// chunked apply: with a per-pass cap smaller than the number of peers, a
|
|
||||||
// single updateNetworkMap applies one batch and reports more==true; the
|
|
||||||
// caller re-runs until convergence. (engine currently holds 0 peers.)
|
|
||||||
t.Run("chunked add converges over multiple passes", func(t *testing.T) {
|
|
||||||
nm := &mgmtProto.NetworkMap{
|
|
||||||
Serial: 6,
|
|
||||||
RemotePeers: []*mgmtProto.RemotePeerConfig{peer1, peer2, peer3},
|
|
||||||
}
|
|
||||||
|
|
||||||
more, err := engine.updateNetworkMap(nm, 1, true)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.True(t, more, "pass 1 should signal more")
|
|
||||||
require.Len(t, engine.peerStore.PeersPubKey(), 1)
|
|
||||||
|
|
||||||
more, err = engine.updateNetworkMap(nm, 1, false)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.True(t, more, "pass 2 should signal more")
|
|
||||||
require.Len(t, engine.peerStore.PeersPubKey(), 2)
|
|
||||||
|
|
||||||
more, err = engine.updateNetworkMap(nm, 1, false)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.False(t, more, "pass 3 should converge")
|
|
||||||
require.Len(t, engine.peerStore.PeersPubKey(), 3)
|
|
||||||
})
|
|
||||||
|
|
||||||
t.Run("chunked remove converges over multiple passes", func(t *testing.T) {
|
|
||||||
nm := &mgmtProto.NetworkMap{
|
|
||||||
Serial: 7,
|
|
||||||
RemotePeers: []*mgmtProto.RemotePeerConfig{peer1}, // remove peer2, peer3
|
|
||||||
}
|
|
||||||
|
|
||||||
more, err := engine.updateNetworkMap(nm, 1, true)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.True(t, more, "pass 1 should signal more (2 to remove, cap 1)")
|
|
||||||
|
|
||||||
more, err = engine.updateNetworkMap(nm, 1, false)
|
|
||||||
require.NoError(t, err)
|
|
||||||
require.False(t, more, "pass 2 should converge")
|
|
||||||
require.Len(t, engine.peerStore.PeersPubKey(), 1)
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) {
|
func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) {
|
||||||
@@ -675,7 +634,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
_, err = engine.updateNetworkMap(testCase.networkMap, maxPeersPerSyncPass, true)
|
err = engine.updateNetworkMap(testCase.networkMap)
|
||||||
assert.NoError(t, err, "shouldn't return error")
|
assert.NoError(t, err, "shouldn't return error")
|
||||||
assert.Equal(t, testCase.expectedSerial, input.inputSerial, "serial should match")
|
assert.Equal(t, testCase.expectedSerial, input.inputSerial, "serial should match")
|
||||||
assert.Len(t, input.clientRoutes, testCase.expectedLen, "clientRoutes len should match")
|
assert.Len(t, input.clientRoutes, testCase.expectedLen, "clientRoutes len should match")
|
||||||
@@ -879,7 +838,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
|
|
||||||
_, err = engine.updateNetworkMap(testCase.networkMap, maxPeersPerSyncPass, true)
|
err = engine.updateNetworkMap(testCase.networkMap)
|
||||||
assert.NoError(t, err, "shouldn't return error")
|
assert.NoError(t, err, "shouldn't return error")
|
||||||
assert.Equal(t, testCase.expectedSerial, input.inputSerial, "serial should match")
|
assert.Equal(t, testCase.expectedSerial, input.inputSerial, "serial should match")
|
||||||
assert.Len(t, input.inputNSGroups, testCase.expectedZonesLen, "zones len should match")
|
assert.Len(t, input.inputNSGroups, testCase.expectedZonesLen, "zones len should match")
|
||||||
|
|||||||
20
client/internal/engine_tunsettings.go
Normal file
20
client/internal/engine_tunsettings.go
Normal file
@@ -0,0 +1,20 @@
|
|||||||
|
package internal
|
||||||
|
|
||||||
|
func (e *Engine) TunSettings() ([]string, []string) {
|
||||||
|
e.syncMsgMux.Lock()
|
||||||
|
routeManager := e.routeManager
|
||||||
|
dnsServer := e.dnsServer
|
||||||
|
e.syncMsgMux.Unlock()
|
||||||
|
|
||||||
|
var routes []string
|
||||||
|
if routeManager != nil {
|
||||||
|
routes = routeManager.CurrentRouteRange()
|
||||||
|
}
|
||||||
|
|
||||||
|
var searchDomains []string
|
||||||
|
if dnsServer != nil {
|
||||||
|
searchDomains = dnsServer.SearchDomains()
|
||||||
|
}
|
||||||
|
|
||||||
|
return routes, searchDomains
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user