mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-01 13:21:29 +02:00
Compare commits
106 Commits
update-pro
...
dependabot
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c82e04a0b2 | ||
|
|
234abd7a08 | ||
|
|
c1f0006012 | ||
|
|
cff49237b6 | ||
|
|
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 |
9
.github/dependabot.yml
vendored
9
.github/dependabot.yml
vendored
@@ -3,8 +3,8 @@ updates:
|
|||||||
- package-ecosystem: "github-actions"
|
- package-ecosystem: "github-actions"
|
||||||
directory: "/"
|
directory: "/"
|
||||||
schedule:
|
schedule:
|
||||||
interval: "daily"
|
interval: "weekly"
|
||||||
open-pull-requests-limit: 15
|
open-pull-requests-limit: 3
|
||||||
groups:
|
groups:
|
||||||
actions:
|
actions:
|
||||||
patterns:
|
patterns:
|
||||||
@@ -22,9 +22,12 @@ updates:
|
|||||||
directories:
|
directories:
|
||||||
- "/"
|
- "/"
|
||||||
schedule:
|
schedule:
|
||||||
interval: "daily"
|
interval: "weekly"
|
||||||
open-pull-requests-limit: 15
|
open-pull-requests-limit: 15
|
||||||
groups:
|
groups:
|
||||||
|
golang-x-packages:
|
||||||
|
patterns:
|
||||||
|
- "golang.org/x/*"
|
||||||
aws-sdk:
|
aws-sdk:
|
||||||
patterns:
|
patterns:
|
||||||
- "github.com/aws/aws-sdk-go-v2/*"
|
- "github.com/aws/aws-sdk-go-v2/*"
|
||||||
|
|||||||
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 }}
|
||||||
|
|||||||
2
.github/workflows/frontend-ui.yml
vendored
2
.github/workflows/frontend-ui.yml
vendored
@@ -86,7 +86,7 @@ jobs:
|
|||||||
${{ runner.os }}-pnpm-
|
${{ runner.os }}-pnpm-
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: pnpm install --frozen-lockfile
|
run: pnpm install --frozen-lockfile --ignore-scripts
|
||||||
|
|
||||||
- name: Generate Wails bindings
|
- name: Generate Wails bindings
|
||||||
run: pnpm run bindings
|
run: pnpm run bindings
|
||||||
|
|||||||
4
.github/workflows/golangci-lint.yml
vendored
4
.github/workflows/golangci-lint.yml
vendored
@@ -45,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
|
||||||
@@ -79,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
|
||||||
|
|||||||
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
|
|
||||||
|
|||||||
@@ -273,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
|
||||||
|
|||||||
@@ -24,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
|
||||||
@@ -39,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
|
||||||
@@ -55,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
|
||||||
|
|||||||
@@ -29,6 +29,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
|
||||||
|
|
||||||
universal_binaries:
|
universal_binaries:
|
||||||
- id: netbird-ui-darwin
|
- id: netbird-ui-darwin
|
||||||
|
|||||||
@@ -234,12 +234,22 @@ cd client/ui
|
|||||||
task dev
|
task dev
|
||||||
```
|
```
|
||||||
|
|
||||||
Pass daemon flags after `--`:
|
Pass daemon flags after `--`, pointing the UI at the socket the daemon serves:
|
||||||
|
|
||||||
```
|
```
|
||||||
task dev -- --daemon-addr=tcp://127.0.0.1:41731
|
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/`):
|
Production build (frontend assets embedded into the binary, output in `client/ui/bin/`):
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|||||||
72
SECURITY.md
72
SECURITY.md
@@ -1,12 +1,70 @@
|
|||||||
# Security Policy
|
# Security Policy
|
||||||
|
|
||||||
NetBird's goal is to provide a secure network. If you find a vulnerability or bug, please report it by opening an issue [here](https://github.com/netbirdio/netbird/issues/new?assignees=&labels=&template=bug-issue-report.md&title=) or by contacting us by email.
|
NetBird's goal is to provide a secure network. The client runs as a privileged service on every machine it is installed on,
|
||||||
|
so we take reports about it seriously and we publish what we fix.
|
||||||
There has yet to be an official bug bounty program for the NetBird project.
|
|
||||||
|
|
||||||
## Supported Versions
|
|
||||||
- We currently support only the latest version
|
|
||||||
|
|
||||||
## Reporting a Vulnerability
|
## Reporting a Vulnerability
|
||||||
|
|
||||||
Please report security issues to `security@netbird.io`
|
**Please do not open a public issue for a security vulnerability.** Public issues are visible to everyone, including before
|
||||||
|
a fix is available.
|
||||||
|
|
||||||
|
Report security issues one of these two ways:
|
||||||
|
|
||||||
|
- **GitHub private vulnerability reporting** — [open a private report](https://github.com/netbirdio/netbird/security/advisories/new)
|
||||||
|
on this repository. This is the preferred route: it keeps the discussion, the draft advisory, and the credit in one place.
|
||||||
|
- **Email** — `security@netbird.io`.
|
||||||
|
|
||||||
|
If the finding affects NetBird Cloud or our hosted infrastructure rather than the open-source code, email us rather than
|
||||||
|
filing a repository report.
|
||||||
|
|
||||||
|
### What to include
|
||||||
|
|
||||||
|
A report is easier to act on when it contains:
|
||||||
|
|
||||||
|
- The affected component (client, management, signal, relay, dashboard) and the version or commit you tested
|
||||||
|
- The platform and configuration, where relevant — operating system, self-hosted or NetBird Cloud, container or host install
|
||||||
|
- What an attacker needs before they can exploit it: network position, an account, local access, a specific privilege level
|
||||||
|
- Steps to reproduce, and a proof of concept if you have one
|
||||||
|
- The impact you believe it has
|
||||||
|
|
||||||
|
Partial reports are still welcome. If you are unsure whether something is a security issue, send it to `security@netbird.io`
|
||||||
|
and let us make that call.
|
||||||
|
|
||||||
|
## What to expect from us
|
||||||
|
|
||||||
|
- **We acknowledge your report** and tell you whether we can reproduce it.
|
||||||
|
- **We work with you on severity and scope.** If we assess it differently than you do, we will explain why rather than
|
||||||
|
silently downgrade it.
|
||||||
|
- **We fix and release**, then publish a [GitHub Security Advisory](https://github.com/netbirdio/netbird/security/advisories)
|
||||||
|
naming the affected version range and the patched version.
|
||||||
|
- **We credit reporters who want to be credited.** Tell us the name or handle you would like used, or that you would rather
|
||||||
|
stay anonymous.
|
||||||
|
- **We keep you in the loop** until the advisory is published.
|
||||||
|
|
||||||
|
We ask that you give us a reasonable opportunity to ship a fix before disclosing the issue publicly, and that you avoid
|
||||||
|
accessing, modifying, or exfiltrating data belonging to other people while testing. Testing against your own installation
|
||||||
|
or your own account is always fine.
|
||||||
|
|
||||||
|
## Supported Versions
|
||||||
|
|
||||||
|
We support the latest release. Security fixes ship in the next version rather than as backports to older releases, so
|
||||||
|
upgrading to the current release is how you get them.
|
||||||
|
|
||||||
|
Release notifications are available by watching [releases](https://github.com/netbirdio/netbird/releases).
|
||||||
|
|
||||||
|
## Published advisories
|
||||||
|
|
||||||
|
Every vulnerability we fix is published as a GitHub Security Advisory on the
|
||||||
|
[advisories page](https://github.com/netbirdio/netbird/security/advisories), including the affected version range, the
|
||||||
|
patched version, and the reporter's credit. Advisories for the Go module are also distributed through the Go vulnerability
|
||||||
|
database, so `govulncheck` will report them against your dependencies.
|
||||||
|
|
||||||
|
## Bug bounty
|
||||||
|
|
||||||
|
There is no official bug bounty program for the NetBird project. We credit reporters in advisories, and we are grateful for
|
||||||
|
the work, but we cannot currently offer payment for reports.
|
||||||
|
|
||||||
|
## Non-security bugs
|
||||||
|
|
||||||
|
For bugs that are not security issues, please use the
|
||||||
|
[issue tracker](https://github.com/netbirdio/netbird/discussions/new/choose).
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|
||||||
@@ -75,6 +76,24 @@ type Client struct {
|
|||||||
connectClient *internal.ConnectClient
|
connectClient *internal.ConnectClient
|
||||||
config *profilemanager.Config
|
config *profilemanager.Config
|
||||||
cacheDir string
|
cacheDir 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, cc *internal.ConnectClient) {
|
||||||
@@ -148,11 +167,16 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
|
|||||||
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, 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)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -247,6 +271,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 +323,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 +340,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 +484,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 +503,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 +553,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")
|
||||||
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -189,6 +189,19 @@ func (pm *ProfileManager) LogoutProfile(id string) error {
|
|||||||
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)
|
||||||
|
|||||||
@@ -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)
|
|
||||||
}
|
|
||||||
309
client/android/session.go
Normal file
309
client/android/session.go
Normal file
@@ -0,0 +1,309 @@
|
|||||||
|
//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, _, cc := c.stateSnapshot()
|
||||||
|
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()
|
||||||
|
|
||||||
|
a := &Auth{ctx: ctx, config: cfg}
|
||||||
|
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,7 +17,9 @@ 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"
|
||||||
)
|
)
|
||||||
@@ -331,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.
|
||||||
|
|||||||
@@ -33,10 +33,15 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
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
|
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
|
jsonServMu sync.Mutex
|
||||||
serverInstance *server.Server
|
serverInstance *server.Server
|
||||||
serverInstanceMu sync.Mutex
|
serverInstanceMu sync.Mutex
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"runtime"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/kardianos/service"
|
"github.com/kardianos/service"
|
||||||
@@ -13,6 +14,8 @@ 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"
|
||||||
@@ -26,6 +29,31 @@ func validateJSONSocketFlags() error {
|
|||||||
return nil
|
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
|
||||||
@@ -37,68 +65,106 @@ func (p *program) Start(svc service.Service) error {
|
|||||||
// 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.
|
||||||
daemonListener, err := listenOnAddress(daemonAddr)
|
if migrated, ok := daemonaddr.MigrateLegacy(daemonAddr); ok {
|
||||||
if err != nil {
|
log.Infof("daemon address %q predates named-pipe support, listening on %q so callers can be identified", daemonAddr, migrated)
|
||||||
return fmt.Errorf("listen daemon interface: %w", err)
|
daemonAddr = migrated
|
||||||
}
|
}
|
||||||
|
|
||||||
var jsonListener *socketListener
|
network, _, err := parseListenAddress(daemonAddr)
|
||||||
if enableJSONSocket {
|
if err != nil {
|
||||||
jsonListener, err = listenOnAddress(jsonSocket)
|
return fmt.Errorf("parse daemon address: %w", err)
|
||||||
if err != nil {
|
}
|
||||||
_ = daemonListener.Close()
|
|
||||||
return fmt.Errorf("listen daemon JSON interface: %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)...)
|
||||||
} else {
|
|
||||||
removeStaleUnixSocketForAddress(jsonSocket)
|
daemonListener, jsonListener, err := listenDaemonSockets()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
defer daemonListener.Close()
|
// Fatal here rather than inside serve, so serve's deferred listener
|
||||||
if jsonListener != nil {
|
// closes run before the process exits.
|
||||||
defer jsonListener.Close()
|
if err := p.serve(daemonListener, jsonListener); err != nil {
|
||||||
}
|
log.Fatalf("failed to %v", err)
|
||||||
|
|
||||||
if err := daemonListener.chmodUnixSocket("daemon"); err != nil {
|
|
||||||
log.Error(err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if jsonListener != nil {
|
|
||||||
if err := jsonListener.chmodUnixSocket("daemon JSON"); err != nil {
|
|
||||||
log.Error(err)
|
|
||||||
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()
|
|
||||||
|
|
||||||
if jsonListener != nil {
|
|
||||||
if err := p.startJSONGateway(jsonListener, daemonAddr); err != nil {
|
|
||||||
log.Fatalf("failed to start daemon JSON server: %v", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
log.Debug("daemon JSON socket disabled")
|
|
||||||
}
|
|
||||||
|
|
||||||
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
|
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 {
|
||||||
@@ -113,8 +179,13 @@ func (p *program) Stop(srv service.Service) error {
|
|||||||
p.cancel()
|
p.cancel()
|
||||||
|
|
||||||
p.jsonServMu.Lock()
|
p.jsonServMu.Lock()
|
||||||
jsonServ := p.jsonServ
|
jsonServ, jsonClient := p.jsonServ, p.jsonClient
|
||||||
p.jsonServMu.Unlock()
|
p.jsonServMu.Unlock()
|
||||||
|
if jsonClient != nil {
|
||||||
|
if err := jsonClient.Close(); err != nil {
|
||||||
|
log.Debugf("close daemon JSON gateway client: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
if jsonServ != nil {
|
if jsonServ != nil {
|
||||||
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 2*time.Second)
|
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||||
if err := jsonServ.Shutdown(shutdownCtx); err != nil {
|
if err := jsonServ.Shutdown(shutdownCtx); err != nil {
|
||||||
|
|||||||
@@ -5,27 +5,123 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
"google.golang.org/grpc/credentials/insecure"
|
|
||||||
|
|
||||||
|
"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"
|
||||||
)
|
)
|
||||||
|
|
||||||
func grpcGatewayEndpoint(addr string) string {
|
// jsonPeerIdentity is the context key under which the connecting HTTP client's
|
||||||
return strings.TrimPrefix(addr, "tcp://")
|
// 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 {
|
func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint string) error {
|
||||||
mux := runtime.NewServeMux()
|
if jsonListener.network == "tcp" {
|
||||||
opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())}
|
log.Warnf("daemon JSON socket is listening on TCP (%s): callers carry no verifiable identity over TCP, "+
|
||||||
if err := proto.RegisterDaemonServiceHandlerFromEndpoint(p.ctx, mux, grpcGatewayEndpoint(daemonEndpoint), opts); err != nil {
|
"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
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -35,10 +131,12 @@ func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint
|
|||||||
BaseContext: func(net.Listener) context.Context {
|
BaseContext: func(net.Listener) context.Context {
|
||||||
return p.ctx
|
return p.ctx
|
||||||
},
|
},
|
||||||
|
ConnContext: jsonConnContext,
|
||||||
}
|
}
|
||||||
|
|
||||||
p.jsonServMu.Lock()
|
p.jsonServMu.Lock()
|
||||||
p.jsonServ = jsonServer
|
p.jsonServ = jsonServer
|
||||||
|
p.jsonClient = conn
|
||||||
p.jsonServMu.Unlock()
|
p.jsonServMu.Unlock()
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
|
|||||||
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -125,6 +126,13 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
|
|||||||
|
|
||||||
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 != "" {
|
if !serviceCmd.PersistentFlags().Changed("json-socket") && params.JSONSocket != "" {
|
||||||
|
|||||||
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...)
|
||||||
|
}
|
||||||
@@ -26,6 +26,14 @@ func listenOnAddress(addr string) (*socketListener, error) {
|
|||||||
return nil, err
|
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" {
|
if network == "unix" {
|
||||||
removeStaleUnixSocket(address)
|
removeStaleUnixSocket(address)
|
||||||
}
|
}
|
||||||
@@ -41,11 +49,11 @@ func listenOnAddress(addr string) (*socketListener, error) {
|
|||||||
func parseListenAddress(addr string) (string, string, error) {
|
func parseListenAddress(addr string) (string, string, error) {
|
||||||
network, address, ok := strings.Cut(addr, "://")
|
network, address, ok := strings.Cut(addr, "://")
|
||||||
if !ok || network == "" || address == "" {
|
if !ok || network == "" || address == "" {
|
||||||
return "", "", fmt.Errorf("address must be in [unix|tcp]://[path|host:port] format: %q", addr)
|
return "", "", fmt.Errorf("address must be in [unix|tcp|npipe]://[path|host:port|name] format: %q", addr)
|
||||||
}
|
}
|
||||||
|
|
||||||
switch network {
|
switch network {
|
||||||
case "unix", "tcp":
|
case "unix", "tcp", "npipe":
|
||||||
return network, address, nil
|
return network, address, nil
|
||||||
default:
|
default:
|
||||||
return "", "", fmt.Errorf("unsupported daemon address protocol: %v", network)
|
return "", "", fmt.Errorf("unsupported daemon address protocol: %v", network)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -351,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,
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -24,11 +24,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// Skew tolerates a small clock difference between the management
|
maxPastHorizon = 30 * 24 * time.Hour
|
||||||
// server and this peer before treating a deadline as "in the past".
|
|
||||||
// Slightly above typical NTP drift; tight enough that the UI doesn't
|
|
||||||
// paint a stale expiry as if it were valid.
|
|
||||||
Skew = 30 * time.Second
|
|
||||||
|
|
||||||
// maxDeadlineHorizon caps how far in the future an accepted deadline
|
// maxDeadlineHorizon caps how far in the future an accepted deadline
|
||||||
// can sit. A timestamp beyond this is almost certainly a protocol
|
// can sit. A timestamp beyond this is almost certainly a protocol
|
||||||
@@ -57,7 +53,7 @@ var (
|
|||||||
ErrDeadlineTooFarFuture = errors.New("session deadline too far in the future")
|
ErrDeadlineTooFarFuture = errors.New("session deadline too far in the future")
|
||||||
|
|
||||||
// ErrDeadlineInPast is returned by Update when the supplied deadline
|
// ErrDeadlineInPast is returned by Update when the supplied deadline
|
||||||
// is more than Skew in the past.
|
// is more than maxPastHorizon in the past.
|
||||||
ErrDeadlineInPast = errors.New("session deadline in the past")
|
ErrDeadlineInPast = errors.New("session deadline in the past")
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -66,15 +62,14 @@ var (
|
|||||||
// for deadline change/clear, PublishEvent for the two warnings); tests pass
|
// for deadline change/clear, PublishEvent for the two warnings); tests pass
|
||||||
// a fake recorder so the same surface is observable without an engine.
|
// a fake recorder so the same surface is observable without an engine.
|
||||||
//
|
//
|
||||||
// The watcher is the single owner of the deadline propagated to the
|
// While the watcher runs, it owns the deadline propagated to the recorder:
|
||||||
// recorder: every set, clear, sanity-check rejection and Close routes the
|
// every set, clear and sanity-check rejection routes the value through
|
||||||
// value through SetSessionExpiresAt, so the SubscribeStatus snapshot the UI
|
// SetSessionExpiresAt, so the SubscribeStatus snapshot the UI reads can
|
||||||
// reads can never drift from the watcher's timer state. (SetSessionExpiresAt
|
// never drift from the watcher's timer state. (SetSessionExpiresAt fans
|
||||||
// fans out its own state-change notification, so no separate notify is
|
// out its own state-change notification, so no separate notify is needed.)
|
||||||
// needed.) The recorder is server-scoped and outlives this engine-scoped
|
// The recorder is server-scoped and outlives this engine-scoped watcher;
|
||||||
// watcher — without the Close-time clear a teardown (Down, or the Down+Up of
|
// Close deliberately leaves the recorder value in place so transient engine
|
||||||
// a profile switch) would leave the next session showing the previous one's
|
// restarts don't blank it — the client run loop clears it on real teardown.
|
||||||
// stale "expires in" value.
|
|
||||||
//
|
//
|
||||||
// PublishEvent's signature mirrors peer.Status.PublishEvent: the watcher
|
// PublishEvent's signature mirrors peer.Status.PublishEvent: the watcher
|
||||||
// composes the metadata internally so the wire format (MetaSession*) is
|
// composes the metadata internally so the wire format (MetaSession*) is
|
||||||
@@ -135,10 +130,13 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher {
|
|||||||
// was disabled).
|
// was disabled).
|
||||||
//
|
//
|
||||||
// Same-value updates are no-ops. A different non-zero value cancels any
|
// Same-value updates are no-ops. A different non-zero value cancels any
|
||||||
// pending timer, resets the "already fired" guard, and arms a new one.
|
// 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
|
// Returns one of the sentinel Err* values when the deadline fails the
|
||||||
// sanity checks (pre-epoch, far future, or in the past beyond Skew).
|
// sanity checks (pre-epoch, far future, or past beyond maxPastHorizon).
|
||||||
// In every error case the watcher first clears its state so it stays
|
// 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.
|
// consistent with what the caller will push into its other sinks (e.g.
|
||||||
// applySessionDeadline forces a zero deadline into the status recorder
|
// applySessionDeadline forces a zero deadline into the status recorder
|
||||||
@@ -163,7 +161,7 @@ func (w *Watcher) Update(deadline time.Time) error {
|
|||||||
case deadline.After(now.Add(maxDeadlineHorizon)):
|
case deadline.After(now.Add(maxDeadlineHorizon)):
|
||||||
w.clearLocked()
|
w.clearLocked()
|
||||||
return fmt.Errorf("%w: %v", ErrDeadlineTooFarFuture, deadline)
|
return fmt.Errorf("%w: %v", ErrDeadlineTooFarFuture, deadline)
|
||||||
case deadline.Before(now.Add(-Skew)):
|
case deadline.Before(now.Add(-maxPastHorizon)):
|
||||||
w.clearLocked()
|
w.clearLocked()
|
||||||
return fmt.Errorf("%w: %v (now=%v)", ErrDeadlineInPast, deadline, now)
|
return fmt.Errorf("%w: %v (now=%v)", ErrDeadlineInPast, deadline, now)
|
||||||
}
|
}
|
||||||
@@ -183,7 +181,9 @@ func (w *Watcher) Update(deadline time.Time) error {
|
|||||||
w.finalFiredAt = time.Time{}
|
w.finalFiredAt = time.Time{}
|
||||||
w.dismissedAt = time.Time{}
|
w.dismissedAt = time.Time{}
|
||||||
|
|
||||||
w.armTimerLocked(deadline)
|
if deadline.After(now) {
|
||||||
|
w.armTimerLocked(deadline)
|
||||||
|
}
|
||||||
recorder := w.recorder
|
recorder := w.recorder
|
||||||
w.mu.Unlock()
|
w.mu.Unlock()
|
||||||
if recorder != nil {
|
if recorder != nil {
|
||||||
@@ -227,30 +227,25 @@ func (w *Watcher) Dismiss() {
|
|||||||
log.Infof("auth session final-warning dismissed for deadline %s", w.current.Format(time.RFC3339))
|
log.Infof("auth session final-warning dismissed for deadline %s", w.current.Format(time.RFC3339))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close stops any pending timer and drops the deadline on the status
|
// Close stops any pending timer. Update calls after Close are ignored.
|
||||||
// recorder. Update calls after Close are ignored. Clearing the recorder
|
// The recorder keeps its deadline: the watcher is engine-scoped and closes
|
||||||
// here is what keeps a teardown (Down, or the Down+Up of a profile switch)
|
// on every engine restart (network change, sleep/wake, stream errors)
|
||||||
// from leaving the next session showing this one's stale "expires in"
|
// while the SSO deadline stays valid across those, so clearing here would
|
||||||
// value — the recorder is server-scoped and outlives this engine-scoped
|
// blank the UI's "expires in" row on every transient reconnect. The
|
||||||
// watcher, so nothing else drops the anchor on teardown.
|
// client run loop clears the server-scoped recorder when it exits for
|
||||||
|
// real (Down, profile switch, permanent login failure).
|
||||||
func (w *Watcher) Close() {
|
func (w *Watcher) Close() {
|
||||||
w.mu.Lock()
|
w.mu.Lock()
|
||||||
|
defer w.mu.Unlock()
|
||||||
if w.closed {
|
if w.closed {
|
||||||
w.mu.Unlock()
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
w.closed = true
|
w.closed = true
|
||||||
w.stopTimerLocked()
|
w.stopTimerLocked()
|
||||||
hadDeadline := !w.current.IsZero()
|
|
||||||
w.current = time.Time{}
|
w.current = time.Time{}
|
||||||
w.firedAt = time.Time{}
|
w.firedAt = time.Time{}
|
||||||
w.finalFiredAt = time.Time{}
|
w.finalFiredAt = time.Time{}
|
||||||
w.dismissedAt = time.Time{}
|
w.dismissedAt = time.Time{}
|
||||||
recorder := w.recorder
|
|
||||||
w.mu.Unlock()
|
|
||||||
if recorder != nil && hadDeadline {
|
|
||||||
recorder.SetSessionExpiresAt(time.Time{})
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// clearLocked drops the tracked deadline and notifies the recorder so
|
// clearLocked drops the tracked deadline and notifies the recorder so
|
||||||
|
|||||||
@@ -224,11 +224,13 @@ func TestNewDeadlineCancelsPriorTimer(t *testing.T) {
|
|||||||
|
|
||||||
func TestRefreshAfterFireArmsNewWarning(t *testing.T) {
|
func TestRefreshAfterFireArmsNewWarning(t *testing.T) {
|
||||||
r := &fakeRecorder{}
|
r := &fakeRecorder{}
|
||||||
lead := 30 * time.Millisecond
|
lead := 150 * time.Millisecond
|
||||||
w := newWatcher(lead, r)
|
w := newWatcher(lead, r)
|
||||||
defer w.Close()
|
defer w.Close()
|
||||||
|
|
||||||
first := time.Now().Add(50 * time.Millisecond)
|
// 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)
|
_ = w.Update(first)
|
||||||
|
|
||||||
// Wait for stateChange + warning of the first cycle.
|
// Wait for stateChange + warning of the first cycle.
|
||||||
@@ -306,7 +308,29 @@ func TestUpdateRejectsTooFarFuture(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUpdateInPastClearsDeadline(t *testing.T) {
|
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{}
|
r := &fakeRecorder{}
|
||||||
w := newWatcher(50*time.Millisecond, r)
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
defer w.Close()
|
defer w.Close()
|
||||||
@@ -318,12 +342,12 @@ func TestUpdateInPastClearsDeadline(t *testing.T) {
|
|||||||
// Drain the stateChange from the seed.
|
// Drain the stateChange from the seed.
|
||||||
waitForEvents(t, r, 1)
|
waitForEvents(t, r, 1)
|
||||||
|
|
||||||
err := w.Update(time.Now().Add(-1 * time.Hour))
|
err := w.Update(time.Now().Add(-31 * 24 * time.Hour))
|
||||||
if !errors.Is(err, ErrDeadlineInPast) {
|
if !errors.Is(err, ErrDeadlineInPast) {
|
||||||
t.Fatalf("want ErrDeadlineInPast, got %v", err)
|
t.Fatalf("want ErrDeadlineInPast, got %v", err)
|
||||||
}
|
}
|
||||||
if !w.Deadline().IsZero() {
|
if !w.Deadline().IsZero() {
|
||||||
t.Fatalf("in-past update must clear the deadline, got %v", w.Deadline())
|
t.Fatalf("rejected ancient-past update must clear the deadline, got %v", w.Deadline())
|
||||||
}
|
}
|
||||||
events := waitForEvents(t, r, 2)
|
events := waitForEvents(t, r, 2)
|
||||||
if events[1].kind != stateChange {
|
if events[1].kind != stateChange {
|
||||||
@@ -331,39 +355,25 @@ func TestUpdateInPastClearsDeadline(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUpdateWithinSkewAccepted(t *testing.T) {
|
|
||||||
r := &fakeRecorder{}
|
|
||||||
w := newWatcher(50*time.Millisecond, r)
|
|
||||||
defer w.Close()
|
|
||||||
|
|
||||||
// 5 seconds in the past is within the 30s Skew tolerance — accept it.
|
|
||||||
d := time.Now().Add(-5 * time.Second)
|
|
||||||
if err := w.Update(d); err != nil {
|
|
||||||
t.Fatalf("within-skew Update should succeed, got %v", err)
|
|
||||||
}
|
|
||||||
if !w.Deadline().Equal(d) {
|
|
||||||
t.Fatalf("expected deadline to be applied, got %v want %v", w.Deadline(), d)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCloseSilencesUpdates(t *testing.T) {
|
func TestCloseSilencesUpdates(t *testing.T) {
|
||||||
r := &fakeRecorder{}
|
r := &fakeRecorder{}
|
||||||
w := newWatcher(50*time.Millisecond, r)
|
w := newWatcher(50*time.Millisecond, r)
|
||||||
w.Close()
|
w.Close()
|
||||||
|
|
||||||
_ = w.Update(time.Now().Add(time.Hour))
|
if err := w.Update(time.Now().Add(time.Hour)); err != nil {
|
||||||
|
t.Fatalf("Update after Close: want nil, got %v", err)
|
||||||
time.Sleep(20 * time.Millisecond)
|
}
|
||||||
if got := r.snapshot(); len(got) != 0 {
|
if got := r.snapshot(); len(got) != 0 {
|
||||||
t.Fatalf("expected no events after Close, got %+v", got)
|
t.Fatalf("expected no events after Close, got %+v", got)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestCloseClearsRecorderDeadline pins the profile-switch fix: a watcher
|
// TestCloseKeepsRecorderDeadline pins the reconnect-flap fix: the watcher
|
||||||
// holding a live deadline must zero the recorder on Close so the next
|
// closes on every engine restart (network change, sleep/wake) while the
|
||||||
// engine's watcher (and the UI reading the shared server-scoped recorder)
|
// SSO deadline stays valid across those, so Close must leave the
|
||||||
// doesn't start out showing the previous session's stale "expires in".
|
// server-scoped recorder's value in place. The client run loop clears the
|
||||||
func TestCloseClearsRecorderDeadline(t *testing.T) {
|
// recorder when it exits for real.
|
||||||
|
func TestCloseKeepsRecorderDeadline(t *testing.T) {
|
||||||
r := &fakeRecorder{}
|
r := &fakeRecorder{}
|
||||||
w := newWatcher(time.Hour, r)
|
w := newWatcher(time.Hour, r)
|
||||||
|
|
||||||
@@ -377,8 +387,8 @@ func TestCloseClearsRecorderDeadline(t *testing.T) {
|
|||||||
|
|
||||||
w.Close()
|
w.Close()
|
||||||
|
|
||||||
if got := r.deadline(); !got.IsZero() {
|
if got := r.deadline(); !got.Equal(d) {
|
||||||
t.Fatalf("recorder deadline after Close = %v, want zero", got)
|
t.Fatalf("recorder deadline after Close = %v, want %v", got, d)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -136,10 +137,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 +261,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 {
|
||||||
@@ -618,6 +625,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),
|
||||||
|
|
||||||
@@ -693,6 +701,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
|
||||||
|
}
|
||||||
@@ -480,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,
|
||||||
@@ -677,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))
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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 {
|
||||||
@@ -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()
|
||||||
@@ -1435,11 +1443,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,6 +224,13 @@ 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
|
||||||
|
|
||||||
|
// latestComponents is the most-recent NetworkMapComponents decoded from
|
||||||
|
// a NetworkMapEnvelope (capability=3 peers only). Held alongside the
|
||||||
|
// NetworkMap that Calculate() produced from it so future incremental
|
||||||
|
// updates have a base to apply changes against. nil for legacy-format
|
||||||
|
// peers. Guarded by syncMsgMux.
|
||||||
|
latestComponents *types.NetworkMapComponents
|
||||||
|
|
||||||
networkMonitor *networkmonitor.NetworkMonitor
|
networkMonitor *networkmonitor.NetworkMonitor
|
||||||
|
|
||||||
sshServer sshServer
|
sshServer sshServer
|
||||||
@@ -551,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)
|
||||||
}
|
}
|
||||||
@@ -652,8 +663,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())
|
||||||
|
|
||||||
@@ -963,8 +990,12 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
|
|||||||
|
|
||||||
e.ApplySessionDeadline(update.GetSessionExpiresAt())
|
e.ApplySessionDeadline(update.GetSessionExpiresAt())
|
||||||
|
|
||||||
if update.NetworkMap != nil && update.NetworkMap.PeerConfig != nil {
|
// Envelope sync responses carry PeerConfig at the top level; legacy
|
||||||
e.handleAutoUpdateVersion(update.NetworkMap.PeerConfig.AutoUpdate)
|
// 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")
|
||||||
@@ -974,12 +1005,47 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
|
|||||||
return 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:
|
||||||
// NetworkMap != nil, checks present -> apply the received checks
|
// NetworkMap != nil, checks present -> apply the received checks
|
||||||
// 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 nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -992,6 +1058,14 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
done = e.phase("persist")
|
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)
|
e.persistSyncResponse(update)
|
||||||
done()
|
done()
|
||||||
|
|
||||||
@@ -1005,6 +1079,19 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
|
|||||||
return 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:
|
||||||
// STUN/TURN and relay servers, flow logging and DNS settings. A nil config is a no-op,
|
// STUN/TURN and relay servers, flow logging and DNS settings. A nil config is a no-op,
|
||||||
// which is the case for sync updates carrying only a network map.
|
// which is the case for sync updates carrying only a network map.
|
||||||
@@ -1164,6 +1251,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,
|
||||||
@@ -2032,6 +2120,7 @@ func (e *Engine) readInitialSettings() ([]*route.Route, *nbdns.Config, bool, err
|
|||||||
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,
|
||||||
@@ -2605,13 +2694,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
|
||||||
}
|
}
|
||||||
@@ -2621,6 +2711,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 {
|
||||||
|
|||||||
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
|
||||||
|
}
|
||||||
@@ -75,4 +75,14 @@ func TestApplySessionDeadline_ThreeState(t *testing.T) {
|
|||||||
require.True(t, e.statusRecorder.GetSessionExpiresAt().IsZero(),
|
require.True(t, e.statusRecorder.GetSessionExpiresAt().IsZero(),
|
||||||
"invalid timestamp must clear the deadline")
|
"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")
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
31
client/internal/ipcauth/creds_stub.go
Normal file
31
client/internal/ipcauth/creds_stub.go
Normal file
@@ -0,0 +1,31 @@
|
|||||||
|
//go:build !linux && !darwin && !freebsd && !windows
|
||||||
|
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"google.golang.org/grpc/credentials"
|
||||||
|
)
|
||||||
|
|
||||||
|
// errUnsupported is returned on platforms with no local peer-identity
|
||||||
|
// primitive, so consumers fail closed instead of guessing an identity.
|
||||||
|
var errUnsupported = errors.New("peer identity is not available on this platform")
|
||||||
|
|
||||||
|
// NewTransportCredentials returns nil: without a peer-identity primitive the
|
||||||
|
// daemon cannot authenticate local callers, and the caller must treat that as
|
||||||
|
// "authorization cannot be enforced".
|
||||||
|
func NewTransportCredentials() credentials.TransportCredentials {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// PeerIdentity always fails on this platform.
|
||||||
|
func PeerIdentity(net.Conn) (Identity, error) {
|
||||||
|
return Identity{}, errUnsupported
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConnIdentity always fails on this platform.
|
||||||
|
func ConnIdentity(net.Conn) (Identity, error) {
|
||||||
|
return Identity{}, errUnsupported
|
||||||
|
}
|
||||||
56
client/internal/ipcauth/creds_unix.go
Normal file
56
client/internal/ipcauth/creds_unix.go
Normal file
@@ -0,0 +1,56 @@
|
|||||||
|
//go:build linux || darwin || freebsd
|
||||||
|
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"google.golang.org/grpc/credentials"
|
||||||
|
)
|
||||||
|
|
||||||
|
// NewTransportCredentials returns gRPC transport credentials that expose the
|
||||||
|
// caller's kernel-authenticated identity via IdentityFromContext. It returns
|
||||||
|
// nil on platforms that have no peer-identity primitive, which the caller must
|
||||||
|
// treat as "authorization cannot be enforced".
|
||||||
|
//
|
||||||
|
// The handshake exchanges no bytes on the wire, so a client dialing with
|
||||||
|
// insecure credentials interoperates with a server using these. That keeps
|
||||||
|
// older CLI and UI binaries working against an upgraded daemon.
|
||||||
|
func NewTransportCredentials() credentials.TransportCredentials {
|
||||||
|
return unixCreds{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConnIdentity extracts the caller's identity from an accepted local IPC
|
||||||
|
// connection. It is shared by the gRPC transport credentials and by the JSON
|
||||||
|
// gateway, which reads the identity of its own HTTP clients.
|
||||||
|
func ConnIdentity(conn net.Conn) (Identity, error) {
|
||||||
|
return PeerIdentity(conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
type unixCreds struct{}
|
||||||
|
|
||||||
|
func (unixCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
|
||||||
|
return conn, AuthInfo{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ServerHandshake extracts the peer identity and fails closed when it cannot
|
||||||
|
// be read, so a connection whose caller is unknown never reaches a handler.
|
||||||
|
func (unixCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
|
||||||
|
id, err := ConnIdentity(conn)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
return conn, AuthInfo{
|
||||||
|
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
|
||||||
|
Identity: id,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (unixCreds) Info() credentials.ProtocolInfo {
|
||||||
|
return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (unixCreds) Clone() credentials.TransportCredentials { return unixCreds{} }
|
||||||
|
|
||||||
|
func (unixCreds) OverrideServerName(string) error { return nil }
|
||||||
194
client/internal/ipcauth/creds_windows.go
Normal file
194
client/internal/ipcauth/creds_windows.go
Normal file
@@ -0,0 +1,194 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"runtime"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
"google.golang.org/grpc/credentials"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
modadvapi32 = windows.NewLazySystemDLL("advapi32.dll")
|
||||||
|
procImpersonateNamedPipeClient = modadvapi32.NewProc("ImpersonateNamedPipeClient")
|
||||||
|
)
|
||||||
|
|
||||||
|
// DefaultPipeSDDL is the security descriptor for the daemon control pipe.
|
||||||
|
//
|
||||||
|
// D:P protected DACL, no inheritance
|
||||||
|
// (A;;GA;;;SY) allow GENERIC_ALL to LocalSystem (the daemon's service account)
|
||||||
|
// (A;;GA;;;WD) allow GENERIC_ALL to Everyone
|
||||||
|
//
|
||||||
|
// Any local caller may connect, as with a Unix socket at 0666; what a caller may
|
||||||
|
// actually do is decided from its token, not from the DACL. Remote callers are not
|
||||||
|
// a concern here: winio.ListenPipe creates the pipe with
|
||||||
|
// FILE_PIPE_REJECT_REMOTE_CLIENTS, so NPFS rejects connections from other machines
|
||||||
|
// before the descriptor is consulted.
|
||||||
|
//
|
||||||
|
// A deny ACE on the NETWORK SID would not add anything and would break callers:
|
||||||
|
// that SID is present in any network-logon token, which includes OpenSSH and WinRM
|
||||||
|
// sessions, so it denies administrators driving the CLI over SSH and denies the
|
||||||
|
// daemon itself when started from such a session.
|
||||||
|
func DefaultPipeSDDL() string {
|
||||||
|
return "D:P(A;;GA;;;SY)(A;;GA;;;WD)"
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewTransportCredentials returns gRPC transport credentials that derive the
|
||||||
|
// caller's identity from the named-pipe client token.
|
||||||
|
//
|
||||||
|
// The client must connect at SECURITY_IDENTIFICATION for the daemon to be able
|
||||||
|
// to read its token, which is what DialNamedPipe does.
|
||||||
|
func NewTransportCredentials() credentials.TransportCredentials {
|
||||||
|
return winpipeCreds{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConnIdentity extracts the caller's identity from an accepted named-pipe
|
||||||
|
// connection by impersonating the pipe client and reading its token. It is
|
||||||
|
// shared by the gRPC transport credentials and by the JSON gateway, which
|
||||||
|
// reads the identity of its own HTTP clients.
|
||||||
|
func ConnIdentity(conn net.Conn) (Identity, error) {
|
||||||
|
// go-winio's pipe connection embeds *win32File, which exposes Fd().
|
||||||
|
fdConn, ok := conn.(interface{ Fd() uintptr })
|
||||||
|
if !ok {
|
||||||
|
return Identity{}, fmt.Errorf("connection %T does not expose a pipe handle", conn)
|
||||||
|
}
|
||||||
|
return pipeClientIdentity(windows.Handle(fdConn.Fd()))
|
||||||
|
}
|
||||||
|
|
||||||
|
type winpipeCreds struct{}
|
||||||
|
|
||||||
|
func (winpipeCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
|
||||||
|
return conn, AuthInfo{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ServerHandshake extracts the connecting client's identity and fails closed
|
||||||
|
// when the handle or token cannot be read, so a connection whose caller is
|
||||||
|
// unknown never reaches a handler.
|
||||||
|
func (winpipeCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
|
||||||
|
id, err := ConnIdentity(conn)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
return conn, AuthInfo{
|
||||||
|
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
|
||||||
|
Identity: id,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (winpipeCreds) Info() credentials.ProtocolInfo {
|
||||||
|
return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (winpipeCreds) Clone() credentials.TransportCredentials { return winpipeCreds{} }
|
||||||
|
|
||||||
|
func (winpipeCreds) OverrideServerName(string) error { return nil }
|
||||||
|
|
||||||
|
// pipeClientIdentity reads the connecting client's user SID, usable group
|
||||||
|
// SIDs, and elevation state by impersonating the pipe client on this thread
|
||||||
|
// and reading the resulting impersonation token.
|
||||||
|
func pipeClientIdentity(handle windows.Handle) (id Identity, err error) {
|
||||||
|
// Impersonation is per-thread, so the goroutine must stay on this thread
|
||||||
|
// until RevertToSelf, otherwise an unrelated goroutine could inherit the
|
||||||
|
// impersonated context.
|
||||||
|
runtime.LockOSThread()
|
||||||
|
|
||||||
|
// The thread only goes back to the runtime's pool once it is provably no
|
||||||
|
// longer impersonating the client. If the revert fails, leaving it locked
|
||||||
|
// makes Go terminate it when this goroutine exits, which costs one thread
|
||||||
|
// and keeps a thread running as the client from ever being reused.
|
||||||
|
clean := false
|
||||||
|
defer func() {
|
||||||
|
if clean {
|
||||||
|
runtime.UnlockOSThread()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
if err = impersonateNamedPipeClient(handle); err != nil {
|
||||||
|
clean = true
|
||||||
|
return Identity{}, fmt.Errorf("impersonate named pipe client: %w", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
// Surface the revert failure only when nothing else failed: leaving
|
||||||
|
// the thread impersonated is worse than the original error.
|
||||||
|
revErr := windows.RevertToSelf()
|
||||||
|
if revErr != nil {
|
||||||
|
if err == nil {
|
||||||
|
err = fmt.Errorf("revert impersonation: %w", revErr)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
clean = true
|
||||||
|
}()
|
||||||
|
|
||||||
|
// openAsSelf=true opens the token with the daemon's own process context
|
||||||
|
// rather than the impersonated client's, so the open cannot fail because
|
||||||
|
// the client lacks access to its own token.
|
||||||
|
var token windows.Token
|
||||||
|
if err = windows.OpenThreadToken(windows.CurrentThread(), windows.TOKEN_QUERY, true, &token); err != nil {
|
||||||
|
return Identity{}, fmt.Errorf("open thread token: %w", err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if cerr := token.Close(); cerr != nil {
|
||||||
|
log.Debugf("close client token: %v", cerr)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
return identityFromToken(token)
|
||||||
|
}
|
||||||
|
|
||||||
|
// identityFromToken reads the user SID, usable group SIDs and elevation state
|
||||||
|
// out of a Windows token.
|
||||||
|
func identityFromToken(token windows.Token) (Identity, error) {
|
||||||
|
user, err := token.GetTokenUser()
|
||||||
|
if err != nil {
|
||||||
|
return Identity{}, fmt.Errorf("read token user: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
groups, err := tokenGroupSIDs(token)
|
||||||
|
if err != nil {
|
||||||
|
return Identity{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return Identity{
|
||||||
|
SID: user.User.Sid.String(),
|
||||||
|
Groups: groups,
|
||||||
|
Elevated: token.IsElevated(),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// tokenGroupSIDs returns the SIDs of the groups the token can actually
|
||||||
|
// exercise. Groups that are disabled or marked deny-only are skipped: a
|
||||||
|
// UAC-filtered administrator carries BUILTIN\Administrators as deny-only, and
|
||||||
|
// treating that as membership would hand every admin account privilege it
|
||||||
|
// cannot currently use.
|
||||||
|
func tokenGroupSIDs(token windows.Token) ([]string, error) {
|
||||||
|
tg, err := token.GetTokenGroups()
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("read token groups: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var sids []string
|
||||||
|
for _, g := range tg.AllGroups() {
|
||||||
|
if g.Attributes&windows.SE_GROUP_ENABLED == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if g.Attributes&windows.SE_GROUP_USE_FOR_DENY_ONLY != 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
sids = append(sids, g.Sid.String())
|
||||||
|
}
|
||||||
|
return sids, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func impersonateNamedPipeClient(h windows.Handle) error {
|
||||||
|
r, _, e := procImpersonateNamedPipeClient.Call(uintptr(h))
|
||||||
|
if r == 0 {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
272
client/internal/ipcauth/forward.go
Normal file
272
client/internal/ipcauth/forward.go
Normal file
@@ -0,0 +1,272 @@
|
|||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/subtle"
|
||||||
|
"encoding/hex"
|
||||||
|
"fmt"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"google.golang.org/grpc/metadata"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Metadata keys the local JSON gateway uses to forward the identity of its own
|
||||||
|
// HTTP client to the daemon. The gateway runs inside the daemon process and
|
||||||
|
// re-dials the daemon over the control socket, so without forwarding every
|
||||||
|
// JSON request would appear to come from the daemon itself.
|
||||||
|
const (
|
||||||
|
// mdFwd marks a request as forwarded by the JSON gateway. It is always
|
||||||
|
// set, even when the gateway could not read its client's identity, so the
|
||||||
|
// daemon can tell "no identity available" apart from "not forwarded".
|
||||||
|
mdFwd = "x-netbird-fwd"
|
||||||
|
mdFwdUID = "x-netbird-fwd-uid" // Unix user ID
|
||||||
|
mdFwdGID = "x-netbird-fwd-gid" // Unix primary group ID
|
||||||
|
mdFwdSID = "x-netbird-fwd-sid" // Windows user SID
|
||||||
|
mdFwdGroup = "x-netbird-fwd-group" // Windows group SID, repeated
|
||||||
|
mdFwdElevated = "x-netbird-fwd-elevated" // Windows, "1" when elevated
|
||||||
|
|
||||||
|
// mdFwdProof proves the forwarded identity was stamped by this process. The
|
||||||
|
// gateway runs inside the daemon, so a secret held in memory is available to
|
||||||
|
// the only legitimate producer and to nothing else.
|
||||||
|
mdFwdProof = "x-netbird-fwd-proof"
|
||||||
|
)
|
||||||
|
|
||||||
|
// forwardKeys is every metadata key the gateway sets. An HTTP client must never
|
||||||
|
// be able to supply one itself: see IsReservedForwardKey.
|
||||||
|
var forwardKeys = []string{mdFwd, mdFwdUID, mdFwdGID, mdFwdSID, mdFwdGroup, mdFwdElevated, mdFwdProof}
|
||||||
|
|
||||||
|
// forwardProof authenticates the gateway's forwarding metadata. It is generated
|
||||||
|
// once per daemon process and never leaves it: it is not written to disk, not
|
||||||
|
// logged, and not sent anywhere except over the daemon's own control socket to
|
||||||
|
// itself.
|
||||||
|
//
|
||||||
|
// Without it, trusting a forwarded identity rests on every layer in front of it
|
||||||
|
// stripping incoming forwarding keys, and on each key's value shape being
|
||||||
|
// distinguishable from an injected one. A single injected group SID or an
|
||||||
|
// injected "elevated" flag has the same shape as a legitimate one, so no
|
||||||
|
// cardinality rule can catch it. Requiring the proof means metadata that did not
|
||||||
|
// come from this process is refused whatever it contains.
|
||||||
|
var forwardProof = mustForwardProof()
|
||||||
|
|
||||||
|
func mustForwardProof() string {
|
||||||
|
var buf [32]byte
|
||||||
|
if _, err := rand.Read(buf[:]); err != nil {
|
||||||
|
// Continuing would leave the forwarded path authenticated by a
|
||||||
|
// predictable value, which is worse than not starting.
|
||||||
|
panic(fmt.Sprintf("generate identity forwarding proof: %v", err))
|
||||||
|
}
|
||||||
|
return hex.EncodeToString(buf[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsReservedForwardKey reports whether a gRPC metadata key belongs to the
|
||||||
|
// gateway's identity forwarding, and therefore must be dropped when it arrives
|
||||||
|
// from outside.
|
||||||
|
//
|
||||||
|
// grpc-gateway maps "Grpc-Metadata-<key>" request headers into gRPC metadata and
|
||||||
|
// joins them ahead of the values its own annotators add. Without dropping these,
|
||||||
|
// an HTTP client could hand the daemon "x-netbird-fwd-uid: 0" and be believed,
|
||||||
|
// because the daemon trusts forwarded metadata when the transport peer is the
|
||||||
|
// (privileged) gateway.
|
||||||
|
func IsReservedForwardKey(key string) bool {
|
||||||
|
key = strings.ToLower(key)
|
||||||
|
return slices.Contains(forwardKeys, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ForwardIdentityMetadata encodes an HTTP client's identity for the JSON
|
||||||
|
// gateway to forward to the daemon. When known is false only the marker is
|
||||||
|
// set, which makes the daemon treat the caller as unidentified rather than as
|
||||||
|
// the daemon itself.
|
||||||
|
func ForwardIdentityMetadata(id Identity, known bool) metadata.MD {
|
||||||
|
md := metadata.MD{}
|
||||||
|
md.Set(mdFwd, "1")
|
||||||
|
md.Set(mdFwdProof, forwardProof)
|
||||||
|
if !known {
|
||||||
|
return md
|
||||||
|
}
|
||||||
|
|
||||||
|
if id.IsWindows() {
|
||||||
|
md.Set(mdFwdSID, id.SID)
|
||||||
|
if len(id.Groups) > 0 {
|
||||||
|
md.Set(mdFwdGroup, id.Groups...)
|
||||||
|
}
|
||||||
|
if id.Elevated {
|
||||||
|
md.Set(mdFwdElevated, "1")
|
||||||
|
}
|
||||||
|
return md
|
||||||
|
}
|
||||||
|
|
||||||
|
md.Set(mdFwdUID, strconv.FormatUint(uint64(id.UID), 10))
|
||||||
|
md.Set(mdFwdGID, strconv.FormatUint(uint64(id.GID), 10))
|
||||||
|
return md
|
||||||
|
}
|
||||||
|
|
||||||
|
// CallerIdentity returns the identity to authorize a request against. For a
|
||||||
|
// direct connection that is the transport peer's kernel identity. For a
|
||||||
|
// request relayed by the local JSON gateway it is the identity the gateway
|
||||||
|
// forwarded, since the transport peer is then the daemon itself.
|
||||||
|
//
|
||||||
|
// A forwarded identity is only honoured when the transport peer is the daemon's
|
||||||
|
// own identity and the metadata carries this process's forwarding proof, so
|
||||||
|
// forged forwarding metadata gains a caller nothing. A forwarded request that
|
||||||
|
// carries no identity is reported as unidentified, never as the daemon.
|
||||||
|
//
|
||||||
|
// The second return value is false when no identity could be established, and
|
||||||
|
// callers MUST fail closed in that case.
|
||||||
|
func CallerIdentity(ctx context.Context) (Identity, bool) {
|
||||||
|
id, ok := IdentityFromContext(ctx)
|
||||||
|
if !ok {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// A forwarding key that arrives more than once did not come from the gateway
|
||||||
|
// alone, so nothing about the request can be trusted to describe its caller.
|
||||||
|
// Refusing outright matters because the alternative reading, "not forwarded",
|
||||||
|
// would authorize the request as the transport peer, which on the gateway's
|
||||||
|
// connection is the daemon itself.
|
||||||
|
if duplicatedForwardKey(ctx) {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
forwarded := isForwarded(ctx)
|
||||||
|
|
||||||
|
// Our own process on the other end of the socket is the JSON gateway, the only
|
||||||
|
// thing that dials the daemon from inside it. Such a call must carry a
|
||||||
|
// forwarded identity; without one there is no caller to authorize, and
|
||||||
|
// treating it as the daemon would authorize whatever reached the JSON socket.
|
||||||
|
// Only Linux reports the peer PID, so this is a belt on top of the gateway's
|
||||||
|
// interceptor rather than the sole guarantee.
|
||||||
|
if id.PID != 0 && int(id.PID) == selfPID && !forwarded {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only the gateway's own connection may speak for someone else. Being
|
||||||
|
// privileged is not enough and not the point: the gateway runs inside the
|
||||||
|
// daemon, so it dials as the daemon's identity whatever user that is, which
|
||||||
|
// also covers a rootless container.
|
||||||
|
if !forwarded || !IsDaemonSelf(id) {
|
||||||
|
return id, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// Speaking for someone else additionally requires the proof only this process
|
||||||
|
// holds. Refusing is the only safe reading: the transport peer here is the
|
||||||
|
// daemon itself, so falling back to it would authorize the request as the
|
||||||
|
// daemon. This is also what makes the forwarded values trustworthy once
|
||||||
|
// accepted, so they need no shape checks of their own.
|
||||||
|
if !authenticForward(ctx) {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
return forwardedIdentity(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// duplicatedForwardKey reports whether any forwarding key carries more than one
|
||||||
|
// value. The gateway's interceptor sets each key exactly once and replaces what
|
||||||
|
// was already there, so a repeat means a second source supplied it.
|
||||||
|
func duplicatedForwardKey(ctx context.Context) bool {
|
||||||
|
md, ok := metadata.FromIncomingContext(ctx)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, key := range forwardKeys {
|
||||||
|
// Group SIDs are legitimately repeated; the rest identify the caller.
|
||||||
|
if key == mdFwdGroup {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if len(md.Get(key)) > 1 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// authenticForward reports whether the request carries this process's forwarding
|
||||||
|
// proof, which only the in-process JSON gateway can supply.
|
||||||
|
func authenticForward(ctx context.Context) bool {
|
||||||
|
md, ok := metadata.FromIncomingContext(ctx)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
got := mdSingle(md, mdFwdProof)
|
||||||
|
return subtle.ConstantTimeCompare([]byte(got), []byte(forwardProof)) == 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// isForwarded reports whether the request carries the JSON gateway marker.
|
||||||
|
func isForwarded(ctx context.Context) bool {
|
||||||
|
md, ok := metadata.FromIncomingContext(ctx)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return mdSingle(md, mdFwd) != ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// forwardedIdentity decodes the identity the JSON gateway attached.
|
||||||
|
func forwardedIdentity(ctx context.Context) (Identity, bool) {
|
||||||
|
md, ok := metadata.FromIncomingContext(ctx)
|
||||||
|
if !ok {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
if sid := mdSingle(md, mdFwdSID); sid != "" {
|
||||||
|
return Identity{
|
||||||
|
SID: sid,
|
||||||
|
// Repeated by design, one value per group, and only reachable once
|
||||||
|
// the forwarding proof has been verified.
|
||||||
|
Groups: md.Get(mdFwdGroup),
|
||||||
|
Elevated: mdSingle(md, mdFwdElevated) == "1",
|
||||||
|
}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
uid, err := strconv.ParseUint(mdSingle(md, mdFwdUID), 10, 32)
|
||||||
|
if err != nil {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
id := Identity{UID: uint32(uid)}
|
||||||
|
if gid, err := strconv.ParseUint(mdSingle(md, mdFwdGID), 10, 32); err == nil {
|
||||||
|
id.GID = uint32(gid)
|
||||||
|
}
|
||||||
|
return id, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// mdSingle returns the value of a forwarded key only when exactly one was
|
||||||
|
// supplied. The gateway's interceptor sets each key exactly once, so more than one
|
||||||
|
// value means something else also supplied it, and the whole identity is treated as
|
||||||
|
// unknown rather than picking a winner. Defence in depth behind the gateway's
|
||||||
|
// header filter.
|
||||||
|
func mdSingle(md metadata.MD, key string) string {
|
||||||
|
if v := md.Get(key); len(v) == 1 {
|
||||||
|
return v[0]
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithForwardedIdentity stamps id onto a context's outgoing metadata for the JSON
|
||||||
|
// gateway's call to the daemon, replacing any forwarding keys already present so
|
||||||
|
// values supplied from outside cannot survive alongside it.
|
||||||
|
//
|
||||||
|
// This is deliberately not done with runtime.WithMetadata: grpc-gateway skips its
|
||||||
|
// annotators entirely when no request header maps to metadata ("if len(pairs) == 0
|
||||||
|
// { return ctx, nil, nil }", runtime/context.go), which an HTTP/1.0 request with no
|
||||||
|
// Host header over a unix socket achieves. The daemon would then see an unmarked
|
||||||
|
// call whose transport peer is the daemon's own identity, and authorize it as the
|
||||||
|
// daemon. A client interceptor runs for every RPC regardless of headers.
|
||||||
|
func WithForwardedIdentity(ctx context.Context, id Identity, known bool) context.Context {
|
||||||
|
md, ok := metadata.FromOutgoingContext(ctx)
|
||||||
|
if !ok {
|
||||||
|
md = metadata.MD{}
|
||||||
|
} else {
|
||||||
|
md = md.Copy()
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, key := range forwardKeys {
|
||||||
|
delete(md, key)
|
||||||
|
}
|
||||||
|
for key, values := range ForwardIdentityMetadata(id, known) {
|
||||||
|
md[key] = values
|
||||||
|
}
|
||||||
|
|
||||||
|
return metadata.NewOutgoingContext(ctx, md)
|
||||||
|
}
|
||||||
214
client/internal/ipcauth/forward_test.go
Normal file
214
client/internal/ipcauth/forward_test.go
Normal file
@@ -0,0 +1,214 @@
|
|||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"google.golang.org/grpc/credentials"
|
||||||
|
"google.golang.org/grpc/metadata"
|
||||||
|
"google.golang.org/grpc/peer"
|
||||||
|
)
|
||||||
|
|
||||||
|
// transportCtx builds a request context as the daemon's transport credentials
|
||||||
|
// would: the identity of whoever opened the socket, plus whatever metadata the
|
||||||
|
// request carried.
|
||||||
|
func transportCtx(id Identity, md metadata.MD) context.Context {
|
||||||
|
ctx := peer.NewContext(context.Background(), &peer.Peer{
|
||||||
|
AuthInfo: AuthInfo{
|
||||||
|
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
|
||||||
|
Identity: id,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if md != nil {
|
||||||
|
ctx = metadata.NewIncomingContext(ctx, md)
|
||||||
|
}
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
root = Identity{UID: 0}
|
||||||
|
unprivUser = Identity{UID: 1000, GID: 1000}
|
||||||
|
)
|
||||||
|
|
||||||
|
// asDaemon pins which identity counts as this process for the duration of a test.
|
||||||
|
// Without it the test binary's own uid decides, which silently changes what
|
||||||
|
// "the gateway" means.
|
||||||
|
func asDaemon(t *testing.T, id Identity) {
|
||||||
|
t.Helper()
|
||||||
|
prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate
|
||||||
|
t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate })
|
||||||
|
selfIdentity, selfKnown = id, true
|
||||||
|
selfMayDelegate = !id.IsPrivileged()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallerIdentity_DirectConnections(t *testing.T) {
|
||||||
|
t.Run("no transport credentials is not an identity", func(t *testing.T) {
|
||||||
|
if _, ok := CallerIdentity(context.Background()); ok {
|
||||||
|
t.Fatal("a caller with no credentials must not be identified")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a direct caller is its transport identity", func(t *testing.T) {
|
||||||
|
id, ok := CallerIdentity(transportCtx(unprivUser, nil))
|
||||||
|
if !ok || id.UID != 1000 {
|
||||||
|
t.Fatalf("got %v ok=%t, want uid 1000", id, ok)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
// The whole point of honouring forwarded metadata only from a privileged
|
||||||
|
// transport peer: an unprivileged caller can set any metadata it likes on its
|
||||||
|
// own connection to the daemon socket.
|
||||||
|
t.Run("an unprivileged caller cannot forge an identity", func(t *testing.T) {
|
||||||
|
asDaemon(t, root)
|
||||||
|
forged := metadata.Pairs(mdFwd, "1", mdFwdUID, "0", mdFwdGID, "0")
|
||||||
|
id, ok := CallerIdentity(transportCtx(unprivUser, forged))
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("caller should still be identified, as itself")
|
||||||
|
}
|
||||||
|
if id.IsPrivileged() || id.UID != 1000 {
|
||||||
|
t.Fatalf("forged metadata was believed: got %v", id)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallerIdentity_GatewayForwarding(t *testing.T) {
|
||||||
|
t.Run("the gateway's client identity is used, not the gateway's own", func(t *testing.T) {
|
||||||
|
asDaemon(t, root)
|
||||||
|
md := ForwardIdentityMetadata(unprivUser, true)
|
||||||
|
id, ok := CallerIdentity(transportCtx(root, md))
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("forwarded identity should be usable")
|
||||||
|
}
|
||||||
|
if id.IsPrivileged() || id.UID != 1000 {
|
||||||
|
t.Fatalf("got %v, want the forwarded uid 1000 and not privileged", id)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a privileged gateway client stays privileged", func(t *testing.T) {
|
||||||
|
asDaemon(t, root)
|
||||||
|
md := ForwardIdentityMetadata(root, true)
|
||||||
|
id, ok := CallerIdentity(transportCtx(root, md))
|
||||||
|
if !ok || !id.IsPrivileged() {
|
||||||
|
t.Fatalf("got %v ok=%t, want a privileged identity", id, ok)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
// A JSON socket the gateway cannot read peer credentials from (a TCP socket,
|
||||||
|
// say) must not make every request look like the daemon itself.
|
||||||
|
t.Run("an unreadable client identity is unknown, not the daemon", func(t *testing.T) {
|
||||||
|
asDaemon(t, root)
|
||||||
|
md := ForwardIdentityMetadata(Identity{}, false)
|
||||||
|
if _, ok := CallerIdentity(transportCtx(root, md)); ok {
|
||||||
|
t.Fatal("a forwarded request with no identity must not be identified")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
// grpc-gateway turns Grpc-Metadata-<key> headers into gRPC metadata and joins
|
||||||
|
// them ahead of its annotators' values. If an HTTP client's header survived
|
||||||
|
// that, this is the shape the daemon would see: the attacker's uid 0 first,
|
||||||
|
// the real uid second. The gateway filters those headers out, and reading a
|
||||||
|
// duplicated key as unknown makes the daemon safe even if it did not.
|
||||||
|
t.Run("a duplicated key from an injected header is not believed", func(t *testing.T) {
|
||||||
|
asDaemon(t, root)
|
||||||
|
md := metadata.MD{}
|
||||||
|
md.Append(mdFwd, "1")
|
||||||
|
md.Append(mdFwdUID, "0") // injected by the HTTP client
|
||||||
|
md.Append(mdFwdUID, "1000") // appended by the gateway's annotator
|
||||||
|
if id, ok := CallerIdentity(transportCtx(root, md)); ok {
|
||||||
|
t.Fatalf("injected uid was accepted: got %v", id)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("a duplicated marker is not believed either", func(t *testing.T) {
|
||||||
|
asDaemon(t, root)
|
||||||
|
md := metadata.MD{}
|
||||||
|
md.Append(mdFwd, "1")
|
||||||
|
md.Append(mdFwd, "1")
|
||||||
|
md.Append(mdFwdUID, "1000")
|
||||||
|
// A repeated marker must not be read as "not forwarded": that would
|
||||||
|
// authorize the request as the transport peer, which on the gateway's
|
||||||
|
// connection is the daemon itself.
|
||||||
|
if id, ok := CallerIdentity(transportCtx(root, md)); ok {
|
||||||
|
t.Fatalf("a duplicated marker was believed: got %v", id)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
// The layers in front of this (the gateway's header matcher, and its
|
||||||
|
// interceptor replacing every forwarding key) are what keep outside metadata
|
||||||
|
// from arriving at all. The proof is what the daemon can check for itself, and
|
||||||
|
// it is the only defence that works for a value whose legitimate shape is
|
||||||
|
// indistinguishable from an injected one: a lone group SID, or "elevated".
|
||||||
|
t.Run("forwarding metadata without this process's proof is refused", func(t *testing.T) {
|
||||||
|
asDaemon(t, root)
|
||||||
|
for name, md := range map[string]metadata.MD{
|
||||||
|
"no proof": metadata.Pairs(mdFwd, "1", mdFwdUID, "0"),
|
||||||
|
"wrong proof": metadata.Pairs(mdFwd, "1", mdFwdUID, "0", mdFwdProof, "deadbeef"),
|
||||||
|
"windows identity without a proof": metadata.Pairs(mdFwd, "1",
|
||||||
|
mdFwdSID, "S-1-5-21-1-2-3-1001", mdFwdGroup, sidAdministrators, mdFwdElevated, "1"),
|
||||||
|
} {
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
if id, ok := CallerIdentity(transportCtx(root, md)); ok {
|
||||||
|
t.Fatalf("unstamped forwarding metadata was believed: got %v", id)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
// A caller that reaches the gateway cannot see the proof, so it cannot append
|
||||||
|
// a group of its own to a genuine forwarded identity: doing so would have to
|
||||||
|
// go through the interceptor, which replaces the whole set.
|
||||||
|
t.Run("a group appended to a stamped identity does not survive the interceptor", func(t *testing.T) {
|
||||||
|
asDaemon(t, root)
|
||||||
|
injected := metadata.MD{}
|
||||||
|
injected.Append(mdFwdGroup, sidAdministrators)
|
||||||
|
|
||||||
|
ctx := WithForwardedIdentity(metadata.NewOutgoingContext(context.Background(), injected),
|
||||||
|
Identity{SID: "S-1-5-21-1-2-3-1001"}, true)
|
||||||
|
out, ok := metadata.FromOutgoingContext(ctx)
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("no outgoing metadata")
|
||||||
|
}
|
||||||
|
if groups := out.Get(mdFwdGroup); len(groups) != 0 {
|
||||||
|
t.Fatalf("injected group survived: %v", groups)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsReservedForwardKey(t *testing.T) {
|
||||||
|
for _, key := range forwardKeys {
|
||||||
|
if !IsReservedForwardKey(key) {
|
||||||
|
t.Errorf("%q must be reserved", key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// grpc-gateway canonicalises header names, so the check has to be
|
||||||
|
// case-insensitive.
|
||||||
|
if !IsReservedForwardKey("X-Netbird-Fwd-Uid") {
|
||||||
|
t.Error("the check must be case-insensitive")
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, key := range []string{"authorization", "x-netbird", "x-netbird-fwd-uid-extra", ""} {
|
||||||
|
if IsReservedForwardKey(key) {
|
||||||
|
t.Errorf("%q must not be reserved", key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestForwardIdentityMetadata_AlwaysMarksForwarded(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
id Identity
|
||||||
|
known bool
|
||||||
|
}{
|
||||||
|
{"known unix identity", unprivUser, true},
|
||||||
|
{"unknown identity", Identity{}, false},
|
||||||
|
{"windows identity", Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}, true},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
md := ForwardIdentityMetadata(tc.id, tc.known)
|
||||||
|
if got := md.Get(mdFwd); len(got) != 1 || got[0] != "1" {
|
||||||
|
t.Fatalf("marker = %v, want exactly one \"1\"", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
127
client/internal/ipcauth/identity.go
Normal file
127
client/internal/ipcauth/identity.go
Normal file
@@ -0,0 +1,127 @@
|
|||||||
|
// Package ipcauth provides the kernel-authenticated identity of a local IPC
|
||||||
|
// (gRPC) caller and the transport credentials that surface it into the gRPC
|
||||||
|
// context, so the daemon can authorize individual RPCs by caller identity.
|
||||||
|
//
|
||||||
|
// On Unix the identity is read from the kernel via SO_PEERCRED (Linux) or
|
||||||
|
// LOCAL_PEERCRED (Darwin/FreeBSD). On Windows it is derived from the
|
||||||
|
// named-pipe client token. Platforms without a peer-identity primitive get no
|
||||||
|
// credentials, and every consumer must fail closed when no identity is
|
||||||
|
// available.
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"slices"
|
||||||
|
|
||||||
|
"google.golang.org/grpc/credentials"
|
||||||
|
"google.golang.org/grpc/peer"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Well-known Windows SIDs that identify a fully privileged principal.
|
||||||
|
const (
|
||||||
|
sidLocalSystem = "S-1-5-18" // NT AUTHORITY\SYSTEM
|
||||||
|
sidLocalService = "S-1-5-19" // NT AUTHORITY\LOCAL SERVICE
|
||||||
|
sidNetworkService = "S-1-5-20" // NT AUTHORITY\NETWORK SERVICE
|
||||||
|
sidAdministrators = "S-1-5-32-544" // BUILTIN\Administrators
|
||||||
|
)
|
||||||
|
|
||||||
|
// Identity is the kernel-authenticated identity of a local IPC caller. The
|
||||||
|
// zero value is not a valid identity: consumers must only use one obtained
|
||||||
|
// with a true ok/nil error return.
|
||||||
|
type Identity struct {
|
||||||
|
// UID and GID are the caller's Unix user ID and primary group ID. Both are
|
||||||
|
// zero on Windows, where SID is authoritative instead.
|
||||||
|
UID uint32
|
||||||
|
GID uint32
|
||||||
|
|
||||||
|
// SID is the caller's Windows security identifier, empty on Unix.
|
||||||
|
SID string
|
||||||
|
|
||||||
|
// Groups holds the caller's Windows group SIDs, captured from the client
|
||||||
|
// token at handshake time. Only groups that are enabled and not
|
||||||
|
// deny-only are captured, so a group listed here is one the caller can
|
||||||
|
// actually exercise. Empty on Unix.
|
||||||
|
Groups []string
|
||||||
|
|
||||||
|
// Elevated reports whether the Windows client token is elevated (running
|
||||||
|
// as administrator, or an administrator with UAC turned off). Always false
|
||||||
|
// on Unix, where privilege is uid 0.
|
||||||
|
Elevated bool
|
||||||
|
|
||||||
|
// PID is the caller's process ID where the platform reports it (Linux's
|
||||||
|
// SO_PEERCRED), and 0 where it does not. It identifies the daemon's own
|
||||||
|
// process dialling itself, which is what the JSON gateway does, and is never
|
||||||
|
// used to grant anything.
|
||||||
|
PID int32
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsWindows reports whether this identity is a Windows principal (SID-based)
|
||||||
|
// rather than a Unix uid/gid principal.
|
||||||
|
func (i Identity) IsWindows() bool {
|
||||||
|
return i.SID != ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsPrivileged reports whether the caller is the platform's administrative
|
||||||
|
// principal, which is what the daemon requires for changes that cross the
|
||||||
|
// user-to-root boundary.
|
||||||
|
//
|
||||||
|
// On Windows the decision comes from the caller's token rather than from
|
||||||
|
// account names or group RIDs: an elevated token, one of the service accounts
|
||||||
|
// the daemon itself may run as, or a token with BUILTIN\Administrators
|
||||||
|
// enabled. A UAC-filtered administrator has that group marked deny-only, and
|
||||||
|
// deny-only groups are dropped when the identity is captured, so such a
|
||||||
|
// caller is correctly reported as unprivileged. Domain group memberships
|
||||||
|
// (Domain Admins and friends) are deliberately not consulted: they say
|
||||||
|
// nothing about what this token may do on this machine.
|
||||||
|
func (i Identity) IsPrivileged() bool {
|
||||||
|
if !i.IsWindows() {
|
||||||
|
return i.UID == 0
|
||||||
|
}
|
||||||
|
|
||||||
|
if i.Elevated {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
switch i.SID {
|
||||||
|
case sidLocalSystem, sidLocalService, sidNetworkService:
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
return slices.Contains(i.Groups, sidAdministrators)
|
||||||
|
}
|
||||||
|
|
||||||
|
// String renders the identity for audit logs and denial messages.
|
||||||
|
func (i Identity) String() string {
|
||||||
|
if i.IsWindows() {
|
||||||
|
return fmt.Sprintf("sid=%s elevated=%t", i.SID, i.Elevated)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("uid=%d gid=%d", i.UID, i.GID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AuthInfo carries the peer Identity as a gRPC credentials.AuthInfo so
|
||||||
|
// handlers can retrieve it from the request context via IdentityFromContext.
|
||||||
|
type AuthInfo struct {
|
||||||
|
credentials.CommonAuthInfo
|
||||||
|
Identity Identity
|
||||||
|
}
|
||||||
|
|
||||||
|
// AuthType identifies the authentication scheme.
|
||||||
|
func (AuthInfo) AuthType() string { return "netbird-ipc-peercred" }
|
||||||
|
|
||||||
|
// IdentityFromContext extracts the caller's kernel-authenticated identity from
|
||||||
|
// the gRPC peer context. The second return value is false when no IPC
|
||||||
|
// transport credentials were negotiated, which happens on a TCP daemon socket
|
||||||
|
// and on platforms without a peer-identity primitive. Callers MUST fail closed
|
||||||
|
// in that case.
|
||||||
|
func IdentityFromContext(ctx context.Context) (Identity, bool) {
|
||||||
|
p, ok := peer.FromContext(ctx)
|
||||||
|
if !ok {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
info, ok := p.AuthInfo.(AuthInfo)
|
||||||
|
if !ok {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
return info.Identity, true
|
||||||
|
}
|
||||||
43
client/internal/ipcauth/peercred_bsd.go
Normal file
43
client/internal/ipcauth/peercred_bsd.go
Normal file
@@ -0,0 +1,43 @@
|
|||||||
|
//go:build darwin || freebsd
|
||||||
|
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PeerIdentity reads the kernel-authenticated identity of the process on the
|
||||||
|
// other end of a Unix socket via LOCAL_PEERCRED. The xucred is recorded by the
|
||||||
|
// kernel at connect() time and carries the peer's uid and its group list, of
|
||||||
|
// which the first entry is the primary group.
|
||||||
|
func PeerIdentity(conn net.Conn) (Identity, error) {
|
||||||
|
uc, ok := conn.(*net.UnixConn)
|
||||||
|
if !ok {
|
||||||
|
return Identity{}, fmt.Errorf("connection is not a unix socket: %T", conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
raw, err := uc.SyscallConn()
|
||||||
|
if err != nil {
|
||||||
|
return Identity{}, fmt.Errorf("raw conn: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var cred *unix.Xucred
|
||||||
|
var credErr error
|
||||||
|
if err := raw.Control(func(fd uintptr) {
|
||||||
|
cred, credErr = unix.GetsockoptXucred(int(fd), unix.SOL_LOCAL, unix.LOCAL_PEERCRED)
|
||||||
|
}); err != nil {
|
||||||
|
return Identity{}, fmt.Errorf("control raw conn: %w", err)
|
||||||
|
}
|
||||||
|
if credErr != nil {
|
||||||
|
return Identity{}, fmt.Errorf("read LOCAL_PEERCRED: %w", credErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
id := Identity{UID: cred.Uid}
|
||||||
|
if cred.Ngroups > 0 {
|
||||||
|
id.GID = cred.Groups[0]
|
||||||
|
}
|
||||||
|
return id, nil
|
||||||
|
}
|
||||||
39
client/internal/ipcauth/peercred_linux.go
Normal file
39
client/internal/ipcauth/peercred_linux.go
Normal file
@@ -0,0 +1,39 @@
|
|||||||
|
//go:build linux
|
||||||
|
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PeerIdentity reads the kernel-authenticated identity of the process on the
|
||||||
|
// other end of a Unix socket via SO_PEERCRED. The credentials are recorded by
|
||||||
|
// the kernel at connect() time and cannot be changed for the life of the
|
||||||
|
// connection, so they are not spoofable by the caller.
|
||||||
|
func PeerIdentity(conn net.Conn) (Identity, error) {
|
||||||
|
uc, ok := conn.(*net.UnixConn)
|
||||||
|
if !ok {
|
||||||
|
return Identity{}, fmt.Errorf("connection is not a unix socket: %T", conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
raw, err := uc.SyscallConn()
|
||||||
|
if err != nil {
|
||||||
|
return Identity{}, fmt.Errorf("raw conn: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var cred *unix.Ucred
|
||||||
|
var credErr error
|
||||||
|
if err := raw.Control(func(fd uintptr) {
|
||||||
|
cred, credErr = unix.GetsockoptUcred(int(fd), unix.SOL_SOCKET, unix.SO_PEERCRED)
|
||||||
|
}); err != nil {
|
||||||
|
return Identity{}, fmt.Errorf("control raw conn: %w", err)
|
||||||
|
}
|
||||||
|
if credErr != nil {
|
||||||
|
return Identity{}, fmt.Errorf("read SO_PEERCRED: %w", credErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
return Identity{UID: cred.Uid, GID: cred.Gid, PID: cred.Pid}, nil
|
||||||
|
}
|
||||||
87
client/internal/ipcauth/pipeserver_windows.go
Normal file
87
client/internal/ipcauth/pipeserver_windows.go
Normal file
@@ -0,0 +1,87 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
)
|
||||||
|
|
||||||
|
// PipeServerTrusted reports an error unless the pipe behind conn was created by a
|
||||||
|
// principal this client may hand secrets to. Clients call it for a pipe whose name
|
||||||
|
// carries no guarantee of its own, which is any name outside the
|
||||||
|
// ProtectedPrefix\Administrators namespace: that namespace already restricts
|
||||||
|
// creation to administrators and LocalSystem, while a plain name can be created by
|
||||||
|
// any local user before the daemon gets there.
|
||||||
|
//
|
||||||
|
// The decision is made from the pipe object's owner, not from the serving process,
|
||||||
|
// because a client cannot open a process running as another user at all, and the
|
||||||
|
// legitimate case is precisely an unprivileged client talking to a privileged
|
||||||
|
// daemon. Trusted owners are the service accounts, BUILTIN\Administrators, and
|
||||||
|
// this client's own user, the last of which is the daemon a user runs themselves
|
||||||
|
// as in netstack mode. A pipe owned by anyone else gets no setup key, pre-shared
|
||||||
|
// key or SSO prompt out of this client.
|
||||||
|
func PipeServerTrusted(conn net.Conn) error {
|
||||||
|
// go-winio's pipe connection embeds *win32File, which exposes Fd().
|
||||||
|
fdConn, ok := conn.(interface{ Fd() uintptr })
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("connection %T does not expose a pipe handle", conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
owner, err := pipeOwnerSID(windows.Handle(fdConn.Fd()))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if !trustedPipeOwner(owner) {
|
||||||
|
return fmt.Errorf("pipe owned by %s, which is neither an administrator nor this user", owner)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// PipeOwnedBySelf reports whether the pipe behind conn was created by this very
|
||||||
|
// user, which is how a client recognises a daemon running as itself. Ownership it
|
||||||
|
// cannot read is reported as false.
|
||||||
|
func PipeOwnedBySelf(conn net.Conn) bool {
|
||||||
|
fdConn, ok := conn.(interface{ Fd() uintptr })
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
owner, err := pipeOwnerSID(windows.Handle(fdConn.Fd()))
|
||||||
|
if err != nil {
|
||||||
|
log.Debugf("read daemon pipe owner: %v", err)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return selfKnown && selfIdentity.SID != "" && owner == selfIdentity.SID
|
||||||
|
}
|
||||||
|
|
||||||
|
// pipeOwnerSID reads the owner of the pipe object a client is connected to. The
|
||||||
|
// handle was opened with GENERIC_READ, which includes READ_CONTROL, so no extra
|
||||||
|
// access is needed.
|
||||||
|
func pipeOwnerSID(handle windows.Handle) (string, error) {
|
||||||
|
sd, err := windows.GetSecurityInfo(handle, windows.SE_KERNEL_OBJECT, windows.OWNER_SECURITY_INFORMATION)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("read pipe security info: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
owner, _, err := sd.Owner()
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("read pipe owner: %w", err)
|
||||||
|
}
|
||||||
|
return owner.String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// trustedPipeOwner reports whether a pipe's owner is a principal a client may
|
||||||
|
// speak to. An elevated process's objects are owned by BUILTIN\Administrators by
|
||||||
|
// default, an unelevated one's by the user, which is why both forms appear here.
|
||||||
|
func trustedPipeOwner(owner string) bool {
|
||||||
|
switch owner {
|
||||||
|
case sidLocalSystem, sidLocalService, sidNetworkService, sidAdministrators:
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return selfKnown && selfIdentity.SID != "" && owner == selfIdentity.SID
|
||||||
|
}
|
||||||
125
client/internal/ipcauth/privileged.go
Normal file
125
client/internal/ipcauth/privileged.go
Normal file
@@ -0,0 +1,125 @@
|
|||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"runtime"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Fields of the ErrorInfo detail the daemon attaches to a PermissionDenied it
|
||||||
|
// raises for an operation that requires root/administrator. Clients match on
|
||||||
|
// Reason and Domain rather than on the message text, and render the summary and
|
||||||
|
// command themselves so the user gets guidance instead of a gRPC error dump.
|
||||||
|
const (
|
||||||
|
// ErrorReasonPrivilegeRequired identifies the detail.
|
||||||
|
ErrorReasonPrivilegeRequired = "PRIVILEGE_REQUIRED"
|
||||||
|
// ErrorDomain scopes the reason to the NetBird daemon.
|
||||||
|
ErrorDomain = "daemon.netbird.io"
|
||||||
|
// ErrorMetaSummary is the one-sentence explanation of what was refused.
|
||||||
|
ErrorMetaSummary = "summary"
|
||||||
|
// ErrorMetaCommand is the command that performs the same operation with the
|
||||||
|
// privileges it needs, ready to copy and run.
|
||||||
|
ErrorMetaCommand = "command"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The identity of the process evaluating callers, captured once because it cannot
|
||||||
|
// change. selfKnown is false when it could not be read, in which case nothing is
|
||||||
|
// ever treated as this process. selfMayDelegate additionally requires this
|
||||||
|
// process to be unprivileged: see IsPrivilegedCaller.
|
||||||
|
var (
|
||||||
|
selfIdentity Identity
|
||||||
|
selfKnown bool
|
||||||
|
selfMayDelegate bool
|
||||||
|
// selfPID is this process's PID, used to recognise the daemon dialling itself.
|
||||||
|
selfPID = os.Getpid()
|
||||||
|
)
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
id, err := CurrentProcessIdentity()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
selfIdentity, selfKnown = id, true
|
||||||
|
// Only an unprivileged daemon delegates its authority to its own identity.
|
||||||
|
// When it is root or LocalSystem, sharing its identity does not mean sharing
|
||||||
|
// its power: on Windows a filtered and a full token carry the same SID, so
|
||||||
|
// matching there would let a non-elevated shell of an administrator account
|
||||||
|
// act as an administrator, which is the boundary the token check exists to
|
||||||
|
// keep.
|
||||||
|
selfMayDelegate = !id.IsPrivileged()
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsDaemonSelf reports whether an identity is this very process. The JSON gateway
|
||||||
|
// runs inside the daemon and re-dials it locally, so this is what distinguishes
|
||||||
|
// the gateway from any other caller, whatever user the daemon runs as.
|
||||||
|
func IsDaemonSelf(id Identity) bool {
|
||||||
|
if !selfKnown || id.IsWindows() != selfIdentity.IsWindows() {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if id.IsWindows() {
|
||||||
|
return id.SID != "" && id.SID == selfIdentity.SID
|
||||||
|
}
|
||||||
|
return id.UID == selfIdentity.UID
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsPrivilegedCaller reports whether an identity may make the changes the daemon
|
||||||
|
// restricts to the platform administrator. This is the daemon's own rule and
|
||||||
|
// cannot be evaluated by a client, which does not know what the daemon runs as.
|
||||||
|
//
|
||||||
|
// Beyond root/administrator it accepts a caller running as the daemon's own
|
||||||
|
// identity when the daemon is itself unprivileged. That keeps a rootless container
|
||||||
|
// working, where there is no uid 0 at all, and a Windows daemon in netstack mode,
|
||||||
|
// which needs no administrator rights. In those setups a caller sharing the
|
||||||
|
// daemon's identity can already rewrite the config files it reads and replace the
|
||||||
|
// binary it runs, so refusing it a config change would protect nothing; and an
|
||||||
|
// unprivileged daemon cannot hand out a root shell in the first place.
|
||||||
|
func IsPrivilegedCaller(id Identity) bool {
|
||||||
|
if id.IsPrivileged() {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return selfMayDelegate && IsDaemonSelf(id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SelfDelegatesTo returns the identity this process delegates its authority to,
|
||||||
|
// and whether it delegates at all. Only an unprivileged daemon does: see
|
||||||
|
// IsPrivilegedCaller. It exists so a refusal can name who may actually perform the
|
||||||
|
// operation, because on such a host root is neither required nor necessarily
|
||||||
|
// available.
|
||||||
|
func SelfDelegatesTo() (Identity, bool) {
|
||||||
|
if !selfKnown || !selfMayDelegate {
|
||||||
|
return Identity{}, false
|
||||||
|
}
|
||||||
|
return selfIdentity, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// PrivilegedActor names the principal a privileged operation requires, for use
|
||||||
|
// in messages shown to the user.
|
||||||
|
func PrivilegedActor() string {
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
return "administrator privileges"
|
||||||
|
}
|
||||||
|
return "root"
|
||||||
|
}
|
||||||
|
|
||||||
|
// ElevatedCommand renders a command so that running it grants the privileges the
|
||||||
|
// operation needs. Windows has no in-line equivalent of sudo, so the command is
|
||||||
|
// returned unchanged and the user is expected to run it from an elevated
|
||||||
|
// terminal.
|
||||||
|
func ElevatedCommand(command string) string {
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
return command
|
||||||
|
}
|
||||||
|
return "sudo " + command
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpCommand renders an elevated `netbird up` with the given flags, preceded by a
|
||||||
|
// `down`. The down is what makes the command work on a connected client: `netbird
|
||||||
|
// up` prints "Already connected" and returns without applying any config flag, so
|
||||||
|
// on its own the command would appear to do nothing. It is a no-op, exit 0, when
|
||||||
|
// the client is not connected.
|
||||||
|
//
|
||||||
|
// ";" rather than "&&" so the line can be pasted into any of the shells a user
|
||||||
|
// might have: PowerShell 5.1, still the default on Windows Server, rejects "&&"
|
||||||
|
// as a syntax error.
|
||||||
|
func UpCommand(flags string) string {
|
||||||
|
return ElevatedCommand("netbird down") + "; " + ElevatedCommand("netbird up "+flags)
|
||||||
|
}
|
||||||
134
client/internal/ipcauth/privileged_test.go
Normal file
134
client/internal/ipcauth/privileged_test.go
Normal file
@@ -0,0 +1,134 @@
|
|||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
// The self rule is the one place privilege is granted to something other than the
|
||||||
|
// platform administrator, so its two guards matter: it must apply only when the
|
||||||
|
// daemon is itself unprivileged, and only to a caller with the daemon's identity.
|
||||||
|
func TestIsPrivilegedCaller_SelfRule(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
// self stands in for the process the daemon runs as.
|
||||||
|
self Identity
|
||||||
|
selfKnown bool
|
||||||
|
caller Identity
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "root is privileged whatever the daemon runs as",
|
||||||
|
self: Identity{UID: 1000},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{UID: 0},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "an unprivileged daemon delegates to its own user (rootless container)",
|
||||||
|
self: Identity{UID: 1000},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{UID: 1000},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "an unprivileged daemon delegates to nobody else",
|
||||||
|
self: Identity{UID: 1000},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{UID: 1001},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// The daemon is root on a normal install, so sharing its identity is
|
||||||
|
// already covered by being root; nothing else may match.
|
||||||
|
name: "a root daemon delegates to nobody",
|
||||||
|
self: Identity{UID: 0},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{UID: 1000},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// Windows netstack mode: the daemon needs no administrator rights.
|
||||||
|
name: "an unprivileged windows daemon delegates to its own SID",
|
||||||
|
self: Identity{SID: "S-1-5-21-1-2-3-1001"},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{SID: "S-1-5-21-1-2-3-1001"},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "an unprivileged windows daemon delegates to no other SID",
|
||||||
|
self: Identity{SID: "S-1-5-21-1-2-3-1001"},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{SID: "S-1-5-21-1-2-3-1002"},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// The UAC boundary: a filtered and a full token of the same account
|
||||||
|
// carry the same SID but not the same power, so an elevated daemon must
|
||||||
|
// never delegate to its own SID.
|
||||||
|
name: "an elevated windows daemon does not delegate to its own SID",
|
||||||
|
self: Identity{SID: "S-1-5-21-1-2-3-500", Elevated: true},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{SID: "S-1-5-21-1-2-3-500"},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "LocalSystem is privileged on its own merits, not by delegation",
|
||||||
|
self: Identity{SID: sidLocalSystem},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{SID: sidLocalSystem},
|
||||||
|
want: true, // LocalSystem is privileged on its own merits
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "identities of different kinds never match",
|
||||||
|
self: Identity{UID: 1000},
|
||||||
|
selfKnown: true,
|
||||||
|
caller: Identity{SID: "S-1-5-21-1-2-3-1001"},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "an unknown self identity delegates to nobody",
|
||||||
|
self: Identity{},
|
||||||
|
selfKnown: false,
|
||||||
|
caller: Identity{UID: 1000},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
prevID, prevKnown, prevDelegate := selfIdentity, selfKnown, selfMayDelegate
|
||||||
|
t.Cleanup(func() { selfIdentity, selfKnown, selfMayDelegate = prevID, prevKnown, prevDelegate })
|
||||||
|
|
||||||
|
selfIdentity, selfKnown = tt.self, tt.selfKnown
|
||||||
|
selfMayDelegate = tt.selfKnown && !tt.self.IsPrivileged()
|
||||||
|
|
||||||
|
if got := IsPrivilegedCaller(tt.caller); got != tt.want {
|
||||||
|
t.Fatalf("IsPrivilegedCaller(%v) with daemon %v = %t, want %t",
|
||||||
|
tt.caller, tt.self, got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The real process must never accidentally delegate: a test binary running as a
|
||||||
|
// normal user is unprivileged, so it may match itself, but nothing else.
|
||||||
|
func TestIsPrivilegedCaller_ThisProcess(t *testing.T) {
|
||||||
|
id, err := CurrentProcessIdentity()
|
||||||
|
if err != nil {
|
||||||
|
t.Skipf("cannot read this process's identity: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// This process is always allowed to act as itself: either it is privileged, or
|
||||||
|
// it is unprivileged and therefore delegates to its own identity.
|
||||||
|
if !IsPrivilegedCaller(id) {
|
||||||
|
t.Errorf("this process %v was refused its own identity", id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A caller that is neither root nor this process must be refused, whatever
|
||||||
|
// this process happens to be.
|
||||||
|
other := Identity{UID: id.UID + 1}
|
||||||
|
if id.IsWindows() {
|
||||||
|
other = Identity{SID: id.SID + "9"}
|
||||||
|
}
|
||||||
|
if IsPrivilegedCaller(other) {
|
||||||
|
t.Errorf("an unrelated identity %v was treated as privileged", other)
|
||||||
|
}
|
||||||
|
}
|
||||||
17
client/internal/ipcauth/self_unix.go
Normal file
17
client/internal/ipcauth/self_unix.go
Normal file
@@ -0,0 +1,17 @@
|
|||||||
|
//go:build !windows
|
||||||
|
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import "os"
|
||||||
|
|
||||||
|
// CurrentProcessIdentity returns this process's identity as the daemon would
|
||||||
|
// see it if this process connected to the local IPC. It lets a client (the UI)
|
||||||
|
// decide up front whether a privileged operation can succeed, without a
|
||||||
|
// round-trip and without duplicating the rules: the answer comes from the same
|
||||||
|
// Identity.IsPrivileged the daemon applies.
|
||||||
|
func CurrentProcessIdentity() (Identity, error) {
|
||||||
|
return Identity{
|
||||||
|
UID: uint32(os.Geteuid()),
|
||||||
|
GID: uint32(os.Getegid()),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
35
client/internal/ipcauth/self_windows.go
Normal file
35
client/internal/ipcauth/self_windows.go
Normal file
@@ -0,0 +1,35 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package ipcauth
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CurrentProcessIdentity returns this process's identity as the daemon would see
|
||||||
|
// it if this process connected to the local IPC. It lets a client (the UI)
|
||||||
|
// decide up front whether a privileged operation can succeed, without a
|
||||||
|
// round-trip and without duplicating the rules: the answer comes from the same
|
||||||
|
// Identity.IsPrivileged the daemon applies to the token it reads off the pipe.
|
||||||
|
func CurrentProcessIdentity() (Identity, error) {
|
||||||
|
// A pseudo-token, so it must not be closed.
|
||||||
|
token := windows.GetCurrentProcessToken()
|
||||||
|
|
||||||
|
user, err := token.GetTokenUser()
|
||||||
|
if err != nil {
|
||||||
|
return Identity{}, fmt.Errorf("read token user: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
groups, err := tokenGroupSIDs(token)
|
||||||
|
if err != nil {
|
||||||
|
return Identity{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return Identity{
|
||||||
|
SID: user.User.Sid.String(),
|
||||||
|
Groups: groups,
|
||||||
|
Elevated: token.IsElevated(),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
@@ -29,6 +29,11 @@ type managedPeer struct {
|
|||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
InactivityThreshold *time.Duration
|
InactivityThreshold *time.Duration
|
||||||
|
// ReconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is
|
||||||
|
// armed. The activity listener creates the wake peer with the overlay /32 only; without the
|
||||||
|
// routed prefixes WireGuard would not steer subnet-bound traffic to the wake endpoint, so an
|
||||||
|
// idle routing peer could never be woken by that traffic. Optional; nil disables the reconcile.
|
||||||
|
ReconcileAllowedIPs func(peerKey string) error
|
||||||
}
|
}
|
||||||
|
|
||||||
// Manager manages lazy connections
|
// Manager manages lazy connections
|
||||||
@@ -56,6 +61,9 @@ type Manager struct {
|
|||||||
peerToHAGroups map[string][]route.HAUniqueID // peer ID -> HA groups they belong to
|
peerToHAGroups map[string][]route.HAUniqueID // peer ID -> HA groups they belong to
|
||||||
haGroupToPeers map[route.HAUniqueID][]string // HA group -> peer IDs in the group
|
haGroupToPeers map[route.HAUniqueID][]string // HA group -> peer IDs in the group
|
||||||
routesMu sync.RWMutex
|
routesMu sync.RWMutex
|
||||||
|
|
||||||
|
// reconcileAllowedIPs re-applies a peer's routed allowed IPs after its wake endpoint is armed.
|
||||||
|
reconcileAllowedIPs func(peerKey string) error
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewManager creates a new lazy connection manager
|
// NewManager creates a new lazy connection manager
|
||||||
@@ -73,6 +81,7 @@ func NewManager(config Config, engineCtx context.Context, peerStore *peerstore.S
|
|||||||
activityManager: activity.NewManager(wgIface),
|
activityManager: activity.NewManager(wgIface),
|
||||||
peerToHAGroups: make(map[string][]route.HAUniqueID),
|
peerToHAGroups: make(map[string][]route.HAUniqueID),
|
||||||
haGroupToPeers: make(map[route.HAUniqueID][]string),
|
haGroupToPeers: make(map[route.HAUniqueID][]string),
|
||||||
|
reconcileAllowedIPs: config.ReconcileAllowedIPs,
|
||||||
}
|
}
|
||||||
|
|
||||||
if wgIface.IsUserspaceBind() {
|
if wgIface.IsUserspaceBind() {
|
||||||
@@ -201,7 +210,7 @@ func (m *Manager) AddPeer(peerCfg lazyconn.PeerConfig) (bool, error) {
|
|||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil {
|
if err := m.armActivityListener(peerCfg); err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -288,7 +297,7 @@ func (m *Manager) DeactivatePeer(peerID peerid.ConnID) {
|
|||||||
|
|
||||||
m.inactivityManager.RemovePeer(mp.peerCfg.PublicKey)
|
m.inactivityManager.RemovePeer(mp.peerCfg.PublicKey)
|
||||||
|
|
||||||
if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil {
|
if err := m.armActivityListener(*mp.peerCfg); err != nil {
|
||||||
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
|
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -465,6 +474,31 @@ func (m *Manager) close() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// shouldDeferIdleForHA checks if peer should stay connected due to HA group requirements
|
// shouldDeferIdleForHA checks if peer should stay connected due to HA group requirements
|
||||||
|
// armRoutedAllowedIPs re-applies the peer's routed allowed IPs onto its freshly armed wake
|
||||||
|
// endpoint. The activity listener creates the wake peer with the overlay /32 only, so without
|
||||||
|
// this the routed prefixes would be missing and traffic to a routed subnet could not wake the
|
||||||
|
// idle routing peer. It is a no-op when no reconciler is configured.
|
||||||
|
// armActivityListener (re)arms the peer's wake endpoint via the activity manager and then
|
||||||
|
// re-applies its routed allowed IPs, so traffic to a routed subnet can wake an idle routing
|
||||||
|
// peer. The routed prefixes must be re-applied after the wake endpoint exists because the
|
||||||
|
// listener creates it with the overlay /32 only.
|
||||||
|
func (m *Manager) armActivityListener(peerCfg lazyconn.PeerConfig) error {
|
||||||
|
if err := m.activityManager.MonitorPeerActivity(peerCfg); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
m.armRoutedAllowedIPs(&peerCfg)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *Manager) armRoutedAllowedIPs(peerCfg *lazyconn.PeerConfig) {
|
||||||
|
if m.reconcileAllowedIPs == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := m.reconcileAllowedIPs(peerCfg.PublicKey); err != nil {
|
||||||
|
peerCfg.Log.Errorf("failed to reconcile routed allowed IPs on wake endpoint: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (m *Manager) shouldDeferIdleForHA(inactivePeers map[string]struct{}, peerID string) bool {
|
func (m *Manager) shouldDeferIdleForHA(inactivePeers map[string]struct{}, peerID string) bool {
|
||||||
m.routesMu.RLock()
|
m.routesMu.RLock()
|
||||||
defer m.routesMu.RUnlock()
|
defer m.routesMu.RUnlock()
|
||||||
@@ -577,7 +611,7 @@ func (m *Manager) onPeerInactivityTimedOut(peerIDs map[string]struct{}) {
|
|||||||
|
|
||||||
mp.peerCfg.Log.Infof("start activity monitor")
|
mp.peerCfg.Log.Infof("start activity monitor")
|
||||||
|
|
||||||
if err := m.activityManager.MonitorPeerActivity(*mp.peerCfg); err != nil {
|
if err := m.armActivityListener(*mp.peerCfg); err != nil {
|
||||||
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
|
mp.peerCfg.Log.Errorf("failed to create activity monitor: %v", err)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -11,12 +11,14 @@ import (
|
|||||||
|
|
||||||
// MobileDependency collect all dependencies for mobile platform
|
// MobileDependency collect all dependencies for mobile platform
|
||||||
type MobileDependency struct {
|
type MobileDependency struct {
|
||||||
// Android only
|
// Android and iOS
|
||||||
TunAdapter device.TunAdapter
|
|
||||||
IFaceDiscover stdnet.ExternalIFaceDiscover
|
|
||||||
NetworkChangeListener listener.NetworkChangeListener
|
NetworkChangeListener listener.NetworkChangeListener
|
||||||
HostDNSAddresses []netip.AddrPort
|
|
||||||
DnsReadyListener dns.ReadyListener
|
// Android only
|
||||||
|
TunAdapter device.TunAdapter
|
||||||
|
IFaceDiscover stdnet.ExternalIFaceDiscover
|
||||||
|
HostDNSAddresses []netip.AddrPort
|
||||||
|
DnsReadyListener dns.ReadyListener
|
||||||
|
|
||||||
// iOS only
|
// iOS only
|
||||||
DnsManager dns.IosDnsManager
|
DnsManager dns.IosDnsManager
|
||||||
|
|||||||
@@ -175,7 +175,9 @@ func TestFlowAggregationOfUnknownProtocols(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestResetAggregationWindow(t *testing.T) {
|
func TestResetAggregationWindow(t *testing.T) {
|
||||||
store := NewAggregatingMemoryStore()
|
now := time.Now()
|
||||||
|
nowFunc := func() time.Time { return now }
|
||||||
|
store := NewAggregatingMemoryStoreWithTimeFunc(nowFunc)
|
||||||
store.StoreEvent(&types.Event{
|
store.StoreEvent(&types.Event{
|
||||||
ID: uuid.New(),
|
ID: uuid.New(),
|
||||||
Timestamp: time.Now(),
|
Timestamp: time.Now(),
|
||||||
@@ -198,6 +200,7 @@ func TestResetAggregationWindow(t *testing.T) {
|
|||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
|
now = now.Add(1 * time.Second)
|
||||||
reset := store.ResetAggregationWindow()
|
reset := store.ResetAggregationWindow()
|
||||||
previousEvents, ok := reset.(*AggregatingMemory)
|
previousEvents, ok := reset.(*AggregatingMemory)
|
||||||
assert.True(t, ok)
|
assert.True(t, ok)
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ type AggregatingMemory struct {
|
|||||||
WindowStart time.Time
|
WindowStart time.Time
|
||||||
WindowEnd time.Time
|
WindowEnd time.Time
|
||||||
rnd *v2.PCG
|
rnd *v2.PCG
|
||||||
|
nowFunc func() time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *Memory) StoreEvent(event *types.Event) {
|
func (m *Memory) StoreEvent(event *types.Event) {
|
||||||
@@ -62,14 +63,19 @@ func (m *Memory) DeleteEvents(ids []uuid.UUID) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func NewAggregatingMemoryStore() *AggregatingMemory {
|
func NewAggregatingMemoryStore() *AggregatingMemory {
|
||||||
return &AggregatingMemory{WindowStart: time.Now(), Memory: Memory{events: make(map[uuid.UUID]*types.Event)}, rnd: v2.NewPCG(rand.Uint64(), rand.Uint64())}
|
return NewAggregatingMemoryStoreWithTimeFunc(defaultNowFunc)
|
||||||
|
}
|
||||||
|
|
||||||
|
// used in tests when deterministic (less random) time intervals are required
|
||||||
|
func NewAggregatingMemoryStoreWithTimeFunc(nowFunc func() time.Time) *AggregatingMemory {
|
||||||
|
return &AggregatingMemory{WindowStart: nowFunc(), Memory: Memory{events: make(map[uuid.UUID]*types.Event)}, nowFunc: nowFunc, rnd: v2.NewPCG(rand.Uint64(), rand.Uint64())}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (am *AggregatingMemory) ResetAggregationWindow() types.FlowEventAggregator {
|
func (am *AggregatingMemory) ResetAggregationWindow() types.FlowEventAggregator {
|
||||||
am.mux.Lock()
|
am.mux.Lock()
|
||||||
defer am.mux.Unlock()
|
defer am.mux.Unlock()
|
||||||
|
|
||||||
now := time.Now()
|
now := am.nowFunc()
|
||||||
toret := AggregatingMemory{WindowStart: am.WindowStart, WindowEnd: now, Memory: Memory{events: am.events}, rnd: v2.NewPCG(rand.Uint64(), rand.Uint64())}
|
toret := AggregatingMemory{WindowStart: am.WindowStart, WindowEnd: now, Memory: Memory{events: am.events}, rnd: v2.NewPCG(rand.Uint64(), rand.Uint64())}
|
||||||
|
|
||||||
am.events = make(map[uuid.UUID]*types.Event)
|
am.events = make(map[uuid.UUID]*types.Event)
|
||||||
@@ -152,3 +158,7 @@ func (am *AggregatingMemory) GetAggregatedEvents() []*types.Event {
|
|||||||
|
|
||||||
return slices.Collect(maps.Values(aggregated)) // could return an iterator instead here
|
return slices.Collect(maps.Values(aggregated)) // could return an iterator instead here
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func defaultNowFunc() time.Time {
|
||||||
|
return time.Now()
|
||||||
|
}
|
||||||
|
|||||||
@@ -30,6 +30,11 @@ import (
|
|||||||
relayClient "github.com/netbirdio/netbird/shared/relay/client"
|
relayClient "github.com/netbirdio/netbird/shared/relay/client"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// wgTimeoutEscalationThreshold is the number of consecutive WireGuard
|
||||||
|
// handshake timeouts after which the rosenpass state for the peer is
|
||||||
|
// considered desynced and gets reset.
|
||||||
|
const wgTimeoutEscalationThreshold = 3
|
||||||
|
|
||||||
// MetricsRecorder is an interface for recording peer connection metrics
|
// MetricsRecorder is an interface for recording peer connection metrics
|
||||||
type MetricsRecorder interface {
|
type MetricsRecorder interface {
|
||||||
RecordConnectionStages(
|
RecordConnectionStages(
|
||||||
@@ -118,6 +123,9 @@ type Conn struct {
|
|||||||
wgWatcher *WGWatcher
|
wgWatcher *WGWatcher
|
||||||
wgWatcherWg sync.WaitGroup
|
wgWatcherWg sync.WaitGroup
|
||||||
wgWatcherCancel context.CancelFunc
|
wgWatcherCancel context.CancelFunc
|
||||||
|
// wgTimeouts counts consecutive WireGuard handshake timeouts without a
|
||||||
|
// successful handshake in between. Guarded by mu.
|
||||||
|
wgTimeouts int
|
||||||
|
|
||||||
// used to store the remote Rosenpass key for Relayed connection in case of connection update from ice
|
// used to store the remote Rosenpass key for Relayed connection in case of connection update from ice
|
||||||
rosenpassRemoteKey []byte
|
rosenpassRemoteKey []byte
|
||||||
@@ -195,7 +203,6 @@ func NewConn(config ConnConfig, services ServiceDependencies) (*Conn, error) {
|
|||||||
statusICE: worker.NewAtomicStatus(),
|
statusICE: worker.NewAtomicStatus(),
|
||||||
dumpState: dumpState,
|
dumpState: dumpState,
|
||||||
endpointUpdater: NewEndpointUpdater(connLog, config.WgConfig, isController(config)),
|
endpointUpdater: NewEndpointUpdater(connLog, config.WgConfig, isController(config)),
|
||||||
wgWatcher: NewWGWatcher(connLog, config.WgConfig.WgInterface, config.Key, dumpState),
|
|
||||||
metricsRecorder: services.MetricsRecorder,
|
metricsRecorder: services.MetricsRecorder,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -663,11 +670,12 @@ func (conn *Conn) onGuardEvent() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Conn) onWGDisconnected() {
|
func (conn *Conn) onWGDisconnected(watcherCtx context.Context) {
|
||||||
conn.mu.Lock()
|
conn.mu.Lock()
|
||||||
defer conn.mu.Unlock()
|
defer conn.mu.Unlock()
|
||||||
|
|
||||||
if conn.ctx.Err() != nil {
|
// watcherCtx guards against a stale watcher tearing down a connection that already superseded it.
|
||||||
|
if conn.ctx.Err() != nil || watcherCtx.Err() != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -683,6 +691,29 @@ func (conn *Conn) onWGDisconnected() {
|
|||||||
default:
|
default:
|
||||||
conn.Log.Debugf("No active connection to close on WG timeout")
|
conn.Log.Debugf("No active connection to close on WG timeout")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
conn.escalateWGTimeoutLocked()
|
||||||
|
}
|
||||||
|
|
||||||
|
// escalateWGTimeoutLocked resets the peer's rosenpass state after repeated
|
||||||
|
// handshake timeouts. With rosenpass enabled, persistent timeouts mean the
|
||||||
|
// preshared keys have desynced; the renewal exchange runs over the dead
|
||||||
|
// tunnel and cannot resync them. Reporting the peer disconnected drops its
|
||||||
|
// rosenpass state, so the next connection configuration programs the
|
||||||
|
// rendezvous key and the tunnel can bootstrap again. Callers must hold mu.
|
||||||
|
func (conn *Conn) escalateWGTimeoutLocked() {
|
||||||
|
if conn.config.RosenpassConfig.PubKey == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
conn.wgTimeouts++
|
||||||
|
if conn.wgTimeouts < wgTimeoutEscalationThreshold || conn.onDisconnected == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
conn.wgTimeouts = 0
|
||||||
|
|
||||||
|
conn.Log.Warnf("%d consecutive WireGuard handshake timeouts, resetting rosenpass state for peer", wgTimeoutEscalationThreshold)
|
||||||
|
conn.onDisconnected(conn.config.WgConfig.RemoteKey)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Conn) updateRelayStatus(relayServerAddr string, rosenpassPubKey []byte, updateTime time.Time) {
|
func (conn *Conn) updateRelayStatus(relayServerAddr string, rosenpassPubKey []byte, updateTime time.Time) {
|
||||||
@@ -802,25 +833,39 @@ func (conn *Conn) isConnectedOnAllWay() (status guard.ConnStatus) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// enableWgWatcherIfNeeded starts a fresh watcher instance per connection attempt, so its
|
||||||
|
// lifecycle stays bound to conn.mu and enable/disable can't race an old goroutine's shutdown.
|
||||||
|
// Caller must hold conn.mu.
|
||||||
func (conn *Conn) enableWgWatcherIfNeeded(enabledTime time.Time) {
|
func (conn *Conn) enableWgWatcherIfNeeded(enabledTime time.Time) {
|
||||||
if !conn.wgWatcher.PrepareInitialHandshake() {
|
if conn.wgWatcher != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
watcher := NewWGWatcher(conn.Log, conn.config.WgConfig.WgInterface, conn.config.Key, conn.dumpState)
|
||||||
|
watcher.PrepareInitialHandshake()
|
||||||
|
|
||||||
wgWatcherCtx, wgWatcherCancel := context.WithCancel(conn.ctx)
|
wgWatcherCtx, wgWatcherCancel := context.WithCancel(conn.ctx)
|
||||||
|
conn.wgWatcher = watcher
|
||||||
conn.wgWatcherCancel = wgWatcherCancel
|
conn.wgWatcherCancel = wgWatcherCancel
|
||||||
|
|
||||||
conn.wgWatcherWg.Add(1)
|
conn.wgWatcherWg.Add(1)
|
||||||
go func() {
|
go func() {
|
||||||
defer conn.wgWatcherWg.Done()
|
defer conn.wgWatcherWg.Done()
|
||||||
conn.wgWatcher.EnableWgWatcher(wgWatcherCtx, enabledTime, conn.onWGDisconnected, conn.onWGHandshakeSuccess)
|
onDisconnected := func() { conn.onWGDisconnected(wgWatcherCtx) }
|
||||||
|
watcher.EnableWgWatcher(wgWatcherCtx, enabledTime, onDisconnected, conn.onWGHandshakeSuccess, conn.onWGCheckSuccess)
|
||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// disableWgWatcherIfNeeded cancels and drops the watcher once no transport is active. It never
|
||||||
|
// waits for the goroutine: the timeout path reentrantly calls back here under conn.mu, so
|
||||||
|
// blocking would deadlock. Caller must hold conn.mu.
|
||||||
func (conn *Conn) disableWgWatcherIfNeeded() {
|
func (conn *Conn) disableWgWatcherIfNeeded() {
|
||||||
if conn.currentConnPriority == conntype.None && conn.wgWatcherCancel != nil {
|
if conn.currentConnPriority != conntype.None || conn.wgWatcher == nil {
|
||||||
conn.wgWatcherCancel()
|
return
|
||||||
conn.wgWatcherCancel = nil
|
|
||||||
}
|
}
|
||||||
|
conn.wgWatcherCancel()
|
||||||
|
conn.wgWatcher = nil
|
||||||
|
conn.wgWatcherCancel = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Conn) newProxy(remoteConn net.Conn) (wgproxy.Proxy, error) {
|
func (conn *Conn) newProxy(remoteConn net.Conn) (wgproxy.Proxy, error) {
|
||||||
@@ -843,7 +888,9 @@ func (conn *Conn) resetEndpoint() {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
conn.Log.Infof("reset wg endpoint")
|
conn.Log.Infof("reset wg endpoint")
|
||||||
conn.wgWatcher.Reset()
|
if conn.wgWatcher != nil {
|
||||||
|
conn.wgWatcher.Reset()
|
||||||
|
}
|
||||||
if err := conn.endpointUpdater.RemoveEndpointAddress(); err != nil {
|
if err := conn.endpointUpdater.RemoveEndpointAddress(); err != nil {
|
||||||
conn.Log.Warnf("failed to remove endpoint address before update: %v", err)
|
conn.Log.Warnf("failed to remove endpoint address before update: %v", err)
|
||||||
}
|
}
|
||||||
@@ -892,6 +939,15 @@ func (conn *Conn) onWGHandshakeSuccess(when time.Time) {
|
|||||||
conn.recordConnectionMetrics()
|
conn.recordConnectionMetrics()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// onWGCheckSuccess is called for every watcher check that observed a fresh
|
||||||
|
// handshake, including handshakes of connections that were already up when
|
||||||
|
// the watcher started.
|
||||||
|
func (conn *Conn) onWGCheckSuccess() {
|
||||||
|
conn.mu.Lock()
|
||||||
|
conn.wgTimeouts = 0
|
||||||
|
conn.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
// recordConnectionMetrics records connection stage timestamps as metrics
|
// recordConnectionMetrics records connection stage timestamps as metrics
|
||||||
func (conn *Conn) recordConnectionMetrics() {
|
func (conn *Conn) recordConnectionMetrics() {
|
||||||
if conn.metricsRecorder == nil {
|
if conn.metricsRecorder == nil {
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
log "github.com/sirupsen/logrus"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/iface"
|
"github.com/netbirdio/netbird/client/iface"
|
||||||
@@ -304,3 +305,84 @@ func TestConn_presharedKey_RosenpassManaged(t *testing.T) {
|
|||||||
t.Fatalf("expected non-nil presharedKey before Rosenpass manages PSK")
|
t.Fatalf("expected non-nil presharedKey before Rosenpass manages PSK")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func newWGTimeoutTestConn(rosenpassEnabled bool, disconnected *[]string) *Conn {
|
||||||
|
cfg := ConnConfig{
|
||||||
|
Key: "LLHf3Ma6z6mdLbriAJbqhX7+nM/B71lgw2+91q3LfhU=",
|
||||||
|
LocalKey: "RRHf3Ma6z6mdLbriAJbqhX7+nM/B71lgw2+91q3LfhU=",
|
||||||
|
WgConfig: WgConfig{RemoteKey: "LLHf3Ma6z6mdLbriAJbqhX7+nM/B71lgw2+91q3LfhU="},
|
||||||
|
}
|
||||||
|
if rosenpassEnabled {
|
||||||
|
cfg.RosenpassConfig = RosenpassConfig{PubKey: []byte("dummykey")}
|
||||||
|
}
|
||||||
|
|
||||||
|
conn := &Conn{
|
||||||
|
ctx: context.Background(),
|
||||||
|
config: cfg,
|
||||||
|
Log: log.WithField("peer", cfg.Key),
|
||||||
|
metricsStages: &MetricsStages{},
|
||||||
|
}
|
||||||
|
conn.SetOnDisconnected(func(remotePeer string) {
|
||||||
|
*disconnected = append(*disconnected, remotePeer)
|
||||||
|
})
|
||||||
|
return conn
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestConn_onWGDisconnected_EscalatesToRosenpassReset: repeated handshake
|
||||||
|
// timeouts with rosenpass enabled mean the preshared keys have desynced. The
|
||||||
|
// renewal exchange runs over the dead tunnel and cannot resync them, so after
|
||||||
|
// wgTimeoutEscalationThreshold consecutive timeouts the conn must report the
|
||||||
|
// peer disconnected, dropping its rosenpass state so the next configuration
|
||||||
|
// programs the rendezvous key.
|
||||||
|
func TestConn_onWGDisconnected_EscalatesToRosenpassReset(t *testing.T) {
|
||||||
|
var disconnected []string
|
||||||
|
conn := newWGTimeoutTestConn(true, &disconnected)
|
||||||
|
|
||||||
|
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
|
||||||
|
conn.onWGDisconnected(conn.ctx)
|
||||||
|
}
|
||||||
|
assert.Empty(t, disconnected, "escalation must not fire below the threshold")
|
||||||
|
|
||||||
|
conn.onWGDisconnected(conn.ctx)
|
||||||
|
assert.Equal(t, []string{conn.config.WgConfig.RemoteKey}, disconnected,
|
||||||
|
"reaching the threshold must report the peer disconnected once")
|
||||||
|
|
||||||
|
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
|
||||||
|
conn.onWGDisconnected(conn.ctx)
|
||||||
|
}
|
||||||
|
assert.Len(t, disconnected, 1, "escalation must restart counting after firing")
|
||||||
|
|
||||||
|
conn.onWGDisconnected(conn.ctx)
|
||||||
|
assert.Len(t, disconnected, 2, "continued timeouts must escalate again")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestConn_onWGDisconnected_CheckSuccessResetsEscalation: a successful
|
||||||
|
// handshake between timeouts means the tunnel recovered; the counter must
|
||||||
|
// start over.
|
||||||
|
func TestConn_onWGDisconnected_CheckSuccessResetsEscalation(t *testing.T) {
|
||||||
|
var disconnected []string
|
||||||
|
conn := newWGTimeoutTestConn(true, &disconnected)
|
||||||
|
|
||||||
|
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
|
||||||
|
conn.onWGDisconnected(conn.ctx)
|
||||||
|
}
|
||||||
|
conn.onWGCheckSuccess()
|
||||||
|
|
||||||
|
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
|
||||||
|
conn.onWGDisconnected(conn.ctx)
|
||||||
|
}
|
||||||
|
assert.Empty(t, disconnected, "handshake success must reset the timeout count")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestConn_onWGDisconnected_NoEscalationWithoutRosenpass: without rosenpass
|
||||||
|
// there is no per-peer key state to reset; repeated timeouts must not report
|
||||||
|
// disconnects.
|
||||||
|
func TestConn_onWGDisconnected_NoEscalationWithoutRosenpass(t *testing.T) {
|
||||||
|
var disconnected []string
|
||||||
|
conn := newWGTimeoutTestConn(false, &disconnected)
|
||||||
|
|
||||||
|
for i := 0; i < wgTimeoutEscalationThreshold*3; i++ {
|
||||||
|
conn.onWGDisconnected(conn.ctx)
|
||||||
|
}
|
||||||
|
assert.Empty(t, disconnected, "escalation must be limited to rosenpass connections")
|
||||||
|
}
|
||||||
|
|||||||
@@ -813,19 +813,14 @@ func (d *Status) SetSessionExpiresAt(deadline time.Time) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetSessionExpiresAt returns the most recently recorded SSO session deadline,
|
// GetSessionExpiresAt returns the most recently recorded SSO session deadline,
|
||||||
// or the zero value when no deadline is tracked. A deadline that has already
|
// or the zero value when no deadline is tracked. A deadline in the past is
|
||||||
// slipped into the past reports as "none": once the session has expired it is
|
// returned as-is: it means the session has expired, and consumers (tray row,
|
||||||
// no longer a meaningful countdown, and the sessionwatch.Watcher does not
|
// CLI status) render it as "expired" rather than hiding it — masking it as
|
||||||
// arm a timer at the deadline itself to clear it (only the two pre-expiry
|
// "none" would blank the UI at the exact moment it should say the session
|
||||||
// warnings). Without this guard the UI would keep painting a stale
|
// ended.
|
||||||
// "expires in …" against a moment that has passed until the next login,
|
|
||||||
// extend, or teardown rewrote the value.
|
|
||||||
func (d *Status) GetSessionExpiresAt() time.Time {
|
func (d *Status) GetSessionExpiresAt() time.Time {
|
||||||
d.mux.Lock()
|
d.mux.Lock()
|
||||||
defer d.mux.Unlock()
|
defer d.mux.Unlock()
|
||||||
if !d.sessionExpiresAt.IsZero() && d.sessionExpiresAt.Before(time.Now()) {
|
|
||||||
return time.Time{}
|
|
||||||
}
|
|
||||||
return d.sessionExpiresAt
|
return d.sessionExpiresAt
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,6 @@ package peer
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"sync"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
@@ -24,14 +23,14 @@ type WGInterfaceStater interface {
|
|||||||
GetStats() (map[string]configurer.WGStats, error)
|
GetStats() (map[string]configurer.WGStats, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WGWatcher is single-shot: one instance per connection attempt, run once, then discarded.
|
||||||
|
// Lifecycle is owned by Conn under conn.mu, so it keeps no "enabled" state to go stale.
|
||||||
type WGWatcher struct {
|
type WGWatcher struct {
|
||||||
log *log.Entry
|
log *log.Entry
|
||||||
wgIfaceStater WGInterfaceStater
|
wgIfaceStater WGInterfaceStater
|
||||||
peerKey string
|
peerKey string
|
||||||
stateDump *stateDump
|
stateDump *stateDump
|
||||||
|
|
||||||
enabled bool
|
|
||||||
muEnabled sync.Mutex
|
|
||||||
// initialHandshake is not thread-safe; never call PrepareInitialHandshake and EnableWgWatcher concurrently.
|
// initialHandshake is not thread-safe; never call PrepareInitialHandshake and EnableWgWatcher concurrently.
|
||||||
initialHandshake time.Time
|
initialHandshake time.Time
|
||||||
|
|
||||||
@@ -48,36 +47,23 @@ func NewWGWatcher(log *log.Entry, wgIfaceStater WGInterfaceStater, peerKey strin
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// PrepareInitialHandshake reserves the watcher and reads the peer's current WireGuard
|
// PrepareInitialHandshake reads the peer's current WireGuard handshake time. It must be
|
||||||
// handshake time. It must be called before the peer is (re)configured on the WireGuard
|
// called before the peer is (re)configured on the WireGuard interface, so the captured
|
||||||
// interface, so the captured baseline reflects the state prior to this connection attempt
|
// baseline reflects the state prior to this connection attempt instead of racing with
|
||||||
// instead of racing with that configuration. Returns ok=false if the watcher is already
|
// that configuration.
|
||||||
// running, in which case EnableWgWatcher must not be called.
|
func (w *WGWatcher) PrepareInitialHandshake() {
|
||||||
func (w *WGWatcher) PrepareInitialHandshake() (ok bool) {
|
|
||||||
w.muEnabled.Lock()
|
|
||||||
if w.enabled {
|
|
||||||
w.muEnabled.Unlock()
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
w.log.Debugf("enable WireGuard watcher")
|
w.log.Debugf("enable WireGuard watcher")
|
||||||
w.enabled = true
|
|
||||||
w.muEnabled.Unlock()
|
|
||||||
|
|
||||||
handshake, _ := w.wgState()
|
handshake, _ := w.wgState()
|
||||||
w.initialHandshake = handshake
|
w.initialHandshake = handshake
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// EnableWgWatcher runs the WireGuard watcher loop using the handshake baseline captured by
|
// EnableWgWatcher runs the WireGuard watcher loop using the handshake baseline captured by
|
||||||
// PrepareInitialHandshake. The watcher runs until ctx is cancelled. Caller is responsible
|
// PrepareInitialHandshake. The watcher runs until ctx is cancelled. Caller is responsible
|
||||||
// for context lifecycle management.
|
// for context lifecycle management. onHandshakeSuccessFn is called only for the first
|
||||||
func (w *WGWatcher) EnableWgWatcher(ctx context.Context, enabledTime time.Time, onDisconnectedFn func(), onHandshakeSuccessFn func(when time.Time)) {
|
// handshake observed by this run, onCheckSuccessFn for every check that observed a fresh
|
||||||
w.periodicHandshakeCheck(ctx, onDisconnectedFn, onHandshakeSuccessFn, enabledTime, w.initialHandshake)
|
// handshake, including the first.
|
||||||
|
func (w *WGWatcher) EnableWgWatcher(ctx context.Context, enabledTime time.Time, onDisconnectedFn func(), onHandshakeSuccessFn func(when time.Time), onCheckSuccessFn func()) {
|
||||||
w.muEnabled.Lock()
|
w.periodicHandshakeCheck(ctx, onDisconnectedFn, onHandshakeSuccessFn, onCheckSuccessFn, enabledTime, w.initialHandshake)
|
||||||
w.enabled = false
|
|
||||||
w.muEnabled.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reset signals the watcher that the WireGuard peer has been reset and a new
|
// Reset signals the watcher that the WireGuard peer has been reset and a new
|
||||||
@@ -90,7 +76,7 @@ func (w *WGWatcher) Reset() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// wgStateCheck help to check the state of the WireGuard handshake and relay connection
|
// wgStateCheck help to check the state of the WireGuard handshake and relay connection
|
||||||
func (w *WGWatcher) periodicHandshakeCheck(ctx context.Context, onDisconnectedFn func(), onHandshakeSuccessFn func(when time.Time), enabledTime time.Time, initialHandshake time.Time) {
|
func (w *WGWatcher) periodicHandshakeCheck(ctx context.Context, onDisconnectedFn func(), onHandshakeSuccessFn func(when time.Time), onCheckSuccessFn func(), enabledTime time.Time, initialHandshake time.Time) {
|
||||||
w.log.Infof("WireGuard watcher started")
|
w.log.Infof("WireGuard watcher started")
|
||||||
|
|
||||||
timer := time.NewTimer(wgHandshakeOvertime)
|
timer := time.NewTimer(wgHandshakeOvertime)
|
||||||
@@ -103,6 +89,7 @@ func (w *WGWatcher) periodicHandshakeCheck(ctx context.Context, onDisconnectedFn
|
|||||||
case <-timer.C:
|
case <-timer.C:
|
||||||
handshake, ok := w.handshakeCheck(lastHandshake)
|
handshake, ok := w.handshakeCheck(lastHandshake)
|
||||||
if !ok {
|
if !ok {
|
||||||
|
// early ctx cancel check return
|
||||||
if ctx.Err() != nil {
|
if ctx.Err() != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -117,6 +104,10 @@ func (w *WGWatcher) periodicHandshakeCheck(ctx context.Context, onDisconnectedFn
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if onCheckSuccessFn != nil && ctx.Err() == nil {
|
||||||
|
onCheckSuccessFn()
|
||||||
|
}
|
||||||
|
|
||||||
lastHandshake = *handshake
|
lastHandshake = *handshake
|
||||||
|
|
||||||
resetTime := time.Until(handshake.Add(checkPeriod))
|
resetTime := time.Until(handshake.Add(checkPeriod))
|
||||||
@@ -147,9 +138,9 @@ func (w *WGWatcher) handshakeCheck(lastHandshake time.Time) (*time.Time, bool) {
|
|||||||
|
|
||||||
w.log.Tracef("previous handshake, handshake: %v, %v", lastHandshake, handshake)
|
w.log.Tracef("previous handshake, handshake: %v, %v", lastHandshake, handshake)
|
||||||
|
|
||||||
// the current know handshake did not change
|
// the current known handshake did not change
|
||||||
if handshake.Equal(lastHandshake) {
|
if handshake.Equal(lastHandshake) {
|
||||||
w.log.Warnf("WireGuard handshake timed out: %v", handshake)
|
w.log.Warnf("WireGuard handshake not updated: %v", handshake)
|
||||||
return nil, false
|
return nil, false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
"github.com/stretchr/testify/require"
|
|
||||||
|
|
||||||
"github.com/netbirdio/netbird/client/iface/configurer"
|
"github.com/netbirdio/netbird/client/iface/configurer"
|
||||||
)
|
)
|
||||||
@@ -24,6 +23,72 @@ func (m *MocWgIface) disconnect() {
|
|||||||
m.stop = true
|
m.stop = true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type mockHandshakeStats struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
handshake time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockHandshakeStats) GetStats() (map[string]configurer.WGStats, error) {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
return map[string]configurer.WGStats{"": {LastHandshake: m.handshake}}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockHandshakeStats) advance() {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
m.handshake = time.Now()
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWGWatcher_CheckSuccessCallback: onCheckSuccessFn must fire for a fresh
|
||||||
|
// handshake even when the watcher started with an existing handshake baseline,
|
||||||
|
// the case where onHandshakeSuccessFn stays silent.
|
||||||
|
func TestWGWatcher_CheckSuccessCallback(t *testing.T) {
|
||||||
|
// checkPeriod bounds how stale a handshake may be before the watcher treats it
|
||||||
|
// as a suspended-machine timeout. The first check fires after wgHandshakeOvertime,
|
||||||
|
// so keep checkPeriod well above any scheduling jitter to avoid a false timeout
|
||||||
|
// converting the expected success into a disconnect on a loaded runner.
|
||||||
|
checkPeriod = 1 * time.Minute
|
||||||
|
wgHandshakeOvertime = 1 * time.Second
|
||||||
|
|
||||||
|
mlog := log.WithField("peer", "tet")
|
||||||
|
// Use an old baseline so advance() yields a strictly newer handshake even on
|
||||||
|
// platforms with coarse clock resolution (Windows), where two time.Now() calls
|
||||||
|
// microseconds apart can return the same instant and read as a timed-out handshake.
|
||||||
|
stats := &mockHandshakeStats{handshake: time.Now().Add(-time.Hour)}
|
||||||
|
watcher := NewWGWatcher(mlog, stats, "", newStateDump("peer", mlog, &Status{}))
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
watcher.PrepareInitialHandshake()
|
||||||
|
|
||||||
|
firstHandshake := make(chan struct{}, 1)
|
||||||
|
checkSuccess := make(chan struct{}, 1)
|
||||||
|
go watcher.EnableWgWatcher(ctx, time.Now(), func() {}, func(when time.Time) {
|
||||||
|
firstHandshake <- struct{}{}
|
||||||
|
}, func() {
|
||||||
|
select {
|
||||||
|
case checkSuccess <- struct{}{}:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
stats.advance()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-checkSuccess:
|
||||||
|
case <-time.After(10 * time.Second):
|
||||||
|
t.Errorf("timeout waiting for check success callback")
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-firstHandshake:
|
||||||
|
t.Errorf("first-handshake callback must not fire for a non-zero baseline")
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestWGWatcher_EnableWgWatcher(t *testing.T) {
|
func TestWGWatcher_EnableWgWatcher(t *testing.T) {
|
||||||
checkPeriod = 5 * time.Second
|
checkPeriod = 5 * time.Second
|
||||||
wgHandshakeOvertime = 1 * time.Second
|
wgHandshakeOvertime = 1 * time.Second
|
||||||
@@ -35,8 +100,7 @@ func TestWGWatcher_EnableWgWatcher(t *testing.T) {
|
|||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
ok := watcher.PrepareInitialHandshake()
|
watcher.PrepareInitialHandshake()
|
||||||
require.True(t, ok, "watcher should not be enabled yet")
|
|
||||||
|
|
||||||
onDisconnected := make(chan struct{}, 1)
|
onDisconnected := make(chan struct{}, 1)
|
||||||
go watcher.EnableWgWatcher(ctx, time.Now(), func() {
|
go watcher.EnableWgWatcher(ctx, time.Now(), func() {
|
||||||
@@ -44,7 +108,7 @@ func TestWGWatcher_EnableWgWatcher(t *testing.T) {
|
|||||||
onDisconnected <- struct{}{}
|
onDisconnected <- struct{}{}
|
||||||
}, func(when time.Time) {
|
}, func(when time.Time) {
|
||||||
mlog.Infof("onHandshakeSuccess: %v", when)
|
mlog.Infof("onHandshakeSuccess: %v", when)
|
||||||
})
|
}, nil)
|
||||||
|
|
||||||
// wait for initial reading
|
// wait for initial reading
|
||||||
time.Sleep(2 * time.Second)
|
time.Sleep(2 * time.Second)
|
||||||
@@ -66,14 +130,13 @@ func TestWGWatcher_ReEnable(t *testing.T) {
|
|||||||
watcher := NewWGWatcher(mlog, mocWgIface, "", newStateDump("peer", mlog, &Status{}))
|
watcher := NewWGWatcher(mlog, mocWgIface, "", newStateDump("peer", mlog, &Status{}))
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
ok := watcher.PrepareInitialHandshake()
|
watcher.PrepareInitialHandshake()
|
||||||
require.True(t, ok, "watcher should not be enabled yet")
|
|
||||||
|
|
||||||
wg := &sync.WaitGroup{}
|
wg := &sync.WaitGroup{}
|
||||||
wg.Add(1)
|
wg.Add(1)
|
||||||
go func() {
|
go func() {
|
||||||
defer wg.Done()
|
defer wg.Done()
|
||||||
watcher.EnableWgWatcher(ctx, time.Now(), func() {}, func(when time.Time) {})
|
watcher.EnableWgWatcher(ctx, time.Now(), func() {}, func(when time.Time) {}, nil)
|
||||||
}()
|
}()
|
||||||
cancel()
|
cancel()
|
||||||
|
|
||||||
@@ -83,13 +146,12 @@ func TestWGWatcher_ReEnable(t *testing.T) {
|
|||||||
ctx, cancel = context.WithCancel(context.Background())
|
ctx, cancel = context.WithCancel(context.Background())
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
ok = watcher.PrepareInitialHandshake()
|
watcher.PrepareInitialHandshake()
|
||||||
require.True(t, ok, "watcher should be re-enabled after the previous run stopped")
|
|
||||||
|
|
||||||
onDisconnected := make(chan struct{}, 1)
|
onDisconnected := make(chan struct{}, 1)
|
||||||
go watcher.EnableWgWatcher(ctx, time.Now(), func() {
|
go watcher.EnableWgWatcher(ctx, time.Now(), func() {
|
||||||
onDisconnected <- struct{}{}
|
onDisconnected <- struct{}{}
|
||||||
}, func(when time.Time) {})
|
}, func(when time.Time) {}, nil)
|
||||||
|
|
||||||
time.Sleep(2 * time.Second)
|
time.Sleep(2 * time.Second)
|
||||||
mocWgIface.disconnect()
|
mocWgIface.disconnect()
|
||||||
|
|||||||
@@ -96,6 +96,7 @@ type ConfigInput struct {
|
|||||||
BlockLANAccess *bool
|
BlockLANAccess *bool
|
||||||
BlockInbound *bool
|
BlockInbound *bool
|
||||||
DisableIPv6 *bool
|
DisableIPv6 *bool
|
||||||
|
SyncMessageVersion *int
|
||||||
|
|
||||||
DisableNotifications *bool
|
DisableNotifications *bool
|
||||||
|
|
||||||
@@ -137,6 +138,7 @@ type Config struct {
|
|||||||
BlockLANAccess bool
|
BlockLANAccess bool
|
||||||
BlockInbound bool
|
BlockInbound bool
|
||||||
DisableIPv6 bool
|
DisableIPv6 bool
|
||||||
|
SyncMessageVersion *int
|
||||||
|
|
||||||
DisableNotifications *bool
|
DisableNotifications *bool
|
||||||
|
|
||||||
@@ -587,6 +589,12 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
|
|||||||
updated = true
|
updated = true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if input.SyncMessageVersion != nil && *input.SyncMessageVersion != *config.SyncMessageVersion {
|
||||||
|
log.Infof("setting SyncMessageVersion to %v", *input.SyncMessageVersion)
|
||||||
|
*config.SyncMessageVersion = *input.SyncMessageVersion
|
||||||
|
updated = true
|
||||||
|
}
|
||||||
|
|
||||||
if input.DisableNotifications != nil && (config.DisableNotifications == nil || *input.DisableNotifications != *config.DisableNotifications) {
|
if input.DisableNotifications != nil && (config.DisableNotifications == nil || *input.DisableNotifications != *config.DisableNotifications) {
|
||||||
if *input.DisableNotifications {
|
if *input.DisableNotifications {
|
||||||
log.Infof("disabling notifications")
|
log.Infof("disabling notifications")
|
||||||
@@ -738,6 +746,13 @@ func (config *Config) applyMDMPolicy(policy *mdm.Policy) {
|
|||||||
// appended for https or ":80" for http. The serviceName parameter is
|
// appended for https or ":80" for http. The serviceName parameter is
|
||||||
// used to contextualise error messages. On success returns the parsed
|
// used to contextualise error messages. On success returns the parsed
|
||||||
// *url.URL; on failure returns a non-nil error.
|
// *url.URL; on failure returns a non-nil error.
|
||||||
|
// ParseServiceURL normalises a service URL exactly as the config layer does when
|
||||||
|
// it stores one, so callers comparing a requested URL against a stored one do not
|
||||||
|
// have to reimplement the scheme validation and default-port handling.
|
||||||
|
func ParseServiceURL(serviceName, serviceURL string) (*url.URL, error) {
|
||||||
|
return parseURL(serviceName, serviceURL)
|
||||||
|
}
|
||||||
|
|
||||||
func parseURL(serviceName, serviceURL string) (*url.URL, error) {
|
func parseURL(serviceName, serviceURL string) (*url.URL, error) {
|
||||||
parsedMgmtURL, err := url.ParseRequestURI(serviceURL)
|
parsedMgmtURL, err := url.ParseRequestURI(serviceURL)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
"runtime"
|
"runtime"
|
||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
|
||||||
log "github.com/sirupsen/logrus"
|
log "github.com/sirupsen/logrus"
|
||||||
|
|
||||||
@@ -439,7 +440,11 @@ func (s *ServiceManager) GetStatePath() string {
|
|||||||
|
|
||||||
activeProf, err := s.GetActiveProfileState()
|
activeProf, err := s.GetActiveProfileState()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Warnf("failed to get active profile state: %v", err)
|
if errors.Is(err, syscall.ENOSYS) {
|
||||||
|
log.Debugf("active profile state unavailable on this platform: %v", err)
|
||||||
|
} else {
|
||||||
|
log.Warnf("failed to get active profile state: %v", err)
|
||||||
|
}
|
||||||
return defaultStatePath
|
return defaultStatePath
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -39,6 +39,7 @@ type rpServer interface {
|
|||||||
|
|
||||||
type Manager struct {
|
type Manager struct {
|
||||||
ifaceName string
|
ifaceName string
|
||||||
|
localWgKey wgtypes.Key
|
||||||
spk []byte
|
spk []byte
|
||||||
ssk []byte
|
ssk []byte
|
||||||
rpKeyHash string
|
rpKeyHash string
|
||||||
@@ -51,8 +52,9 @@ type Manager struct {
|
|||||||
wgIface PresharedKeySetter
|
wgIface PresharedKeySetter
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewManager creates a new Rosenpass manager
|
// NewManager creates a new Rosenpass manager. localWgKey is the local
|
||||||
func NewManager(preSharedKey *wgtypes.Key, wgIfaceName string) (*Manager, error) {
|
// WireGuard public key, used to derive the per-peer rendezvous key.
|
||||||
|
func NewManager(preSharedKey *wgtypes.Key, wgIfaceName string, localWgKey wgtypes.Key) (*Manager, error) {
|
||||||
public, secret, err := rp.GenerateKeyPair()
|
public, secret, err := rp.GenerateKeyPair()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -62,6 +64,7 @@ func NewManager(preSharedKey *wgtypes.Key, wgIfaceName string) (*Manager, error)
|
|||||||
log.Tracef("generated new rosenpass key pair with public key %s", rpKeyHash)
|
log.Tracef("generated new rosenpass key pair with public key %s", rpKeyHash)
|
||||||
return &Manager{
|
return &Manager{
|
||||||
ifaceName: wgIfaceName,
|
ifaceName: wgIfaceName,
|
||||||
|
localWgKey: localWgKey,
|
||||||
rpKeyHash: rpKeyHash,
|
rpKeyHash: rpKeyHash,
|
||||||
spk: public,
|
spk: public,
|
||||||
ssk: secret,
|
ssk: secret,
|
||||||
@@ -73,7 +76,7 @@ func NewManager(preSharedKey *wgtypes.Key, wgIfaceName string) (*Manager, error)
|
|||||||
// nil receiver in addPeer -> m.rpWgHandler.AddPeer. generateConfig will
|
// nil receiver in addPeer -> m.rpWgHandler.AddPeer. generateConfig will
|
||||||
// replace it with a fresh handler on each Run() to clear stale peer
|
// replace it with a fresh handler on each Run() to clear stale peer
|
||||||
// state from previous engine sessions.
|
// state from previous engine sessions.
|
||||||
rpWgHandler: NewNetbirdHandler(),
|
rpWgHandler: NewNetbirdHandler((*[32]byte)(preSharedKey), localWgKey),
|
||||||
lock: sync.Mutex{},
|
lock: sync.Mutex{},
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
@@ -161,7 +164,7 @@ func (m *Manager) generateConfig() (rp.Config, error) {
|
|||||||
cfg.Peers = []rp.PeerConfig{}
|
cfg.Peers = []rp.PeerConfig{}
|
||||||
|
|
||||||
m.lock.Lock()
|
m.lock.Lock()
|
||||||
m.rpWgHandler = NewNetbirdHandler()
|
m.rpWgHandler = NewNetbirdHandler(m.preSharedKey, m.localWgKey)
|
||||||
if m.wgIface != nil {
|
if m.wgIface != nil {
|
||||||
m.rpWgHandler.SetInterface(m.wgIface)
|
m.rpWgHandler.SetInterface(m.wgIface)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -85,7 +85,7 @@ func newTestManager(spkFirstByte byte, mock *mockServer) *Manager {
|
|||||||
ssk: make([]byte, 32),
|
ssk: make([]byte, 32),
|
||||||
rpKeyHash: "test-hash",
|
rpKeyHash: "test-hash",
|
||||||
rpPeerIDs: make(map[string]*rp.PeerID),
|
rpPeerIDs: make(map[string]*rp.PeerID),
|
||||||
rpWgHandler: NewNetbirdHandler(),
|
rpWgHandler: NewNetbirdHandler(nil, wgtypes.Key{0x01}),
|
||||||
server: mock,
|
server: mock,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -255,7 +255,7 @@ func TestAddPeer_NilServer_ReturnsErrorNoCrash(t *testing.T) {
|
|||||||
// issue #4341 cannot occur in the window between NewManager and Run().
|
// issue #4341 cannot occur in the window between NewManager and Run().
|
||||||
func TestNewManager_PreInitializesHandler(t *testing.T) {
|
func TestNewManager_PreInitializesHandler(t *testing.T) {
|
||||||
psk := wgtypes.Key{}
|
psk := wgtypes.Key{}
|
||||||
m, err := NewManager(&psk, "wt0")
|
m, err := NewManager(&psk, "wt0", wgtypes.Key{0x01})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.NotNil(t, m.rpWgHandler, "rpWgHandler must be initialized in NewManager")
|
require.NotNil(t, m.rpWgHandler, "rpWgHandler must be initialized in NewManager")
|
||||||
}
|
}
|
||||||
@@ -329,10 +329,10 @@ func TestIsPresharedKeyInitialized_AddedButNotHandshaken_ReturnsFalse(t *testing
|
|||||||
require.False(t, m.IsPresharedKeyInitialized(wgKey))
|
require.False(t, m.IsPresharedKeyInitialized(wgKey))
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- NetbirdHandler.outputKey ----------------------------------------------
|
// --- NetbirdHandler.applyKey ----------------------------------------------
|
||||||
|
|
||||||
func TestHandler_OutputKey_FirstCallUsesUpdateOnlyFalse(t *testing.T) {
|
func TestHandler_ApplyKey_FirstCallUsesUpdateOnlyFalse(t *testing.T) {
|
||||||
h := NewNetbirdHandler()
|
h := NewNetbirdHandler(nil, wgtypes.Key{0x01})
|
||||||
iface := &mockIface{}
|
iface := &mockIface{}
|
||||||
h.SetInterface(iface)
|
h.SetInterface(iface)
|
||||||
|
|
||||||
@@ -348,8 +348,8 @@ func TestHandler_OutputKey_FirstCallUsesUpdateOnlyFalse(t *testing.T) {
|
|||||||
require.Equal(t, wgKey.String(), iface.calls[0].peerKey)
|
require.Equal(t, wgKey.String(), iface.calls[0].peerKey)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHandler_OutputKey_SubsequentCallsUseUpdateOnlyTrue(t *testing.T) {
|
func TestHandler_ApplyKey_SubsequentCallsUseUpdateOnlyTrue(t *testing.T) {
|
||||||
h := NewNetbirdHandler()
|
h := NewNetbirdHandler(nil, wgtypes.Key{0x01})
|
||||||
iface := &mockIface{}
|
iface := &mockIface{}
|
||||||
h.SetInterface(iface)
|
h.SetInterface(iface)
|
||||||
|
|
||||||
@@ -364,8 +364,8 @@ func TestHandler_OutputKey_SubsequentCallsUseUpdateOnlyTrue(t *testing.T) {
|
|||||||
require.True(t, iface.calls[1].updateOnly, "subsequent rotations must use updateOnly=true")
|
require.True(t, iface.calls[1].updateOnly, "subsequent rotations must use updateOnly=true")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHandler_OutputKey_NilInterface_NoCrashNoCall(t *testing.T) {
|
func TestHandler_ApplyKey_NilInterface_NoCrashNoCall(t *testing.T) {
|
||||||
h := NewNetbirdHandler()
|
h := NewNetbirdHandler(nil, wgtypes.Key{0x01})
|
||||||
// no SetInterface — iface remains nil
|
// no SetInterface — iface remains nil
|
||||||
pid := rp.PeerID{0x03}
|
pid := rp.PeerID{0x03}
|
||||||
h.AddPeer(pid, "wt0", rp.Key(wgtypes.Key{}))
|
h.AddPeer(pid, "wt0", rp.Key(wgtypes.Key{}))
|
||||||
@@ -374,8 +374,8 @@ func TestHandler_OutputKey_NilInterface_NoCrashNoCall(t *testing.T) {
|
|||||||
h.HandshakeCompleted(pid, rp.Key{})
|
h.HandshakeCompleted(pid, rp.Key{})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHandler_OutputKey_UnknownPeer_NoCall(t *testing.T) {
|
func TestHandler_ApplyKey_UnknownPeer_NoCall(t *testing.T) {
|
||||||
h := NewNetbirdHandler()
|
h := NewNetbirdHandler(nil, wgtypes.Key{0x01})
|
||||||
iface := &mockIface{}
|
iface := &mockIface{}
|
||||||
h.SetInterface(iface)
|
h.SetInterface(iface)
|
||||||
|
|
||||||
@@ -384,7 +384,7 @@ func TestHandler_OutputKey_UnknownPeer_NoCall(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandler_RemovePeer_ClearsInitializedState(t *testing.T) {
|
func TestHandler_RemovePeer_ClearsInitializedState(t *testing.T) {
|
||||||
h := NewNetbirdHandler()
|
h := NewNetbirdHandler(nil, wgtypes.Key{0x01})
|
||||||
iface := &mockIface{}
|
iface := &mockIface{}
|
||||||
h.SetInterface(iface)
|
h.SetInterface(iface)
|
||||||
|
|
||||||
@@ -398,7 +398,7 @@ func TestHandler_RemovePeer_ClearsInitializedState(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandler_SetInterfaceAfterAddPeer_StillReceivesKey(t *testing.T) {
|
func TestHandler_SetInterfaceAfterAddPeer_StillReceivesKey(t *testing.T) {
|
||||||
h := NewNetbirdHandler()
|
h := NewNetbirdHandler(nil, wgtypes.Key{0x01})
|
||||||
pid := rp.PeerID{0x05}
|
pid := rp.PeerID{0x05}
|
||||||
wgKey := wgtypes.Key{0xEE}
|
wgKey := wgtypes.Key{0xEE}
|
||||||
h.AddPeer(pid, "wt0", rp.Key(wgKey))
|
h.AddPeer(pid, "wt0", rp.Key(wgKey))
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user