Merge branch 'main' into loopback-wg-proxy

This commit is contained in:
Viktor Liu
2026-09-30 05:50:22 +09:00
committed by GitHub
223 changed files with 20038 additions and 12222 deletions
+22
View File
@@ -46,3 +46,25 @@ updates:
wireguard:
patterns:
- "golang.zx2c4.com/wireguard*"
# Base images of the source-build Dockerfiles, pinned by digest (Chainguard
# publishes only :latest for free). Dockerfile.release files feed goreleaser
# and keep the published images as they are, so their bases are left alone.
- package-ecosystem: "docker"
directories:
- "/upload-server"
schedule:
interval: "weekly"
open-pull-requests-limit: 3
groups:
base-images:
patterns:
- "*"
ignore:
- dependency-name: "gcr.io/distroless/base"
# Go minor and major versions move with the rest of the repository;
# patch releases and new digests of the pinned tag still come through.
- dependency-name: "golang"
update-types:
- "version-update:semver-minor"
- "version-update:semver-major"
+3 -3
View File
@@ -38,12 +38,12 @@ jobs:
persist-credentials: false
- name: Set up Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v7
with:
node-version: "22"
- name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with:
version: 11
@@ -79,7 +79,7 @@ jobs:
run: echo "path=$(pnpm store path --silent)" >> "$GITHUB_OUTPUT"
- name: Cache pnpm store
uses: actions/cache@v4
uses: actions/cache@v6
with:
path: ${{ steps.pnpm-store.outputs.path }}
key: ${{ runner.os }}-pnpm-${{ hashFiles('client/ui/frontend/pnpm-lock.yaml') }}
+46
View File
@@ -80,3 +80,49 @@ jobs:
skip-save-cache: true
cache-invalidation-interval: 0
args: --timeout=20m
# Separate job rather than extra rows in the matrix above: those rows pick a
# GOOS by picking a runner OS, while android/ios are cross-compiled from
# ubuntu — an `include` entry with os: ubuntu-latest would merge into the
# Linux row instead of adding one. The package path is restricted because a
# whole-repo run under GOOS=android pulls *_linux.go files into packages that
# have no android counterpart.
golangci-mobile:
strategy:
fail-fast: false
matrix:
include:
- goos: android
goarch: arm64
packages: ./client/android/...
display_name: Android
- goos: ios
goarch: arm64
packages: ./client/ios/...
display_name: iOS
name: ${{ matrix.display_name }}
runs-on: ubuntu-latest
timeout-minutes: 25
env:
CGO_ENABLED: 0
GOOS: ${{ matrix.goos }}
GOARCH: ${{ matrix.goarch }}
steps:
- name: Checkout code
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: "go.mod"
cache: false
- name: golangci-lint
uses: golangci/golangci-lint-action@82606bf257cbaff209d206a39f5134f0cfbfd2ee #v9.2.1
with:
version: latest
install-mode: binary
skip-cache: true
skip-save-cache: true
cache-invalidation-interval: 0
args: --timeout=20m ${{ matrix.packages }}
@@ -0,0 +1,64 @@
name: Mobile
on:
push:
branches:
- main
- "release-*"
pull_request:
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}-${{ github.head_ref || github.actor_id }}
cancel-in-progress: true
jobs:
android_build:
name: "Android / Build"
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
goarch: [arm64, arm, amd64, "386"]
env:
CGO_ENABLED: 0
GOOS: android
GOARCH: ${{ matrix.goarch }}
steps:
- name: Checkout repository
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: "go.mod"
- name: Build Android bridge
run: go build ./client/android/...
- name: Vet Android bridge
if: matrix.goarch == 'arm64'
run: go vet ./client/android/...
ios_build:
name: "iOS / Build"
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
goarch: [arm64, amd64]
env:
CGO_ENABLED: 0
GOOS: ios
GOARCH: ${{ matrix.goarch }}
steps:
- name: Checkout repository
uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
with:
persist-credentials: false
- name: Install Go
uses: actions/setup-go@924ae3a1cded613372ab5595356fb5720e22ba16 # v6.5.0
with:
go-version-file: "go.mod"
# No `go vet` counterpart: every ios target requires external (cgo)
# linking, which needs an Xcode toolchain the runner does not have.
- name: Build iOS SDK
run: go build ./client/ios/...
+153 -10
View File
@@ -191,6 +191,17 @@ jobs:
# requires a changelog. Generated, not committed (see .gitignore).
# chglog is a go.mod tool directive, so go.sum pins it and its deps.
run: bash release_files/rpm-changelog.sh
- name: Fill the RPM ISA provide version
# nfpm cannot emit rpmbuild's ISA provide and GoReleaser cannot template it.
run: bash release_files/rpm-provides.sh
- name: Set up Node.js
uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0
with:
node-version: '22'
- name: Install proxy web dependencies for license collection
# proxy/collect-licenses.sh reads the UI's license terms from node_modules.
working-directory: proxy/web
run: npm ci --ignore-scripts
- name: Set up QEMU
uses: docker/setup-qemu-action@06116385d9baf250c9f4dcb4858b16962ea869c3 #v4.1.0
- name: Set up Docker Buildx
@@ -230,14 +241,18 @@ jobs:
uses: goreleaser/goreleaser-action@5daf1e915a5f0af01ddbcd89a43b8061ff4f1a89 # v7.2.2
with:
version: ${{ env.GORELEASER_VER }}
args: release --clean ${{ env.flags }}
args: release --config .goreleaser.generated.yaml --clean ${{ env.flags }}
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
HOMEBREW_TAP_GITHUB_TOKEN: ${{ secrets.HOMEBREW_TAP_GITHUB_TOKEN }}
UPLOAD_DEBIAN_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
UPLOAD_YUM_SECRET: ${{ secrets.PKG_UPLOAD_SECRET }}
GPG_RPM_KEY_FILE: ${{ env.GPG_RPM_KEY_FILE }}
NFPM_NETBIRD_RPM_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
# One per nfpm id: GoReleaser looks the passphrase up as NFPM_<ID>_PASSPHRASE.
NFPM_NETBIRD_RPM_AMD64_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
NFPM_NETBIRD_RPM_ARM64_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
NFPM_NETBIRD_RPM_ARM_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
NFPM_NETBIRD_RPM_386_PASSPHRASE: ${{ secrets.GPG_RPM_PASSPHRASE }}
SKIP_PUBLISH: ${{ env.SKIP_PUBLISH }}
SKIP_DOCKER_PUSH: ${{ env.SKIP_DOCKER_PUSH }}
- name: Verify RPM signatures
@@ -294,10 +309,12 @@ jobs:
tag_and_push() {
local src="$1" img_name tag dst variant=""
img_name="${src%%:*}"
# Client variants share a repository, so keep their tag suffixes.
# Variants share a repository with their default image, so keep
# their tag suffixes. Order matters: the first matching pattern wins.
case "$src" in
*-rootless-ubi-amd64) variant="-rootless-ubi" ;;
*-rootless-amd64) variant="-rootless" ;;
*-ubi-amd64) variant="-ubi" ;;
esac
for tag in $(resolve_tags); do
dst="${img_name}:${tag}${variant}"
@@ -363,6 +380,132 @@ jobs:
path: dist/netbird_darwin**
retention-days: 7
# Certify and publish the rootless UBI client image in the Red Hat Ecosystem
# Catalog. Stable tags only: goreleaser pushes <version>-rootless-ubi to
# ghcr.io in the release job above, and preflight submits every architecture
# of that manifest list to Pyxis. Auto-publish on the component makes the new
# version public once certification passes.
redhat_certification:
name: "Red Hat / Certify rootless UBI image"
needs: release
if: |
github.repository == 'netbirdio/netbird' &&
startsWith(github.ref, 'refs/tags/v') &&
!contains(github.ref_name, '-')
runs-on: ubuntu-24.04
permissions:
contents: read
env:
PREFLIGHT_VERSION: "1.21.0"
# sha256 of preflight-linux-amd64 from the 1.21.0 GitHub release.
# Red Hat publishes no checksum file, so the value is pinned here.
PREFLIGHT_SHA256: "5e653135503c72f8702bbe31d7643197d12937c68086879133dd6b9650a9a449"
IMAGE_REPOSITORY: "ghcr.io/netbirdio/netbird"
# Component "NetBird Client Container Image (rootless)" in Partner Connect.
# Override with the REDHAT_CERT_COMPONENT_ID repository variable if it changes.
DEFAULT_COMPONENT_ID: "6aa3ca4b4676aefdf07aaa97"
steps:
- name: Resolve image reference
id: image
env:
INPUT_VERSION: ${{ github.ref_name }}
run: |
set -euo pipefail
version="${INPUT_VERSION#v}"
if [[ ! "$version" =~ ^[0-9]+\.[0-9]+\.[0-9]+$ ]]; then
echo "::error::Only stable x.y.z versions are certified, got '${INPUT_VERSION}'"
exit 1
fi
echo "version=${version}" >> "$GITHUB_OUTPUT"
echo "ref=${IMAGE_REPOSITORY}:${version}-rootless-ubi" >> "$GITHUB_OUTPUT"
- name: Verify the multi-arch image is on ghcr.io
env:
IMAGE_REF: ${{ steps.image.outputs.ref }}
run: |
set -euo pipefail
docker buildx imagetools inspect "$IMAGE_REF" --raw > manifest.json
for arch in amd64 arm64; do
if ! jq -e --arg a "$arch" '.manifests[] | select(.platform.architecture == $a)' manifest.json > /dev/null; then
echo "::error::${IMAGE_REF} has no ${arch} manifest"
exit 1
fi
done
echo "Manifest list for ${IMAGE_REF}:"
jq -r '.manifests[] | "\(.platform.os)/\(.platform.architecture) \(.digest)"' manifest.json
- name: Install preflight
run: |
set -euo pipefail
curl -fsSL --proto '=https' --proto-redir '=https' -o preflight \
"https://github.com/redhat-openshift-ecosystem/openshift-preflight/releases/download/${PREFLIGHT_VERSION}/preflight-linux-amd64"
echo "${PREFLIGHT_SHA256} preflight" | sha256sum -c -
chmod +x preflight
./preflight --version
- name: Run preflight checks and submit to Red Hat
env:
IMAGE_REF: ${{ steps.image.outputs.ref }}
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
PFLT_CERTIFICATION_COMPONENT_ID: ${{ vars.REDHAT_CERT_COMPONENT_ID || env.DEFAULT_COMPONENT_ID }}
PFLT_ARTIFACTS: artifacts
PFLT_LOGFILE: artifacts/preflight.log
PFLT_LOGLEVEL: info
PFLT_JUNIT: "true"
run: |
set -euo pipefail
# No --platform: preflight walks the manifest list and submits every
# architecture in one run, grouped under one manifest-list digest.
./preflight check container "$IMAGE_REF" --submit
- name: Fail if any check did not pass
run: |
set -euo pipefail
shopt -s nullglob
results=(artifacts/results.json artifacts/*/results.json)
if [[ ${#results[@]} -eq 0 ]]; then
echo "::error::preflight produced no results.json"
exit 1
fi
status=0
for f in "${results[@]}"; do
arch="$(basename "$(dirname "$f")")"
passed="$(jq -r '.passed' "$f")"
failed="$(jq -r '[.results.failed[]?.name] | join(", ")' "$f")"
echo "${arch}: passed=${passed} ${failed:+failed checks: ${failed}}"
[[ "$passed" == "true" ]] || status=1
done
exit $status
- name: Upload preflight artifacts
if: always()
uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
with:
name: redhat-preflight-${{ steps.image.outputs.version }}
path: artifacts/
retention-days: 30
- name: Wait for Pyxis to mark both architectures certified
env:
VERSION: ${{ steps.image.outputs.version }}
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
COMPONENT_ID: ${{ vars.REDHAT_CERT_COMPONENT_ID || env.DEFAULT_COMPONENT_ID }}
run: |
set -euo pipefail
tag="${VERSION}-rootless-ubi"
url="https://catalog.redhat.com/api/containers/v1/projects/certification/id/${COMPONENT_ID}/images?page_size=100"
for attempt in $(seq 1 20); do
certified="$(curl -fsS --proto '=https' --proto-redir '=https' -H "X-API-KEY: ${PFLT_PYXIS_API_TOKEN}" "$url" \
| jq -r --arg t "$tag" '[.data[] | select(.repositories[]?.tags[]?.name == $t) | select(.certified == true) | .architecture] | unique | join(",")')"
echo "attempt ${attempt}: certified architectures for ${tag}: ${certified:-none}"
if [[ "$certified" == "amd64,arm64" ]]; then
echo "Both architectures certified. Auto-publish is enabled on the component, so the catalog updates on its own."
exit 0
fi
sleep 30
done
echo "::warning::Pyxis has not marked both architectures certified after 10 minutes. Check https://connect.redhat.com/component/view/${COMPONENT_ID}/images"
release_ui:
runs-on: ubuntu-latest
outputs:
@@ -417,12 +560,12 @@ jobs:
run: git --no-pager diff --exit-code
- name: Set up Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v7
with:
node-version: '22'
- name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with:
version: 11
@@ -554,12 +697,12 @@ jobs:
run: git --no-pager diff --exit-code
- name: Set up Node.js
uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4
uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0
with:
node-version: '22'
- name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with:
version: 11
@@ -651,11 +794,11 @@ jobs:
- name: check git status
run: git --no-pager diff --exit-code
- name: Set up Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v7
with:
node-version: '22'
- name: Set up pnpm
uses: pnpm/action-setup@a3252b78c470c02df07e9d59298aecedc3ccdd6d # v3.0.0
uses: pnpm/action-setup@ea17c68df8912ef543352723c149a84f56e3d413 # v6.1.0
with:
version: 11
- name: Install wails3 CLI
@@ -774,7 +917,7 @@ jobs:
run: 7z x -o"${{ github.workspace }}/NSIS_Plugins" "${{ github.workspace }}/ShellExecAsUser_amd64-Unicode.7z"
- name: Set up Go for wails3 CLI
uses: actions/setup-go@v5
uses: actions/setup-go@v6
with:
go-version-file: "go.mod"
cache: false
+1 -1
View File
@@ -32,7 +32,7 @@ jobs:
persist-credentials: false
- name: Set up Node.js
uses: actions/setup-node@v4
uses: actions/setup-node@v7
with:
node-version: "22"
+3
View File
@@ -38,4 +38,7 @@ management/server/types/testdata/
# generated by chglog in the release workflow, embedded into the RPM
changelog.yml
# generated by rpm-provides.sh, the config GoReleaser actually runs
.goreleaser.generated.yaml
.chglog.yml
+96 -6
View File
@@ -40,6 +40,32 @@ builds:
tags:
- load_wgnt_from_rsrc
# Single-arch builds: nfpm provides is not templated, so the RPM splits per arch.
- &netbird_rpm_build
id: netbird-rpm-amd64
dir: client
binary: netbird
env: [CGO_ENABLED=0]
goos: [linux]
goarch: [amd64]
ldflags:
- -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 }}"
tags:
- load_wgnt_from_rsrc
- <<: *netbird_rpm_build
id: netbird-rpm-arm64
goarch: [arm64]
- <<: *netbird_rpm_build
id: netbird-rpm-arm
goarch: [arm]
- <<: *netbird_rpm_build
id: netbird-rpm-386
goarch: [386]
- id: netbird-static
dir: client
binary: netbird
@@ -223,17 +249,22 @@ nfpms:
postinstall: "release_files/post_install.sh"
preremove: "release_files/pre_remove.sh"
- maintainer: Netbird <dev@netbird.io>
- &netbird_rpm
maintainer: Netbird <dev@netbird.io>
description: Netbird client.
homepage: https://netbird.io/
license: BSD-3-Clause
vendor: NetBird
id: netbird_rpm
id: netbird_rpm_amd64
bindir: /usr/bin
builds:
- netbird
ids:
- netbird-rpm-amd64
formats:
- rpm
# Red Hat certification (RPM Version Handling) requires rpmbuild's ISA
# provide, which nfpm does not emit. The version is filled in by the release job.
provides:
- "netbird(x86-64) = @RPM_EVR@"
# The client verifies TLS to management and signal against the system trust
# store. Red Hat software certification (RPM Dependency Tracking) also
# rejects packages that declare no dependencies at all.
@@ -263,6 +294,27 @@ nfpms:
packager: NetBird <dev@netbird.io>
signature:
key_file: '{{ if index .Env "GPG_RPM_KEY_FILE" }}{{ .Env.GPG_RPM_KEY_FILE }}{{ end }}'
- <<: *netbird_rpm
id: netbird_rpm_arm64
ids:
- netbird-rpm-arm64
provides:
- "netbird(aarch-64) = @RPM_EVR@"
- <<: *netbird_rpm
id: netbird_rpm_arm
ids:
- netbird-rpm-arm
provides:
- "netbird(armv6hl-32) = @RPM_EVR@"
- <<: *netbird_rpm
id: netbird_rpm_386
ids:
- netbird-rpm-386
provides:
- "netbird(x86-32) = @RPM_EVR@"
dockers_v2:
- id: netbird
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
@@ -425,7 +477,7 @@ dockers_v2:
tags:
- "{{ .Version }}"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}latest{{ end }}"
dockerfile: upload-server/Dockerfile
dockerfile: upload-server/Dockerfile.release
platforms:
- linux/amd64
- linux/arm64
@@ -481,6 +533,41 @@ dockers_v2:
"org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}"
"maintainer": "dev@netbird.io"
- id: proxy-ubi
disable: "{{ .Env.SKIP_DOCKER_PUSH }}"
ids:
- netbird-proxy
images:
- netbirdio/reverse-proxy
- ghcr.io/netbirdio/reverse-proxy
tags:
- "{{ .Version }}-ubi"
- "{{ if eq .Env.SKIP_PUBLISH \"false\" }}ubi-latest{{ end }}"
dockerfile: proxy/Dockerfile.ubi
platforms:
- linux/amd64
- linux/arm64
build_args:
VERSION: "{{ .Version }}"
RELEASE: "{{ .Timestamp }}"
hooks:
pre:
- cmd: 'sh proxy/collect-licenses.sh "{{ .ContextDir }}/licenses" amd64 arm64'
env:
- GOOS=linux
- CGO_ENABLED=0
labels:
"org.opencontainers.image.created": "{{.Date}}"
"org.opencontainers.image.version": "{{.Version}}"
"org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}"
annotations:
"org.opencontainers.image.created": "{{.Date}}"
"org.opencontainers.image.title": "{{.ProjectName}}"
"org.opencontainers.image.version": "{{.Version}}"
"org.opencontainers.image.revision": "{{.FullCommit}}"
"org.opencontainers.image.source": "{{.GitURL}}"
"maintainer": "dev@netbird.io"
brews:
- ids:
@@ -513,7 +600,10 @@ uploads:
- name: yum
skip: "{{ .Env.SKIP_PUBLISH }}"
ids:
- netbird_rpm
- netbird_rpm_amd64
- netbird_rpm_arm64
- netbird_rpm_arm
- netbird_rpm_386
mode: archive
target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }}
username: dev@wiretrustee.com
+50
View File
@@ -115,6 +115,56 @@ export NETBIRD_DOMAIN=netbird.example.com; curl -fsSL https://github.com/netbird
See a complete [architecture overview](https://docs.netbird.io/about-netbird/how-netbird-works#architecture) for details.
### Reporting bugs and requesting features
NetBird uses a discussion-first workflow. Bug reports and feature requests start in
[Discussions](https://github.com/netbirdio/netbird/discussions), not as issues.
| What you want to do | Where to go |
| --- | --- |
| Report a bug, regression, or unexpected behavior | [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) |
| Request a feature or share an idea | [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) |
| Ask about setup, configuration, or self-hosting | [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support) |
| Report a security vulnerability | [Security policy](https://github.com/netbirdio/netbird/security/policy), never a public thread |
Our team and maintainers triage discussions, ask follow-up questions, check for duplicates,
and reproduce bugs. Validated reports are promoted to issues. This keeps the issue tracker a clear
answer to one question: what is the team working on.
Please search existing discussions and issues first, including closed ones. If something similar
already exists, upvote it and add your details there instead of opening a duplicate.
For bug reports, include your NetBird version, operating system, deployment type (Cloud,
self-hosted, Kubernetes, or Docker), reproduction steps, expected and actual behavior, and a debug
bundle where relevant:
```shell
netbird version
netbird status -d -A
netbird debug for 1m -A -S -U
```
`-U` uploads the bundle and prints a file key you can paste instead of attaching the archive.
`-A` anonymizes the output, which matters on a public thread. It masks most identifying details
but is not full redaction, so read the bundle before posting it. Two levels are available:
| Level | How to select | What it masks |
| --- | --- | --- |
| `default` | `-A` / `--anonymize` | Public IP addresses, IPv6 ULA addresses, MAC addresses, and domains other than `netbird.io`, `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`. IPv4 private, CGNAT, and link-local ranges are kept |
| `strict` | `--anonymize-level strict` (implies `-A`) | The above, plus IPv4 private, CGNAT, and link-local ranges, peer names in front of `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`, and WireGuard public keys. Labels under `netbird.io` are kept, since it only hosts infrastructure |
See [collecting a debug bundle](https://docs.netbird.io/help/troubleshooting-client#debug-bundle)
and the [CLI reference](https://docs.netbird.io/get-started/cli#debug-for) for details.
See [How to use Discussions, Issues, and Pull Requests](https://github.com/netbirdio/netbird/discussions/6075)
for the full workflow, or [SUPPORT.md](SUPPORT.md) for a shorter version.
### Contributing
Contributions are welcome. Read [CONTRIBUTING.md](CONTRIBUTING.md) first. NetBird works ticket
first, anything that changes behavior needs an issue the team has agreed on before you open a pull
request.
### Community projects
- [NetBird installer script](https://github.com/physk/netbird-installer)
- [netbird-tui](https://github.com/n0pashkov/netbird-tui) - terminal UI for managing NetBird peers, routes, and settings
+121
View File
@@ -0,0 +1,121 @@
# Getting help with NetBird
Where to go depends on what you need. If you are not sure, start with
[Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support)
and we will move it.
## Before you post
1. Search existing [discussions](https://github.com/netbirdio/netbird/discussions) and
[issues](https://github.com/netbirdio/netbird/issues), including closed ones.
2. Check the [documentation](https://docs.netbird.io) and the troubleshooting guides for
[clients](https://docs.netbird.io/help/troubleshooting-client) and
[self-hosted deployments](https://docs.netbird.io/selfhosted/troubleshooting).
3. Remove or anonymize sensitive information from logs, screenshots, and configuration.
If a discussion already covers your problem, upvote it and add your details there rather than
opening a duplicate. Extra reproduction detail, affected versions, and deployment notes are
useful even on an existing thread.
## Community support
Free, for everyone. Covers the NetBird client, open source self-hosted deployments, and general
questions.
| What you want to do | Where to go |
| --- | --- |
| Report a bug, regression, or unexpected behavior | [Issue Triage](https://github.com/netbirdio/netbird/discussions/new?category=issue-triage) |
| Request a feature or share an idea | [Ideas & Feature Requests](https://github.com/netbirdio/netbird/discussions/new?category=ideas-feature-requests) |
| Ask about setup, configuration, or self-hosting | [Q&A / Support](https://github.com/netbirdio/netbird/discussions/new?category=q-a-support) |
| Chat with the community | [Slack](https://docs.netbird.io/slack-url) |
## Paid support
For NetBird Cloud customers and commercial-license self-hosted deployments, covering the
dashboard, control plane, billing, and subscriptions, see
[reporting bugs and issues](https://docs.netbird.io/help/report-bug-issues).
## Security
Do not report security vulnerabilities in public issues or discussions, and do not post secrets,
private keys, internal hostnames, or sensitive logs. Use the
[security policy](https://github.com/netbirdio/netbird/security/policy).
## What makes a report we can act on
For a bug, the most useful reports include:
- NetBird version, and component versions where applicable
- Operating system or environment
- Deployment type: NetBird Cloud, self-hosted, Kubernetes, Docker, or local development
- Current behavior and expected behavior
- The smallest set of steps that reproduces the problem
- Logs, status output, screenshots, or a debug bundle when relevant
- Whether this worked before, and the last known working version
For client reports, these commands usually give us what we need:
```shell
netbird version
netbird status -d -A
netbird debug for 1m -A -S -U
```
`-A` (`--anonymize`) replaces sensitive values consistently across every file in the bundle, so
it stays readable while masking most identifying details. It is not a guarantee of full redaction:
internal address ranges survive at the default level, and interface names, indexes, MTUs, and
flags are never anonymized. Read the bundle before posting it publicly. Two levels are
available:
| Level | How to select | What it masks |
| --- | --- | --- |
| `default` | `-A` / `--anonymize`, or `--anonymize-level default` | Public IP addresses, IPv6 ULA addresses, MAC addresses, and domains other than `netbird.io`, `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`. IPv4 private, CGNAT, and link-local ranges are kept, and interface names are not anonymized |
| `strict` | `--anonymize-level strict` (implies `-A`) | The above, plus IPv4 private, CGNAT, and link-local ranges, peer names in front of `netbird.cloud`, `netbird.selfhosted`, and `netbird.stage`, and WireGuard public keys. Labels under `netbird.io` are kept, since it only hosts infrastructure |
Use `strict` when internal addressing or peer naming is itself sensitive. Either way, private
keys and SSH keys are never included, and the packet capture (`capture.pcap`) is left out of
anonymized bundles because it holds raw decrypted packets.
`-U` (`--upload-bundle`) uploads the bundle and returns a file key you can paste into the thread
instead of attaching an archive. Retention is controlled by the upload service; check its policy
before uploading, and configure cleanup for self-hosted deployments.
For more detail, see [troubleshooting client issues](https://docs.netbird.io/help/troubleshooting-client),
which explains [what a debug bundle contains](https://docs.netbird.io/help/troubleshooting-client#debug-bundle),
and the [CLI reference](https://docs.netbird.io/get-started/cli#debug-for).
Intermittent problems are still worth reporting. They just need enough detail to investigate:
trigger, frequency, timing, timestamps, and any related logs.
For a feature request, describe the problem before the solution: what you are trying to
accomplish, who is affected and how often, why the current behavior or workaround is not enough,
and what you would like to see instead.
## What happens after you post
Our team, maintainers, or community members may ask for missing details, link related
threads, merge duplicates, move your post to a better category, or try to reproduce the problem.
Not every discussion becomes an issue. Some are answered in Q&A, some turn out to be
configuration problems, and some need more information before engineering can act. A
well-answered discussion is still a useful outcome.
When a report is confirmed and actionable, a maintainer opens a validated issue linked back to
the discussion, in whichever repository the fix belongs to. You do not need to know which
repository that is. Routing is part of triage.
## A note on issues
Issues in this repository are maintainer-curated work items. Every open issue is something a
maintainer or contributor can pick up and act on. Issues opened without a linked validated
discussion may be closed and redirected here.
Maintainers can still open issues directly for work found internally, such as regressions caught
during development, planned maintenance, or release blockers.
## Related reading
- [How to use Discussions, Issues, and Pull Requests](https://github.com/netbirdio/netbird/discussions/6075)
- [Moving to a discussion-first approach](https://github.com/netbirdio/netbird/discussions/6074)
- [CONTRIBUTING.md](CONTRIBUTING.md) for opening pull requests
- [CODE_OF_CONDUCT.md](CODE_OF_CONDUCT.md)
+2
View File
@@ -31,6 +31,8 @@ const (
// PasswordRequiredMarker tells Java to prompt for a password and retry. It is
// a string because gomobile flattens errors to their message, so a sentinel
// value would not survive the binding.
//
//nolint:gosec // G101 false positive: a sentinel marker, not a credential
const PasswordRequiredMarker = "netbird-ssh-password-required"
// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation,
+50
View File
@@ -7,6 +7,7 @@ import (
"fmt"
"net/http"
"runtime"
"slices"
"strings"
"sync"
@@ -25,6 +26,30 @@ var serviceCmd = &cobra.Command{
const defaultJSONSocket = "unix:///var/run/netbird-http.sock"
// forbiddenServiceEnvVars are the environment variables the service is never
// registered with, keyed in upper case since these are Windows names. Each one
// decides where the daemon resolves something it then uses with the privileges
// of the account it runs under — LocalSystem on Windows, root elsewhere: the
// executables it runs (PATH, PATHEXT, COMSPEC, SystemRoot, windir) or the
// directory it writes temporary files in (TEMP, TMP). The daemon needs none of
// them, and the utilities it shells out to are resolved by absolute path.
var forbiddenServiceEnvVars = map[string]struct{}{
"PATH": {},
"PATHEXT": {},
"SYSTEMROOT": {},
"WINDIR": {},
"COMSPEC": {},
"TEMP": {},
"TMP": {},
}
// forbiddenServiceEnvPrefixes are the dynamic-loader families, refused whole
// rather than by name: LD_PRELOAD, DYLD_INSERT_LIBRARIES and their siblings all
// reach the loader of the process, the set differs per platform and libc, and
// new members arrive with new OS releases. Listing them one by one is a list
// that is wrong the moment it is written.
var forbiddenServiceEnvPrefixes = []string{"LD_", "DYLD_"}
var (
serviceName string
serviceEnvVars []string
@@ -127,8 +152,33 @@ func parseServiceEnvVars(envVars []string) (map[string]string, error) {
return nil, fmt.Errorf("empty environment variable key in: %s", env)
}
if isForbiddenServiceEnvVar(key) {
return nil, fmt.Errorf("environment variable %s cannot be set on the service: it decides where the service resolves the executables, libraries or temporary files it uses", key)
}
envMap[key] = value
}
return envMap, nil
}
// isForbiddenServiceEnvVar reports whether name is one the service must not be
// registered with.
//
// The names are matched case-insensitively only on Windows, where they are the
// same variable however they are spelled. Elsewhere the environment is
// case-sensitive, so Path and PATH are two different variables and only the
// exact spelling is the one the loader reads.
func isForbiddenServiceEnvVar(name string) bool {
if runtime.GOOS == "windows" {
name = strings.ToUpper(name)
}
if _, forbidden := forbiddenServiceEnvVars[name]; forbidden {
return true
}
return slices.ContainsFunc(forbiddenServiceEnvPrefixes, func(prefix string) bool {
return strings.HasPrefix(name, prefix)
})
}
+50 -6
View File
@@ -14,6 +14,7 @@ import (
"github.com/netbirdio/netbird/client/configs"
"github.com/netbirdio/netbird/client/internal/daemonaddr"
"github.com/netbirdio/netbird/client/internal/elevate"
"github.com/netbirdio/netbird/util"
)
@@ -43,10 +44,33 @@ func serviceParamsPath() string {
// loadServiceParams reads saved service parameters from disk.
// Returns nil with no error if the file does not exist.
//
// The file is read by an elevated install and decides the arguments and the
// environment of the service it then registers, so it is used only when its
// ownership and permissions are the ones saveServiceParams leaves behind. That
// restricted ACL is applied when the file is written, which is not necessarily
// before it is first read, so this is checked rather than assumed. A file that
// fails the check is treated as absent, and the install proceeds with its
// defaults.
func loadServiceParams() (*serviceParams, error) {
path := serviceParamsPath()
data, err := os.ReadFile(path)
// Resolve links first so the checks apply to the file that is actually read.
// Since the check covers every directory above it as well, nobody who fails
// it can swap the file between here and the read below.
resolved, err := filepath.EvalSymlinks(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil //nolint:nilnil
}
return nil, fmt.Errorf("resolve service params %s: %w", path, err)
}
if err := elevate.CheckOnlyOwnerWritable(resolved); err != nil {
return nil, fmt.Errorf("refusing to read service params from %s: %w", resolved, err)
}
data, err := os.ReadFile(resolved)
if err != nil {
if os.IsNotExist(err) {
return nil, nil //nolint:nilnil
@@ -182,10 +206,16 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
// If --service-env was explicitly set to empty, all saved env vars are cleared.
// If --service-env was not set, saved env vars are used entirely.
func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
// A forbidden name explicitly passed on the command line is an error the
// operator is told about, but one restored from a file written by an older
// version is dropped: an install that refuses to run would leave the host
// without a daemon over a variable nobody is asking for any more.
saved := dropForbiddenServiceEnvVars(cmd, params.ServiceEnvVars)
if !cmd.Flags().Changed("service-env") {
if len(params.ServiceEnvVars) > 0 {
if len(saved) > 0 {
// No explicit env vars: rebuild serviceEnvVars from saved params.
serviceEnvVars = envMapToSlice(params.ServiceEnvVars)
serviceEnvVars = envMapToSlice(saved)
}
return
}
@@ -204,13 +234,13 @@ func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
return
}
if len(params.ServiceEnvVars) == 0 {
if len(saved) == 0 {
return
}
// Merge saved values underneath explicit ones.
merged := make(map[string]string, len(params.ServiceEnvVars)+len(explicit))
maps.Copy(merged, params.ServiceEnvVars)
merged := make(map[string]string, len(saved)+len(explicit))
maps.Copy(merged, saved)
maps.Copy(merged, explicit) // explicit wins on conflict
serviceEnvVars = envMapToSlice(merged)
}
@@ -233,6 +263,20 @@ var resetParamsCmd = &cobra.Command{
},
}
// dropForbiddenServiceEnvVars returns the saved entries that may still be
// registered on the service, reporting every one it leaves behind.
func dropForbiddenServiceEnvVars(cmd *cobra.Command, saved map[string]string) map[string]string {
kept := make(map[string]string, len(saved))
for key, value := range saved {
if isForbiddenServiceEnvVar(key) {
cmd.PrintErrf("Warning: ignoring saved service environment variable %s: it decides where the service resolves the executables, libraries or temporary files it uses\n", key)
continue
}
kept[key] = value
}
return kept
}
// envMapToSlice converts a map of env vars to a KEY=VALUE slice.
func envMapToSlice(m map[string]string) []string {
s := make([]string, 0, len(m))
+54
View File
@@ -9,6 +9,7 @@ import (
"go/token"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
@@ -353,6 +354,59 @@ func TestApplyServiceEnvParams_NotChanged(t *testing.T) {
assert.Equal(t, map[string]string{"FROM_SAVED": "val"}, result)
}
func TestParseServiceEnvVars_RejectsForbiddenNames(t *testing.T) {
for _, env := range []string{"PATH=C:\\somewhere", "LD_PRELOAD=/tmp/lib.so", "DYLD_FALLBACK_LIBRARY_PATH=/tmp"} {
_, err := parseServiceEnvVars([]string{"KEEP=me", env})
require.Errorf(t, err, "%s selects what the service resolves and must be refused", env)
}
}
func TestIsForbiddenServiceEnvVar(t *testing.T) {
// The loader families are matched by prefix, so a name nobody has heard of
// yet is refused too.
for _, name := range []string{
"PATH", "PATHEXT", "COMSPEC", "SYSTEMROOT", "WINDIR", "TEMP", "TMP",
"LD_PRELOAD", "LD_AUDIT", "DYLD_INSERT_LIBRARIES", "DYLD_FALLBACK_FRAMEWORK_PATH",
} {
assert.Truef(t, isForbiddenServiceEnvVar(name), "%s must be refused", name)
}
// The prefix must not swallow names that merely start with the same letters.
for _, name := range []string{"NB_LOG_LEVEL", "NB_WG_DEBUG", "HTTPS_PROXY", "LDAP_URL", "DYLDX"} {
assert.Falsef(t, isForbiddenServiceEnvVar(name), "%s has no reason to be refused", name)
}
// On Windows a variable is the same one however it is spelled; elsewhere
// Path and PATH are two variables and only the exact one is read.
if runtime.GOOS == "windows" {
assert.True(t, isForbiddenServiceEnvVar("Path"))
assert.True(t, isForbiddenServiceEnvVar("ld_preload"))
} else {
assert.False(t, isForbiddenServiceEnvVar("Path"))
assert.False(t, isForbiddenServiceEnvVar("ld_preload"))
}
}
func TestApplyServiceEnvParams_DropsForbiddenSavedNames(t *testing.T) {
origServiceEnvVars := serviceEnvVars
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
serviceEnvVars = nil
cmd := &cobra.Command{}
cmd.Flags().StringSlice("service-env", nil, "")
saved := &serviceParams{
ServiceEnvVars: map[string]string{"PATH": "C:\\attacker", "NB_LOG_FORMAT": "json"},
}
applyServiceEnvParams(cmd, saved)
result, err := parseServiceEnvVars(serviceEnvVars)
require.NoError(t, err, "a saved PATH must be dropped rather than fail the install")
assert.Equal(t, map[string]string{"NB_LOG_FORMAT": "json"}, result)
}
func TestApplyServiceEnvParams_ExplicitEmptyClears(t *testing.T) {
origServiceEnvVars := serviceEnvVars
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
+57
View File
@@ -0,0 +1,57 @@
//go:build !windows && !ios && !android
package cmd
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/configs"
)
// The Windows equivalent of this is the ACL check in
// elevate.CheckOnlyOwnerWritable, covered by that package's own tests; here the
// point is that loadServiceParams asks the question at all.
func TestLoadServiceParams_RefusesWorldWritableFile(t *testing.T) {
tmpDir := t.TempDir()
original := configs.StateDir
t.Cleanup(func() { configs.StateDir = original })
configs.StateDir = tmpDir
path := filepath.Join(tmpDir, serviceParamsFile)
require.NoError(t, os.WriteFile(path, []byte(`{"log_level":"debug"}`), 0o666))
// WriteFile is subject to the umask, so set the bits that matter explicitly.
require.NoError(t, os.Chmod(path, 0o666))
params, err := loadServiceParams()
require.Error(t, err, "a service.json anyone can rewrite must not be trusted")
assert.Nil(t, params)
require.NoError(t, os.Chmod(path, 0o600))
params, err = loadServiceParams()
require.NoError(t, err)
require.NotNil(t, params)
assert.Equal(t, "debug", params.LogLevel)
}
func TestLoadServiceParams_RefusesWorldWritableDirectory(t *testing.T) {
tmpDir := t.TempDir()
stateDir := filepath.Join(tmpDir, "state")
require.NoError(t, os.Mkdir(stateDir, 0o777))
require.NoError(t, os.Chmod(stateDir, 0o777))
original := configs.StateDir
t.Cleanup(func() { configs.StateDir = original })
configs.StateDir = stateDir
require.NoError(t, os.WriteFile(filepath.Join(stateDir, serviceParamsFile), []byte(`{}`), 0o600))
params, err := loadServiceParams()
require.Error(t, err, "a service.json in a directory anyone can replace entries in must not be trusted")
assert.Nil(t, params)
}
@@ -9,6 +9,7 @@ import (
log "github.com/sirupsen/logrus"
nberrors "github.com/netbirdio/netbird/client/errors"
"github.com/netbirdio/netbird/client/internal/wincmd"
)
type action string
@@ -91,7 +92,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
if action == addRule {
args = append(args, extraArgs...)
}
netshCmd := GetSystem32Command("netsh")
netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
return cmd.Run()
@@ -100,7 +101,7 @@ func manageFirewallRule(ruleName string, action action, extraArgs ...string) err
func isWindowsFirewallReachable() bool {
args := []string{"advfirewall", "show", "allprofiles", "state"}
netshCmd := GetSystem32Command("netsh")
netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
@@ -117,23 +118,10 @@ func isWindowsFirewallReachable() bool {
func isFirewallRuleActive(ruleName string) bool {
args := []string{"advfirewall", "firewall", "show", "rule", "name=" + ruleName}
netshCmd := GetSystem32Command("netsh")
netshCmd := wincmd.System32("netsh")
cmd := exec.Command(netshCmd, args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
_, err := cmd.Output()
return err == nil
}
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
func GetSystem32Command(command string) string {
_, err := exec.LookPath(command)
if err == nil {
return command
}
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
return "C:\\windows\\system32\\" + command + ".exe"
}
+226
View File
@@ -0,0 +1,226 @@
package configurer
import (
"net"
"net/netip"
"slices"
"sync"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
// allowedIPStore mirrors the allowed IPs configured on each peer of a device.
//
// A configurer is the only writer of its device's peer set, so the mirror is authoritative
// by construction. It spares the paths that have to rewrite one peer's allowed IPs a full
// device dump just to recover prefixes the process already configured itself.
//
// An allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away
// from whichever peer held it before, and the configurer leaves that handover to the device
// rather than removing the prefix from the previous holder itself. The store tracks the
// owner of each prefix and performs the same handover, so rewriting one peer's list never
// takes a prefix back from the peer that owns it now.
//
// Its own lock guards the map alone, not the device write it accompanies. Consistency
// between the two rests on the caller serializing every configurer call, which WGIface
// does with its mutex; two unserialized writers would interleave a device write with the
// record of a different one.
//
// An operator reconfiguring the device out of band, through `wg set` or the UAPI socket,
// is the one way the mirror can still go stale. A peer missing from it falls back to the
// device, which reseats that peer's prefixes and their ownership; a peer that is present
// does not, so one recorded from empty while the device already held prefixes keeps only
// what was recorded, and the next endpoint removal drops the rest.
type allowedIPStore struct {
mu sync.RWMutex
peers map[wgtypes.Key][]netip.Prefix
owners map[netip.Prefix]wgtypes.Key
}
func newAllowedIPStore() *allowedIPStore {
return &allowedIPStore{
peers: make(map[wgtypes.Key][]netip.Prefix),
owners: make(map[netip.Prefix]wgtypes.Key),
}
}
// get returns the prefixes recorded for a peer, and whether the peer is known at all.
// The caller receives a copy and may retain or modify it freely.
func (s *allowedIPStore) get(key wgtypes.Key) ([]netip.Prefix, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
prefixes, ok := s.peers[key]
if !ok {
return nil, false
}
return slices.Clone(prefixes), true
}
// set replaces the prefixes recorded for a peer.
func (s *allowedIPStore) set(key wgtypes.Key, prefixes []netip.Prefix) {
s.mu.Lock()
defer s.mu.Unlock()
k := key
s.releaseLocked(k)
normalized := normalizePrefixes(prefixes)
for _, prefix := range normalized {
s.claimLocked(k, prefix)
}
s.peers[k] = normalized
}
// add records prefixes on a peer without dropping the ones already there, matching the
// union semantics of a peer update that does not replace its allowed IPs. It records the
// peer if it is not known yet, so it belongs to the operations that create a peer on the
// device rather than to the update-only ones.
func (s *allowedIPStore) add(key wgtypes.Key, prefixes []netip.Prefix) {
s.mu.Lock()
defer s.mu.Unlock()
s.mergeLocked(key, prefixes)
}
// addExisting is add for an update-only device operation. Such an operation is a silent
// no-op when the peer is absent, so recording a peer here would leave the store claiming
// prefixes the device never took, and the peer would then be recreated by the next endpoint
// removal, stealing those allowed IPs from the peer that legitimately holds them.
func (s *allowedIPStore) addExisting(key wgtypes.Key, prefixes []netip.Prefix) {
s.mu.Lock()
defer s.mu.Unlock()
k := key
if _, ok := s.peers[k]; !ok {
return
}
s.mergeLocked(k, prefixes)
}
// ensure records a peer with no prefixes unless it is already known. A device operation
// that is not update-only creates the peer when it is absent, so it has to be recorded even
// when it configures nothing else; otherwise the peer exists on the device while the store
// treats it as unknown, and a prefix later handed over to it is not accounted for.
func (s *allowedIPStore) ensure(key wgtypes.Key) {
s.mu.Lock()
defer s.mu.Unlock()
k := key
if _, ok := s.peers[k]; !ok {
s.peers[k] = nil
}
}
// forget drops every prefix recorded for a peer.
func (s *allowedIPStore) forget(key wgtypes.Key) {
s.mu.Lock()
defer s.mu.Unlock()
k := key
s.releaseLocked(k)
delete(s.peers, k)
}
// reset drops every peer, mirroring a device reconfiguration that replaces the peer set.
func (s *allowedIPStore) reset() {
s.mu.Lock()
defer s.mu.Unlock()
s.peers = make(map[wgtypes.Key][]netip.Prefix)
s.owners = make(map[netip.Prefix]wgtypes.Key)
}
// mergeLocked unions normalized prefixes into a peer and transfers their ownership.
// The caller must hold s.mu for writing.
func (s *allowedIPStore) mergeLocked(k wgtypes.Key, prefixes []netip.Prefix) {
merged := s.peers[k]
for _, prefix := range prefixes {
prefix = normalizePrefix(prefix)
s.claimLocked(k, prefix)
if !slices.Contains(merged, prefix) {
merged = append(merged, prefix)
}
}
s.peers[k] = merged
}
// claimLocked hands a prefix over to a peer, taking it from its previous owner the way the
// device does when the same prefix is configured on a second peer.
func (s *allowedIPStore) claimLocked(k wgtypes.Key, prefix netip.Prefix) {
if owner, ok := s.owners[prefix]; ok && owner != k {
s.peers[owner] = slices.DeleteFunc(s.peers[owner], func(p netip.Prefix) bool {
return p == prefix
})
}
s.owners[prefix] = k
}
// releaseLocked drops a peer's claim on every prefix it currently holds.
func (s *allowedIPStore) releaseLocked(k wgtypes.Key) {
for _, prefix := range s.peers[k] {
if s.owners[prefix] == k {
delete(s.owners, prefix)
}
}
}
// normalizePrefix puts a prefix into the form the store recognises it by. It clears the
// host bits, which a device does on its own, so a caller passing 10.20.0.1/16 still matches
// the 10.20.0.0/16 read back from the device; and it unmaps a v4-mapped prefix so that it
// compares equal to, and marshals like, the plain v4 prefix for the same network.
//
// Masking comes first because it also decides the address family: only a prefix at least 96
// bits long keeps the mapped marker through the mask, so a shorter prefix inside the mapped
// range is a genuine v6 prefix and unmapping it would yield an invalid v4 prefix.
func normalizePrefix(prefix netip.Prefix) netip.Prefix {
masked := prefix.Masked()
addr := masked.Addr()
if !addr.Is4In6() {
return masked
}
return netip.PrefixFrom(addr.Unmap(), masked.Bits()-96)
}
// normalizePrefixes returns a normalized copy without changing the caller's slice.
func normalizePrefixes(prefixes []netip.Prefix) []netip.Prefix {
normalized := make([]netip.Prefix, len(prefixes))
for i, prefix := range prefixes {
normalized[i] = normalizePrefix(prefix)
}
return normalized
}
// ipNetsToPrefixes converts addresses read back from a device. Unmap keeps a v4-mapped v6
// address comparable to the plain v4 prefix the configurer was given.
func ipNetsToPrefixes(ipNets []net.IPNet) []netip.Prefix {
prefixes := make([]netip.Prefix, 0, len(ipNets))
for _, ipNet := range ipNets {
addr, ok := netip.AddrFromSlice(ipNet.IP)
if !ok {
continue
}
ones, maskBits := ipNet.Mask.Size()
// A device may report a v4 prefix as a v4-mapped address. Align the address form with
// the mask rather than unmapping on sight: a 32 bit mask always describes v4, while a
// 128 bit mask describes v4 only when it covers the mapped prefix, so a genuine v6
// prefix inside the mapped range stays v6 instead of being dropped as invalid.
if addr.Is4In6() {
switch {
case maskBits == 32:
addr = addr.Unmap()
case maskBits == 128 && ones >= 96:
addr, ones = addr.Unmap(), ones-96
}
}
prefix := netip.PrefixFrom(addr, ones)
if !prefix.IsValid() {
continue
}
prefixes = append(prefixes, prefix.Masked())
}
return prefixes
}
+263
View File
@@ -0,0 +1,263 @@
package configurer
import (
"net"
"net/netip"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
// The store keys on the parsed key, so the tests use two distinct ones rather than names.
var (
testPeer = wgtypes.Key{1}
otherPeer = wgtypes.Key{2}
)
func TestAllowedIPStoreUnknownPeer(t *testing.T) {
s := newAllowedIPStore()
prefixes, ok := s.get(testPeer)
assert.False(t, ok, "an unconfigured peer must be reported as unknown, not as one without prefixes")
assert.Nil(t, prefixes, "an unknown peer has no prefixes")
}
func TestAllowedIPStoreAddUnions(t *testing.T) {
s := newAllowedIPStore()
overlay := netip.MustParsePrefix("100.64.0.1/32")
routed := netip.MustParsePrefix("10.20.0.0/16")
s.set(testPeer, []netip.Prefix{overlay})
// A peer update does not replace allowed IPs, and a repeated prefix must not be doubled.
s.add(testPeer, []netip.Prefix{overlay, routed})
prefixes, ok := s.get(testPeer)
require.True(t, ok, "peer must be known after set")
assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "add must union rather than replace")
}
func TestAllowedIPStoreGetReturnsCopy(t *testing.T) {
s := newAllowedIPStore()
overlay := netip.MustParsePrefix("100.64.0.1/32")
s.set(testPeer, []netip.Prefix{overlay})
prefixes, ok := s.get(testPeer)
require.True(t, ok, "peer must be known after set")
prefixes[0] = netip.MustParsePrefix("0.0.0.0/0")
stored, _ := s.get(testPeer)
assert.Equal(t, []netip.Prefix{overlay}, stored, "a caller mutating the returned slice must not corrupt the store")
}
func TestAllowedIPStoreForgetAndReset(t *testing.T) {
s := newAllowedIPStore()
s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")})
s.set(otherPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")})
s.forget(testPeer)
_, ok := s.get(testPeer)
assert.False(t, ok, "a forgotten peer must be unknown")
_, ok = s.get(otherPeer)
assert.True(t, ok, "forgetting one peer must not touch the others")
s.reset()
_, ok = s.get(otherPeer)
assert.False(t, ok, "reset must drop every peer")
}
func TestIPNetsToPrefixes(t *testing.T) {
tests := []struct {
name string
ipNet net.IPNet
want string
}{
{
name: "v4",
ipNet: net.IPNet{IP: net.IP{10, 20, 0, 0}, Mask: net.CIDRMask(16, 32)},
want: "10.20.0.0/16",
},
{
name: "v4 mapped under a 128 bit mask",
ipNet: net.IPNet{IP: net.ParseIP("10.20.0.0"), Mask: net.CIDRMask(112, 128)},
want: "10.20.0.0/16",
},
{
name: "v6",
ipNet: net.IPNet{IP: net.ParseIP("fd00::"), Mask: net.CIDRMask(64, 128)},
want: "fd00::/64",
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got := ipNetsToPrefixes([]net.IPNet{tc.ipNet})
require.Len(t, got, 1, "the address must be converted, not dropped")
assert.Equal(t, tc.want, got[0].String(), "converted prefix")
})
}
}
func TestIPNetsToPrefixesRoundTrip(t *testing.T) {
prefixes := []netip.Prefix{
netip.MustParsePrefix("100.64.0.1/32"),
netip.MustParsePrefix("10.20.0.0/16"),
netip.MustParsePrefix("fd00::/64"),
}
assert.Equal(t, prefixes, ipNetsToPrefixes(prefixesToIPNets(prefixes)),
"prefixes handed to a device must come back unchanged")
}
func TestAllowedIPStoreNormalizesMappedPrefixes(t *testing.T) {
s := newAllowedIPStore()
v4 := netip.MustParsePrefix("10.20.0.0/16")
mapped := netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112)
s.set(testPeer, []netip.Prefix{mapped})
// A v4 rule only matches a v4-mapped address once it has been unmapped, so the store must
// hold the plain form and recognise the two spellings as the same prefix.
s.add(testPeer, []netip.Prefix{v4})
prefixes, ok := s.get(testPeer)
require.True(t, ok, "peer must be known after set")
assert.Equal(t, []netip.Prefix{v4}, prefixes, "a mapped prefix must be stored unmapped and not duplicated")
}
func TestNormalizePrefix(t *testing.T) {
v4 := netip.MustParsePrefix("10.20.0.0/16")
v6 := netip.MustParsePrefix("fd00::/64")
assert.Equal(t, v4, normalizePrefix(v4), "a plain v4 prefix is unchanged")
assert.Equal(t, v6, normalizePrefix(v6), "a real v6 prefix is unchanged")
assert.Equal(t, v4, normalizePrefix(netip.PrefixFrom(netip.AddrFrom16(v4.Addr().As16()), 112)),
"a mapped prefix under a 128 bit mask becomes plain v4")
// A prefix shorter than /96 inside the mapped range is a genuine v6 prefix. Unmapping it
// would pair a v4 address with a v6 sized mask, which is invalid, and the store would then
// record a zero prefix that can never recreate the allowed IP.
for _, tc := range []string{"::ffff:0:0/64", "::ffff:1.2.3.4/80", "::ffff:1.2.3.4/95"} {
got := normalizePrefix(netip.MustParsePrefix(tc))
assert.True(t, got.IsValid(), "%s must normalize to a valid prefix", tc)
assert.False(t, got.Addr().Is4(), "%s must stay v6", tc)
}
}
func TestAllowedIPStoreAddExistingDoesNotCreate(t *testing.T) {
s := newAllowedIPStore()
routed := netip.MustParsePrefix("10.20.0.0/16")
// An update-only device operation on an absent peer is a silent no-op, so nothing may be
// recorded for a peer the store does not already know.
s.addExisting(testPeer, []netip.Prefix{routed})
_, ok := s.get(testPeer)
assert.False(t, ok, "addExisting must not record an unknown peer")
overlay := netip.MustParsePrefix("100.64.0.1/32")
s.set(testPeer, []netip.Prefix{overlay})
s.addExisting(testPeer, []netip.Prefix{routed})
prefixes, _ := s.get(testPeer)
assert.Equal(t, []netip.Prefix{overlay, routed}, prefixes, "addExisting must union onto a known peer")
}
func TestAllowedIPStoreHandsPrefixOverToTheNewOwner(t *testing.T) {
s := newAllowedIPStore()
routed := netip.MustParsePrefix("10.20.0.0/16")
other := otherPeer
s.set(testPeer, []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32"), routed})
s.set(other, []netip.Prefix{netip.MustParsePrefix("100.64.0.2/32")})
// The device takes an allowed IP away from its previous holder when it is configured on
// another peer, so the store must do the same rather than list it under both.
s.addExisting(other, []netip.Prefix{routed})
previous, _ := s.get(testPeer)
assert.NotContains(t, previous, routed, "the previous owner must lose the prefix")
current, _ := s.get(other)
assert.Contains(t, current, routed, "the new owner must hold the prefix")
}
func TestAllowedIPStoreForgetReleasesOwnership(t *testing.T) {
s := newAllowedIPStore()
routed := netip.MustParsePrefix("10.20.0.0/16")
s.set(testPeer, []netip.Prefix{routed})
s.forget(testPeer)
s.set(otherPeer, []netip.Prefix{routed})
// A forgotten peer must not be resurrected as a key in the peer map by a later claim.
_, ok := s.get(testPeer)
assert.False(t, ok, "the forgotten peer must stay unknown")
current, _ := s.get(otherPeer)
assert.Equal(t, []netip.Prefix{routed}, current, "the new owner must hold the prefix")
}
func TestNormalizePrefixClearsHostBits(t *testing.T) {
// A device stores a prefix masked, so a caller passing host bits must still match what a
// device fallback seeded, otherwise that prefix could never be removed by value.
assert.Equal(t, netip.MustParsePrefix("10.20.0.0/16"),
normalizePrefix(netip.MustParsePrefix("10.20.0.1/16")), "host bits must be cleared")
assert.Equal(t, netip.MustParsePrefix("fd00::/64"),
normalizePrefix(netip.MustParsePrefix("fd00::1/64")), "host bits must be cleared for v6")
}
func TestIPNetsToPrefixesKeepsV6InTheMappedRange(t *testing.T) {
// ::ffff:0:0/64 reads as v4-mapped but is a genuine v6 prefix: unmapping it would leave a
// v4 address under a 64 bit mask, which is invalid, and the allowed IP would be dropped.
got := ipNetsToPrefixes([]net.IPNet{{
IP: net.ParseIP("::ffff:0:0"),
Mask: net.CIDRMask(64, 128),
}})
require.Len(t, got, 1, "the prefix must be converted, not dropped")
assert.False(t, got[0].Addr().Is4(), "a v6 prefix in the mapped range must not become v4")
assert.Equal(t, 64, got[0].Bits(), "the prefix length must survive the conversion")
}
func TestPrefixesToIPNetsNormalizes(t *testing.T) {
// net.IPNet prints a v4-mapped address as v4 but takes the length from its 16 byte
// mask, so an unnormalized ::ffff:10.1.2.3/64 reaches a userspace device as 10.1.2.3/0,
// an allowed IP that matches every v4 address.
tests := []struct {
name string
given string
want string
}{
{name: "mapped below /96", given: "::ffff:10.1.2.3/64", want: "::/64"},
{name: "mapped at /112", given: "::ffff:10.1.2.3/112", want: "10.1.0.0/16"},
{name: "host bits are cleared", given: "10.20.0.1/16", want: "10.20.0.0/16"},
{name: "v6 is untouched", given: "fd00::1/64", want: "fd00::/64"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got := prefixesToIPNets([]netip.Prefix{netip.MustParsePrefix(tc.given)})
require.Len(t, got, 1, "the prefix must be converted, not dropped")
assert.Equal(t, tc.want, got[0].String(), "what the device is given")
assert.NotEqual(t, 0, mustOnes(t, got[0]), "a device must never be given a zero length allowed IP")
})
}
}
func mustOnes(t *testing.T, ipNet net.IPNet) int {
t.Helper()
ones, _ := ipNet.Mask.Size()
return ones
}
// TestPrefixesToIPNetsAgreesWithTheStore pins the property the store depends on: what a
// device is given and what is recorded for it are the same prefix.
func TestPrefixesToIPNetsAgreesWithTheStore(t *testing.T) {
for _, given := range []string{"::ffff:10.1.2.3/64", "::ffff:10.1.2.3/112", "10.20.0.1/16", "fd00::1/64"} {
prefix := netip.MustParsePrefix(given)
toDevice := prefixesToIPNets([]netip.Prefix{prefix})
recorded := normalizePrefix(prefix)
assert.Equal(t, recorded.String(), toDevice[0].String(),
"%s must reach the device in the form the store records", given)
}
}
+8 -2
View File
@@ -19,12 +19,18 @@ func buildPresharedKeyConfig(peerKey wgtypes.Key, psk wgtypes.Key, updateOnly bo
}
}
// prefixesToIPNets converts prefixes on their way to a device. It is the only place that
// conversion happens, so it also normalizes: the device is then given the same form the
// store records, and a v4-mapped prefix cannot reach net.IPNet, which prints such an
// address as v4 while taking the length from its 16 byte mask and so turns
// ::ffff:10.1.2.3/64 into 10.1.2.3/0 — an allowed IP matching every v4 address.
func prefixesToIPNets(prefixes []netip.Prefix) []net.IPNet {
ipNets := make([]net.IPNet, len(prefixes))
for i, prefix := range prefixes {
normalized := normalizePrefix(prefix)
ipNets[i] = net.IPNet{
IP: prefix.Addr().AsSlice(), // Convert netip.Addr to net.IP
Mask: net.CIDRMask(prefix.Bits(), prefix.Addr().BitLen()), // Create subnet mask
IP: normalized.Addr().AsSlice(),
Mask: net.CIDRMask(normalized.Bits(), normalized.Addr().BitLen()),
}
}
return ipNets
+74 -34
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"net"
"net/netip"
"slices"
"time"
log "github.com/sirupsen/logrus"
@@ -18,16 +19,22 @@ import (
type KernelConfigurer struct {
deviceName string
statsCache *statsCache
allowedIPs *allowedIPStore
}
// NewKernelConfigurer creates a configurer with an empty allowed IP mirror
// and a statistics cache for the named kernel device.
func NewKernelConfigurer(deviceName string) *KernelConfigurer {
c := &KernelConfigurer{
deviceName: deviceName,
allowedIPs: newAllowedIPStore(),
}
c.statsCache = newStatsCache(statsCacheTTL, c.fetchStats)
return c
}
// ConfigureInterface sets the device key, port and firewall mark, replacing all peers.
// The allowed IP mirror is reset only after the device accepts the configuration.
func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error {
log.Debugf("adding Wireguard private key")
key, err := wgtypes.ParseKey(privateKey)
@@ -46,6 +53,8 @@ func (c *KernelConfigurer) ConfigureInterface(privateKey string, port int) error
if err != nil {
return fmt.Errorf(`received error "%w" while configuring interface %s with port %d`, err, c.deviceName, port)
}
c.allowedIPs.reset()
return nil
}
@@ -58,9 +67,20 @@ func (c *KernelConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, upda
}
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
return c.configure(cfg)
if err := c.configure(cfg); err != nil {
return err
}
// Without updateOnly this creates the peer when it is absent, so the store has to
// know about it even though no allowed IP was configured.
if !updateOnly {
c.allowedIPs.ensure(parsedPeerKey)
}
return nil
}
// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set.
// Prefixes assigned to this peer are transferred from their previous owners.
func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
@@ -83,19 +103,23 @@ func (c *KernelConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
if err != nil {
return fmt.Errorf(`received error "%w" while updating peer on interface %s with settings: allowed ips %s, endpoint %s`, err, c.deviceName, allowedIps, endpoint.String())
}
c.allowedIPs.add(peerKeyParsed, allowedIps)
return nil
}
// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured.
// Neither the netlink API nor the userspace one can clear an endpoint in place, so the peer
// is removed and re-added with the allowed IPs it already had.
func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return err
}
// Get the existing peer to preserve its allowed IPs
existingPeer, err := c.getPeer(c.deviceName, peerKey)
allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil {
return fmt.Errorf("get peer: %w", err)
return err
}
removePeerCfg := wgtypes.PeerConfig{
@@ -104,26 +128,27 @@ func (c *KernelConfigurer) RemoveEndpointAddress(peerKey string) error {
}
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{removePeerCfg}}); err != nil {
return fmt.Errorf(`error removing peer %s from interface %s: %w`, peerKey, c.deviceName, err)
return fmt.Errorf("remove peer %s from interface %s: %w", peerKey, c.deviceName, err)
}
//Re-add the peer without the endpoint but same AllowedIPs
reAddPeerCfg := wgtypes.PeerConfig{
PublicKey: peerKeyParsed,
AllowedIPs: existingPeer.AllowedIPs,
AllowedIPs: prefixesToIPNets(allowedIPs),
ReplaceAllowedIPs: true,
}
if err := c.configure(wgtypes.Config{Peers: []wgtypes.PeerConfig{reAddPeerCfg}}); err != nil {
c.allowedIPs.forget(peerKeyParsed)
return fmt.Errorf(
`error re-adding peer %s to interface %s with allowed IPs %v: %w`,
peerKey, c.deviceName, existingPeer.AllowedIPs, err,
"re-add peer %s to interface %s with allowed IPs %v: %w",
peerKey, c.deviceName, allowedIPs, err,
)
}
return nil
}
// RemovePeer removes a peer and forgets its allowed IPs after a successful device write.
func (c *KernelConfigurer) RemovePeer(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
@@ -142,15 +167,13 @@ func (c *KernelConfigurer) RemovePeer(peerKey string) error {
if err != nil {
return fmt.Errorf(`received error "%w" while removing peer %s from interface %s`, err, peerKey, c.deviceName)
}
c.allowedIPs.forget(peerKeyParsed)
return nil
}
// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op.
func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
ipNet := net.IPNet{
IP: allowedIP.Addr().AsSlice(),
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return err
@@ -159,7 +182,7 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
PublicKey: peerKeyParsed,
UpdateOnly: true,
ReplaceAllowedIPs: false,
AllowedIPs: []net.IPNet{ipNet},
AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
}
config := wgtypes.Config{
@@ -169,52 +192,69 @@ func (c *KernelConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix)
if err != nil {
return fmt.Errorf(`received error "%w" while adding allowed Ip to peer on interface %s with settings: allowed ips %s`, err, c.deviceName, allowedIP)
}
c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP})
return nil
}
// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs.
// A prefix not assigned to the peer is a no-op.
func (c *KernelConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
ipNet := net.IPNet{
IP: allowedIP.Addr().AsSlice(),
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return fmt.Errorf("parse peer key: %w", err)
}
existingPeer, err := c.getPeer(c.deviceName, peerKey)
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil {
return fmt.Errorf("get peer: %w", err)
return err
}
newAllowedIPs := existingPeer.AllowedIPs
for i, existingAllowedIP := range existingPeer.AllowedIPs {
if existingAllowedIP.String() == ipNet.String() {
newAllowedIPs = append(existingPeer.AllowedIPs[:i], existingPeer.AllowedIPs[i+1:]...) //nolint:gocritic
break
}
idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
if idx < 0 {
return nil
}
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed,
UpdateOnly: true,
ReplaceAllowedIPs: true,
AllowedIPs: newAllowedIPs,
AllowedIPs: prefixesToIPNets(newAllowedIPs),
}
config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer},
}
err = c.configure(config)
if err != nil {
if err := c.configure(config); err != nil {
return fmt.Errorf("remove allowed IP %s on interface %s: %w", allowedIP, c.deviceName, err)
}
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
return nil
}
func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device
// only for a peer the store has not seen. Dumping the device costs a netlink round trip
// proportional to the whole network map, and this runs on every relay and ICE transition.
func (c *KernelConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) {
if prefixes, ok := c.allowedIPs.get(peerKey); ok {
return prefixes, nil
}
existingPeer, err := c.getPeer(c.deviceName, peerKey)
if err != nil {
return nil, fmt.Errorf("get peer: %w", err)
}
prefixes := ipNetsToPrefixes(existingPeer.AllowedIPs)
c.allowedIPs.set(peerKey, prefixes)
return prefixes, nil
}
// getPeer scans the device for one peer. wgtypes.Key is an array, so the comparison is a
// plain equality: Key.String would base64 encode into a fresh allocation for every peer.
func (c *KernelConfigurer) getPeer(ifaceName string, peerPubKey wgtypes.Key) (wgtypes.Peer, error) {
wg, err := wgctrl.New()
if err != nil {
return wgtypes.Peer{}, fmt.Errorf("wgctl: %w", err)
@@ -231,7 +271,7 @@ func (c *KernelConfigurer) getPeer(ifaceName, peerPubKey string) (wgtypes.Peer,
return wgtypes.Peer{}, fmt.Errorf("get device %s: %w", ifaceName, err)
}
for _, peer := range wgDevice.Peers {
if peer.PublicKey.String() == peerPubKey {
if peer.PublicKey == peerPubKey {
return peer, nil
}
}
+120 -92
View File
@@ -8,6 +8,7 @@ import (
"net/netip"
"os"
"runtime"
"slices"
"strconv"
"strings"
"time"
@@ -41,31 +42,38 @@ type WGUSPConfigurer struct {
deviceName string
activityRecorder *bind.ActivityRecorder
statsCache *statsCache
allowedIPs *allowedIPStore
uapiListener net.Listener
}
// NewUSPConfigurer creates a userspace configurer and starts its UAPI listener.
func NewUSPConfigurer(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
wgCfg := &WGUSPConfigurer{
device: device,
deviceName: deviceName,
activityRecorder: activityRecorder,
allowedIPs: newAllowedIPStore(),
}
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
wgCfg.startUAPI()
return wgCfg
}
// NewUSPConfigurerNoUAPI creates a userspace configurer without a UAPI listener.
func NewUSPConfigurerNoUAPI(device *device.Device, deviceName string, activityRecorder *bind.ActivityRecorder) *WGUSPConfigurer {
wgCfg := &WGUSPConfigurer{
device: device,
deviceName: deviceName,
activityRecorder: activityRecorder,
allowedIPs: newAllowedIPStore(),
}
wgCfg.statsCache = newStatsCache(statsCacheTTL, wgCfg.fetchStats)
return wgCfg
}
// ConfigureInterface sets the device key, port and firewall mark, replacing all peers.
// The allowed IP mirror is reset only after the device accepts the configuration.
func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error {
log.Debugf("adding Wireguard private key")
key, err := wgtypes.ParseKey(privateKey)
@@ -80,7 +88,12 @@ func (c *WGUSPConfigurer) ConfigureInterface(privateKey string, port int) error
ListenPort: &port,
}
return c.device.IpcSet(toWgUserspaceString(config))
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
return err
}
c.allowedIPs.reset()
return nil
}
// SetPresharedKey sets the preshared key for a peer.
@@ -92,14 +105,38 @@ func (c *WGUSPConfigurer) SetPresharedKey(peerKey string, psk wgtypes.Key, updat
}
cfg := buildPresharedKeyConfig(parsedPeerKey, psk, updateOnly)
return c.device.IpcSet(toWgUserspaceString(cfg))
if err := c.device.IpcSet(toWgUserspaceString(cfg)); err != nil {
return err
}
// Without updateOnly this creates the peer when it is absent, so the store has to
// know about it even though no allowed IP was configured.
if !updateOnly {
c.allowedIPs.ensure(parsedPeerKey)
}
return nil
}
// UpdatePeer creates or updates a peer, merging allowed IPs with its existing set.
// It validates the endpoint before writing and records changes after a successful write.
func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return err
}
// Everything that can fail is done before the device is touched, so a failure here
// cannot leave the device holding a peer that the activity recorder and the allowed
// IP store never learned about.
var addrPort netip.AddrPort
if endpoint != nil {
addr, err := netip.ParseAddr(endpoint.IP.String())
if err != nil {
return fmt.Errorf("parse endpoint address: %w", err)
}
addrPort = netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port))
}
peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed,
ReplaceAllowedIPs: false,
@@ -119,47 +156,27 @@ func (c *WGUSPConfigurer) UpdatePeer(peerKey string, allowedIps []netip.Prefix,
}
if endpoint != nil {
addr, err := netip.ParseAddr(endpoint.IP.String())
if err != nil {
return fmt.Errorf("failed to parse endpoint address: %w", err)
}
addrPort := netip.AddrPortFrom(addr.Unmap(), uint16(endpoint.Port))
c.activityRecorder.UpsertAddress(peerKey, addrPort)
}
c.allowedIPs.add(peerKeyParsed, allowedIps)
return nil
}
// RemoveEndpointAddress clears the endpoint of a peer while keeping it configured.
// The UAPI cannot clear an endpoint in place, so the peer is removed and re-added with the
// allowed IPs it already had.
func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return fmt.Errorf("parse peer key: %w", err)
}
ipcStr, err := c.device.IpcGet()
allowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil {
return fmt.Errorf("get IPC config: %w", err)
return err
}
// Parse current status to get allowed IPs for the peer
stats, err := parseStatus(c.deviceName, ipcStr)
if err != nil {
return fmt.Errorf("parse IPC config: %w", err)
}
var allowedIPs []net.IPNet
found := false
for _, peer := range stats.Peers {
if peer.PublicKey == peerKey {
allowedIPs = peer.AllowedIPs
found = true
break
}
}
if !found {
return fmt.Errorf("peer %s not found", peerKey)
}
// remove the peer from the WireGuard configuration
peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed,
Remove: true,
@@ -169,14 +186,13 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
Peers: []wgtypes.PeerConfig{peer},
}
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
return fmt.Errorf("failed to remove peer: %s", ipcErr)
return fmt.Errorf("remove peer: %w", ipcErr)
}
// Build the peer config
peer = wgtypes.PeerConfig{
PublicKey: peerKeyParsed,
ReplaceAllowedIPs: true,
AllowedIPs: allowedIPs,
AllowedIPs: prefixesToIPNets(allowedIPs),
}
config = wgtypes.Config{
@@ -184,12 +200,15 @@ func (c *WGUSPConfigurer) RemoveEndpointAddress(peerKey string) error {
}
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
return fmt.Errorf("remove endpoint address: %w", err)
c.allowedIPs.forget(peerKeyParsed)
return fmt.Errorf("re-add peer without endpoint: %w", err)
}
return nil
}
// RemovePeer removes a peer, then clears its activity and allowed IP records.
// A failed device write leaves both records intact.
func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
@@ -204,18 +223,17 @@ func (c *WGUSPConfigurer) RemovePeer(peerKey string) error {
config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer},
}
ipcErr := c.device.IpcSet(toWgUserspaceString(config))
c.activityRecorder.Remove(peerKey)
return ipcErr
}
func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
ipNet := net.IPNet{
IP: allowedIP.Addr().AsSlice(),
Mask: net.CIDRMask(allowedIP.Bits(), allowedIP.Addr().BitLen()),
if ipcErr := c.device.IpcSet(toWgUserspaceString(config)); ipcErr != nil {
return ipcErr
}
c.activityRecorder.Remove(peerKey)
c.allowedIPs.forget(peerKeyParsed)
return nil
}
// AddAllowedIP adds a prefix to an existing peer; an absent peer is a silent no-op.
func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) error {
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return err
@@ -224,79 +242,89 @@ func (c *WGUSPConfigurer) AddAllowedIP(peerKey string, allowedIP netip.Prefix) e
PublicKey: peerKeyParsed,
UpdateOnly: true,
ReplaceAllowedIPs: false,
AllowedIPs: []net.IPNet{ipNet},
AllowedIPs: prefixesToIPNets([]netip.Prefix{allowedIP}),
}
config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer},
}
return c.device.IpcSet(toWgUserspaceString(config))
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
return err
}
c.allowedIPs.addExisting(peerKeyParsed, []netip.Prefix{allowedIP})
return nil
}
// RemoveAllowedIP removes a prefix while preserving the peer's other allowed IPs.
// It returns ErrAllowedIPNotFound if the prefix is not assigned to the peer.
func (c *WGUSPConfigurer) RemoveAllowedIP(peerKey string, allowedIP netip.Prefix) error {
ipc, err := c.device.IpcGet()
if err != nil {
return err
}
peerKeyParsed, err := wgtypes.ParseKey(peerKey)
if err != nil {
return fmt.Errorf("parse peer key: %w", err)
}
currentAllowedIPs, err := c.peerAllowedIPs(peerKeyParsed)
if err != nil {
return err
}
hexKey := hex.EncodeToString(peerKeyParsed[:])
lines := strings.Split(ipc, "\n")
idx := slices.Index(currentAllowedIPs, normalizePrefix(allowedIP))
if idx < 0 {
return ErrAllowedIPNotFound
}
newAllowedIPs := slices.Delete(currentAllowedIPs, idx, idx+1)
peer := wgtypes.PeerConfig{
PublicKey: peerKeyParsed,
UpdateOnly: true,
ReplaceAllowedIPs: true,
AllowedIPs: []net.IPNet{},
AllowedIPs: prefixesToIPNets(newAllowedIPs),
}
foundPeer := false
removedAllowedIP := false
ip := allowedIP.String()
for _, line := range lines {
line = strings.TrimSpace(line)
// If we're within the details of the found peer and encounter another public key,
// this means we're starting another peer's details. So, reset the flag.
if strings.HasPrefix(line, "public_key=") && foundPeer {
foundPeer = false
}
// Identify the peer with the specific public key
if line == fmt.Sprintf("public_key=%s", hexKey) {
foundPeer = true
}
// If we're within the details of the found peer and find the specific allowed IP, skip this line
if foundPeer && line == "allowed_ip="+ip {
removedAllowedIP = true
continue
}
// Append the line to the output string
if foundPeer && strings.HasPrefix(line, "allowed_ip=") {
allowedIPStr := strings.TrimPrefix(line, "allowed_ip=")
_, ipNet, err := net.ParseCIDR(allowedIPStr)
if err != nil {
return err
}
peer.AllowedIPs = append(peer.AllowedIPs, *ipNet)
}
}
if !removedAllowedIP {
return ErrAllowedIPNotFound
}
config := wgtypes.Config{
Peers: []wgtypes.PeerConfig{peer},
}
return c.device.IpcSet(toWgUserspaceString(config))
if err := c.device.IpcSet(toWgUserspaceString(config)); err != nil {
return fmt.Errorf("remove allowed IP %s: %w", allowedIP, err)
}
c.allowedIPs.set(peerKeyParsed, newAllowedIPs)
return nil
}
// peerAllowedIPs returns the allowed IPs configured for a peer, reading them from the device
// only for a peer the store has not seen. Reading them back means dumping and parsing the
// whole device configuration, and this runs on every relay and ICE transition.
func (c *WGUSPConfigurer) peerAllowedIPs(peerKey wgtypes.Key) ([]netip.Prefix, error) {
if prefixes, ok := c.allowedIPs.get(peerKey); ok {
return prefixes, nil
}
ipcStr, err := c.device.IpcGet()
if err != nil {
return nil, fmt.Errorf("get IPC config: %w", err)
}
stats, err := parseStatus(c.deviceName, ipcStr)
if err != nil {
return nil, fmt.Errorf("parse IPC config: %w", err)
}
// parseStatus reports keys in their textual form, so the comparison needs it once.
wanted := peerKey.String()
for _, peer := range stats.Peers {
if peer.PublicKey != wanted {
continue
}
prefixes := ipNetsToPrefixes(peer.AllowedIPs)
c.allowedIPs.set(peerKey, prefixes)
return prefixes, nil
}
return nil, ErrPeerNotFound
}
func (c *WGUSPConfigurer) FullStats() (*Stats, error) {
@@ -0,0 +1,318 @@
package configurer
import (
"net"
"net/netip"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
wgconn "golang.zx2c4.com/wireguard/conn"
wgdevice "golang.zx2c4.com/wireguard/device"
"golang.zx2c4.com/wireguard/tun/tuntest"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/netbirdio/netbird/client/iface/bind"
)
// newTestUSPConfigurer builds a configurer over a real wireguard-go device backed by an
// in-memory TUN. The device stays down, so no socket is opened and no privileges are needed.
func newTestUSPConfigurer(t *testing.T) *WGUSPConfigurer {
t.Helper()
tun := tuntest.NewChannelTUN()
dev := wgdevice.NewDevice(tun.TUN(), wgconn.NewDefaultBind(), wgdevice.NewLogger(wgdevice.LogLevelSilent, ""))
t.Cleanup(dev.Close)
c := NewUSPConfigurerNoUAPI(dev, "wgtest0", bind.NewActivityRecorder())
key, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate device private key")
require.NoError(t, c.ConfigureInterface(key.String(), 0), "configure test device")
return c
}
// seedPeers adds count peers, each with a /32 overlay address, and returns their public keys.
func seedPeers(t *testing.T, c *WGUSPConfigurer, count int) []string {
t.Helper()
keys := make([]string, 0, count)
for i := 0; i < count; i++ {
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
pub := priv.PublicKey().String()
addr := netip.PrefixFrom(netip.AddrFrom4([4]byte{100, 64, byte(i >> 8), byte(i)}), 32)
require.NoError(t, c.UpdatePeer(pub, []netip.Prefix{addr}, 25*time.Second, nil, nil), "add peer")
keys = append(keys, pub)
}
return keys
}
func peerAllowedIPs(t *testing.T, c *WGUSPConfigurer, peerKey string) []string {
t.Helper()
stats, err := c.FullStats()
require.NoError(t, err, "read device stats")
for _, p := range stats.Peers {
if p.PublicKey != peerKey {
continue
}
got := make([]string, 0, len(p.AllowedIPs))
for _, ipNet := range p.AllowedIPs {
got = append(got, ipNet.String())
}
return got
}
t.Fatalf("peer %s not found on device", peerKey)
return nil
}
// TestRemoveEndpointAddressPreservesRoutedAllowedIPs covers the prefixes the route manager
// attaches to a routing peer through AddAllowedIP. Those are not known to the peer.Conn that
// triggers the endpoint removal, so dropping them here would silently blackhole every route
// behind that peer on each relay or ICE disconnect.
func TestRemoveEndpointAddressPreservesRoutedAllowedIPs(t *testing.T) {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, 3)[1]
routed := []netip.Prefix{
netip.MustParsePrefix("10.20.0.0/16"),
netip.MustParsePrefix("192.168.7.0/24"),
}
for _, prefix := range routed {
require.NoError(t, c.AddAllowedIP(peerKey, prefix), "add routed prefix")
}
before := peerAllowedIPs(t, c, peerKey)
require.Len(t, before, 3, "peer should hold its overlay address plus both routed prefixes")
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey),
"allowed IPs must survive the endpoint removal unchanged")
}
// TestRemoveEndpointAddressDoesNotScaleWithPeerCount is the regression guard for the actual
// defect: clearing one peer's endpoint used to dump and parse the whole device, so its cost
// grew with the size of the network map. On a routing peer with thousands of peers that dump
// runs on every relay and ICE transition, under the interface lock.
func TestRemoveEndpointAddressDoesNotScaleWithPeerCount(t *testing.T) {
measure := func(peerCount int) float64 {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, peerCount)[peerCount/2]
return testing.AllocsPerRun(5, func() {
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
})
}
small := measure(64)
large := measure(1024)
assert.Less(t, large, small*2,
"clearing one endpoint allocated %.0f objects with 1024 peers against %.0f with 64: the cost still scales with the peer count",
large, small)
}
// TestRemoveEndpointAddressFallsBackToDevice covers a peer the store never saw, which is what
// an out-of-band reconfiguration of the device leaves behind. The device stays the source of
// truth in that case, so the allowed IPs must still be preserved.
func TestRemoveEndpointAddressFallsBackToDevice(t *testing.T) {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, 3)[1]
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix")
before := peerAllowedIPs(t, c, peerKey)
c.allowedIPs.reset()
require.NoError(t, c.RemoveEndpointAddress(peerKey), "remove endpoint address")
assert.ElementsMatch(t, before, peerAllowedIPs(t, c, peerKey),
"allowed IPs recovered from the device must be preserved")
recovered, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
assert.True(t, ok, "the fallback must seed the store so the next call skips the device dump")
assert.Len(t, recovered, 2, "seeded prefixes")
}
func TestRemoveAllowedIPKeepsTheOtherPrefixes(t *testing.T) {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, 3)[0]
routed := netip.MustParsePrefix("10.20.0.0/16")
require.NoError(t, c.AddAllowedIP(peerKey, routed), "add routed prefix")
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("192.168.7.0/24")), "add routed prefix")
require.NoError(t, c.RemoveAllowedIP(peerKey, routed), "remove routed prefix")
assert.ElementsMatch(t, []string{"100.64.0.0/32", "192.168.7.0/24"}, peerAllowedIPs(t, c, peerKey),
"only the removed prefix should be gone")
assert.ErrorIs(t, c.RemoveAllowedIP(peerKey, routed), ErrAllowedIPNotFound,
"removing a prefix that is no longer configured must be reported")
}
// TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt covers the lazy connection window documented
// in #6863: AddAllowedIP is update-only, a silent no-op when the peer is absent, so it must not
// leave the store claiming prefixes the device never took. RemoveEndpointAddress re-adds a peer
// without update-only, so a phantom entry would create a peer the device had dropped, and a
// created peer would steal those allowed IPs from whichever peer legitimately holds them.
func TestAddAllowedIPOnAbsentPeerDoesNotResurrectIt(t *testing.T) {
c := newTestUSPConfigurer(t)
seedPeers(t, c, 2)
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
absent := priv.PublicKey().String()
require.NoError(t, c.AddAllowedIP(absent, netip.MustParsePrefix("10.20.0.0/16")),
"update-only add on an absent peer is a silent no-op")
stats, err := c.FullStats()
require.NoError(t, err, "read device stats")
require.Len(t, stats.Peers, 2, "the absent peer must not have been created by AddAllowedIP")
assert.ErrorIs(t, c.RemoveEndpointAddress(absent), ErrPeerNotFound,
"clearing the endpoint of a peer the device does not have must fail")
stats, err = c.FullStats()
require.NoError(t, err, "read device stats")
assert.Len(t, stats.Peers, 2, "no peer may be created while clearing an endpoint")
}
// TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer covers WireGuard's rule that an
// allowed IP belongs to exactly one peer: configuring a prefix on a peer takes it away from
// whichever peer held it before. UpdatePeer relies on that rule rather than removing the prefix
// from the previous holder itself, so a prefix handed over between peers must not come back.
func TestRemoveEndpointAddressDoesNotStealAPrefixFromAnotherPeer(t *testing.T) {
c := newTestUSPConfigurer(t)
keys := seedPeers(t, c, 2)
peerA, peerB := keys[0], keys[1]
routed := netip.MustParsePrefix("10.20.0.0/16")
require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A")
require.Contains(t, peerAllowedIPs(t, c, peerA), routed.String(), "A must hold the prefix")
// The route moves to B. The device takes it away from A on its own.
require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B")
require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix")
require.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(), "the device must have taken it from A")
require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint")
assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(),
"clearing A's endpoint must not take the prefix back from B")
assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(),
"B must still hold the prefix")
}
// TestPresharedKeyCreatedPeerTakesPartInPrefixHandover covers a peer created by a preshared
// key write rather than by a peer update. Rosenpass applies a peer's first key without
// updateOnly, which creates the peer on the device, so a store that ignored that operation
// would treat the peer as unknown and would not account for a prefix later handed over to it.
func TestPresharedKeyCreatedPeerTakesPartInPrefixHandover(t *testing.T) {
c := newTestUSPConfigurer(t)
peerA := seedPeers(t, c, 1)[0]
routed := netip.MustParsePrefix("10.20.0.0/16")
require.NoError(t, c.AddAllowedIP(peerA, routed), "give the prefix to A")
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
peerB := priv.PublicKey().String()
psk, err := wgtypes.GenerateKey()
require.NoError(t, err, "generate preshared key")
require.NoError(t, c.SetPresharedKey(peerB, psk, false), "a first key creates the peer")
require.NoError(t, c.AddAllowedIP(peerB, routed), "hand the prefix over to B")
require.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must hold the prefix")
require.NoError(t, c.RemoveEndpointAddress(peerA), "clear A's endpoint")
assert.NotContains(t, peerAllowedIPs(t, c, peerA), routed.String(),
"clearing A's endpoint must not take the prefix back from B")
assert.Contains(t, peerAllowedIPs(t, c, peerB), routed.String(), "B must still hold the prefix")
}
// TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice is the end to end form of the
// conversion: a v4-mapped prefix must not reach the device as a zero length allowed IP,
// which would route every v4 address to that peer.
func TestUpdatePeerDoesNotWidenAMappedPrefixOnTheDevice(t *testing.T) {
c := newTestUSPConfigurer(t)
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
peerKey := priv.PublicKey().String()
mapped := netip.MustParsePrefix("::ffff:10.1.2.3/112")
require.NoError(t, c.UpdatePeer(peerKey, []netip.Prefix{mapped}, 25*time.Second, nil, nil), "add peer")
onDevice := peerAllowedIPs(t, c, peerKey)
assert.NotContains(t, onDevice, "0.0.0.0/0", "the device must not be given a catch-all allowed IP")
assert.Equal(t, []string{"10.1.0.0/16"}, onDevice, "the device holds the normalized prefix")
recorded, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
require.True(t, ok, "the peer must be recorded")
require.Len(t, recorded, 1, "one prefix recorded")
assert.Equal(t, onDevice[0], recorded[0].String(), "device and store must agree")
}
// TestUpdatePeerWithAnUnusableEndpointTouchesNothing pins the ordering: the endpoint is
// parsed before the device is configured, so a failure cannot leave the device holding a
// peer that the store never learned about, with the prefix handover skipped along with it.
func TestUpdatePeerWithAnUnusableEndpointTouchesNothing(t *testing.T) {
c := newTestUSPConfigurer(t)
seedPeers(t, c, 2)
priv, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err, "generate peer private key")
peerKey := priv.PublicKey().String()
// A three byte address has no textual form netip can parse back.
endpoint := &net.UDPAddr{IP: net.IP{1, 2, 3}, Port: 51820}
require.Error(t, c.UpdatePeer(peerKey, []netip.Prefix{netip.MustParsePrefix("10.30.0.0/16")},
25*time.Second, endpoint, nil), "an unusable endpoint must fail the update")
stats, err := c.FullStats()
require.NoError(t, err, "read device stats")
assert.Len(t, stats.Peers, 2, "the peer must not have reached the device")
_, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
assert.False(t, ok, "the peer must not have been recorded either")
}
// TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses covers a removal that never reached the
// device. A single peer removal is one write, so a failure leaves the peer on the device
// exactly as it was, and the record still describes it; dropping it would only force the
// next caller to read the whole device back for an answer it already had.
func TestRemovePeerKeepsTheRecordWhenTheDeviceRefuses(t *testing.T) {
c := newTestUSPConfigurer(t)
peerKey := seedPeers(t, c, 1)[0]
require.NoError(t, c.AddAllowedIP(peerKey, netip.MustParsePrefix("10.20.0.0/16")), "add routed prefix")
before, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
require.True(t, ok, "the peer must be recorded before the removal")
require.Len(t, before, 2, "overlay address plus routed prefix")
// A closed device refuses every write, which is the shape of any failed removal.
c.device.Close()
require.Error(t, c.RemovePeer(peerKey), "the removal must report the failure")
after, ok := c.allowedIPs.get(mustParseKey(t, peerKey))
require.True(t, ok, "a peer still on the device must stay recorded")
assert.Equal(t, before, after, "the record must describe the peer the device kept")
}
// mustParseKey turns the textual key the configurer API takes into the form the store
// keys on.
func mustParseKey(t *testing.T, key string) wgtypes.Key {
t.Helper()
parsed, err := wgtypes.ParseKey(key)
require.NoError(t, err, "parse peer key")
return parsed
}
+2 -15
View File
@@ -6,27 +6,14 @@ import (
"fmt"
"os/exec"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/wincmd"
)
func (w *WGIface) Destroy() error {
netshCmd := GetSystem32Command("netsh")
netshCmd := wincmd.System32("netsh")
out, err := exec.Command(netshCmd, "interface", "set", "interface", w.Name(), "admin=disable").CombinedOutput()
if err != nil {
return fmt.Errorf("failed to remove interface %s: %w - %s", w.Name(), err, out)
}
return nil
}
// GetSystem32Command checks if a command can be found in the system path and returns it. In case it can't find it
// in the path it will return the full path of a command assuming C:\windows\system32 as the base path.
func GetSystem32Command(command string) string {
_, err := exec.LookPath(command)
if err == nil {
return command
}
log.Tracef("Command %s not found in PATH, using C:\\windows\\system32\\%s.exe path", command, command)
return "C:\\windows\\system32\\" + command + ".exe"
}
+42
View File
@@ -0,0 +1,42 @@
package daemonaddr
import (
"os"
"strconv"
log "github.com/sirupsen/logrus"
)
const (
// EnvMaxRecvMsgSize overrides the default gRPC max receive message size for
// connections to the daemon. Value is in bytes.
EnvMaxRecvMsgSize = "NB_DAEMON_GRPC_MAX_MSG_SIZE"
// defaultMaxRecvMsgSize is the max gRPC receive message size used for daemon
// connections when EnvMaxRecvMsgSize is unset or invalid. It overrides the
// gRPC library default of 4 MB, which a detailed status already exceeds on a
// network of a few thousand peers.
defaultMaxRecvMsgSize = 1024 * 1024 * 16
)
// MaxRecvMsgSize returns the max gRPC receive message size for daemon connections
// from the environment, or defaultMaxRecvMsgSize (16 MB) if unset or invalid.
func MaxRecvMsgSize() int {
val := os.Getenv(EnvMaxRecvMsgSize)
if val == "" {
return defaultMaxRecvMsgSize
}
size, err := strconv.Atoi(val)
if err != nil {
log.Warnf("invalid %s value %q, using default: %v", EnvMaxRecvMsgSize, val, err)
return defaultMaxRecvMsgSize
}
if size <= 0 {
log.Warnf("invalid %s value %d, must be positive, using default", EnvMaxRecvMsgSize, size)
return defaultMaxRecvMsgSize
}
return size
}
+112
View File
@@ -0,0 +1,112 @@
package daemonaddr
import (
"context"
"net"
"os"
"strings"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/proto"
)
func TestMaxRecvMsgSize(t *testing.T) {
tests := []struct {
name string
envValue string
expected int
}{
{name: "unset returns default", envValue: "", expected: defaultMaxRecvMsgSize},
{name: "non-numeric returns default", envValue: "abc", expected: defaultMaxRecvMsgSize},
{name: "negative returns default", envValue: "-1", expected: defaultMaxRecvMsgSize},
{name: "zero returns default", envValue: "0", expected: defaultMaxRecvMsgSize},
{name: "valid value is used", envValue: "33554432", expected: 33554432},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
// Set first so the previous value is restored on cleanup, then unset to
// exercise the absent case.
t.Setenv(EnvMaxRecvMsgSize, tc.envValue)
if tc.envValue == "" {
require.NoError(t, os.Unsetenv(EnvMaxRecvMsgSize), "unset the override")
}
assert.Equal(t, tc.expected, MaxRecvMsgSize(), "max receive message size")
})
}
}
// bigStatusServer answers Status with a response larger than gRPC's 4 MB default
// receive limit, which is what a detailed status on a large network looks like.
type bigStatusServer struct {
proto.UnimplementedDaemonServiceServer
payload string
}
func (s *bigStatusServer) Status(context.Context, *proto.StatusRequest) (*proto.StatusResponse, error) {
return &proto.StatusResponse{Status: s.payload}, nil
}
func startBigStatusServer(t *testing.T, payload string) string {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err, "listen on loopback")
srv := grpc.NewServer()
proto.RegisterDaemonServiceServer(srv, &bigStatusServer{payload: payload})
go func() {
_ = srv.Serve(listener)
}()
t.Cleanup(srv.Stop)
return "tcp://" + listener.Addr().String()
}
func TestDialTargetAcceptsAStatusOverTheGrpcDefault(t *testing.T) {
payload := strings.Repeat("x", 5*1024*1024)
addr := startBigStatusServer(t, payload)
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
target, opts := DialTarget(addr)
conn, err := grpc.NewClient(target, opts...)
require.NoError(t, err, "dial the daemon")
t.Cleanup(func() { _ = conn.Close() })
resp, err := proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{})
require.NoError(t, err, "a detailed status must not be rejected for its size")
assert.Len(t, resp.GetStatus(), len(payload), "the whole response must arrive")
}
// TestDialTargetRaisesTheDefaultLimit is the negative control: the same response
// over a connection carrying gRPC's own defaults is refused, which is the failure
// reported by `netbird status -d` on a large deployment.
func TestDialTargetRaisesTheDefaultLimit(t *testing.T) {
payload := strings.Repeat("x", 5*1024*1024)
addr := startBigStatusServer(t, payload)
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
conn, err := grpc.NewClient(
strings.TrimPrefix(addr, "tcp://"),
grpc.WithTransportCredentials(insecure.NewCredentials()),
)
require.NoError(t, err, "dial with the library defaults")
t.Cleanup(func() { _ = conn.Close() })
_, err = proto.NewDaemonServiceClient(conn).Status(ctx, &proto.StatusRequest{})
require.Error(t, err, "the library default must reject this response")
assert.Equal(t, codes.ResourceExhausted, status.Code(err), "gRPC rejects an oversized message")
}
+4 -1
View File
@@ -36,7 +36,10 @@ const (
// 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())}
opts := []grpc.DialOption{
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithDefaultCallOptions(grpc.MaxCallRecvMsgSize(MaxRecvMsgSize())),
}
if name, ok := strings.CutPrefix(addr, pipeScheme); ok {
paths := PipePaths(name)
+11
View File
@@ -6,6 +6,17 @@ import (
"path/filepath"
)
// CheckOnlyOwnerWritable reports an error unless path, and every directory
// leading to it, is owned by an account that can already act with the privileges
// the caller holds, and is writable by nobody else.
//
// Exported for callers outside elevation that read a file while privileged and
// then act on what it says: the same question this package asks of an
// executable, asked of a configuration file.
func CheckOnlyOwnerWritable(path string) error {
return checkOnlyOwnerWritable(path)
}
// trustedSelf returns the path of this executable, provided it is one we are
// willing to have run as root.
//
+5 -1
View File
@@ -1040,7 +1040,11 @@ func (e *Engine) handleSync(update *mgmProto.SyncResponse) error {
// 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)
// With the firewall disabled there is no ACL manager to program, so
// RoutesFirewallRules would be built and then dropped. On a peer that
// routes many network resources that is the single most expensive
// step of the sync.
result, err := nbnetworkmap.EnvelopeToNetworkMap(e.ctx, envelope, localKey, dnsName, e.config.DisableFirewall)
if err != nil {
return fmt.Errorf("decode network map envelope: %w", err)
}
@@ -0,0 +1,76 @@
package profilemanager
import (
"fmt"
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// Regression test: a concurrent Get and Set of the ActiveProfileState will
// fail on Windows since the write is a temp file renamed over an open file.
// Windows will refuse to replace a file another handle holds open by default.
func TestActiveProfileState_ReadsDoNotBreakAConcurrentWrite(t *testing.T) {
withTempConfigDir(t, func(configDir string) {
withPatchedGlobals(t, configDir, func() {
sm := &ServiceManager{}
require.NoError(t, sm.CreateDefaultProfile())
require.NoError(t, sm.SetActiveProfileStateToDefault())
const switched = ID("0123456789abcdef0123456789abcdef")
const rounds = 50
var wg sync.WaitGroup
errs := make(chan error, 128)
for i := 0; i < 8; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for r := 0; r < rounds; r++ {
state, err := sm.GetActiveProfileState()
if err != nil {
errs <- fmt.Errorf("read: %w", err)
return
}
if state.ID != defaultProfileName && state.ID != switched {
errs <- fmt.Errorf("read: active profile is %q, which no writer wrote", state.ID)
return
}
}
}()
}
for i := 0; i < 2; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for r := 0; r < rounds; r++ {
id := switched
if r%2 == 0 {
id = defaultProfileName
}
if err := sm.SetActiveProfileState(&ActiveProfileState{ID: id, Username: "testuser"}); err != nil {
errs <- fmt.Errorf("switch: %w", err)
return
}
}
}()
}
wg.Wait()
close(errs)
for err := range errs {
assert.NoError(t, err, "a switch and a read of the active profile state must not collide")
}
state, err := sm.GetActiveProfileState()
require.NoError(t, err)
assert.Contains(t, []ID{defaultProfileName, switched}, state.ID,
"the file holds whichever switch landed last, not a mix of the two")
})
})
}
@@ -0,0 +1,30 @@
// Package wincmd locates the Windows utilities the client shells out to.
package wincmd
import (
"path/filepath"
log "github.com/sirupsen/logrus"
"golang.org/x/sys/windows"
)
// defaultSystem32Dir is where the system directory is on every supported
// install, used only when the API that reports it fails.
const defaultSystem32Dir = `C:\Windows\System32`
// System32 returns the full path of a Windows utility under the system
// directory.
//
// PATH is deliberately not consulted. The daemon runs as LocalSystem with an
// environment of its own, so whoever can place an entry in that PATH chooses
// which binary runs with those privileges. The system directory is read from
// the API rather than from %SystemRoot% for the same reason.
func System32(command string) string {
sysDir, err := windows.GetSystemDirectory()
if err != nil {
log.Warnf("Failed to locate the Windows system directory, falling back to %s: %v", defaultSystem32Dir, err)
sysDir = defaultSystem32Dir
}
return filepath.Join(sysDir, command+".exe")
}
@@ -0,0 +1,31 @@
package wincmd
import (
"os"
"path/filepath"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestSystem32IgnoresPATH(t *testing.T) {
// A directory holding something that would win a PATH lookup, in front of
// everything else: the daemon runs as LocalSystem, so a PATH entry must not
// be able to decide what it executes.
planted := t.TempDir()
require.NoError(t, os.WriteFile(filepath.Join(planted, "netsh.exe"), []byte("not really netsh"), 0o600))
t.Setenv("PATH", planted+string(os.PathListSeparator)+os.Getenv("PATH"))
got := System32("netsh")
assert.True(t, filepath.IsAbs(got), "the path must be absolute, got %q", got)
assert.NotContains(t, got, planted, "a PATH entry must not be consulted")
assert.True(t, strings.EqualFold(filepath.Base(got), "netsh.exe"), "unexpected file name in %q", got)
// The system directory is what Windows reports it to be, not %SystemRoot%,
// which the same caller could have set alongside PATH.
t.Setenv("SystemRoot", planted)
assert.Equal(t, got, System32("netsh"), "%SystemRoot% must not move the lookup")
}
@@ -1,6 +1,6 @@
import * as DropdownMenuPrimitive from "@radix-ui/react-dropdown-menu";
import { cva } from "class-variance-authority";
import { Check, ChevronRight, Circle } from "lucide-react";
import { Check, ChevronRight } from "lucide-react";
import * as React from "react";
import { cn } from "@/lib/cn";
@@ -159,19 +159,23 @@ const DropdownMenuRadioItem = React.forwardRef<
<DropdownMenuPrimitive.RadioItem
ref={ref}
className={cn(
"relative flex cursor-default select-none items-center rounded-sm py-1.5 pl-8 pr-2 text-sm outline-none",
"text-nb-gray-200 transition-colors hover:bg-nb-gray-900 hover:text-nb-gray-50 focus-visible:bg-nb-gray-900 focus-visible:text-nb-gray-50",
"my-0.5 flex cursor-default select-none items-center gap-2 rounded-md px-2 py-2 outline-none",
"text-xs font-semibold text-nb-gray-200 transition-colors",
"data-[highlighted]:bg-nb-gray-850 data-[highlighted]:text-nb-gray-50",
"data-[disabled]:pointer-events-none data-[disabled]:opacity-50",
className,
)}
{...props}
>
<span className={"absolute left-2 flex h-3.5 w-3.5 items-center justify-center"}>
{children}
<span
aria-hidden={"true"}
className={"ml-auto flex w-4 shrink-0 items-center justify-center"}
>
<DropdownMenuPrimitive.ItemIndicator>
<Circle className={"h-2 w-2 fill-current"} />
<Check size={14} className={"text-netbird"} />
</DropdownMenuPrimitive.ItemIndicator>
</span>
{children}
</DropdownMenuPrimitive.RadioItem>
));
DropdownMenuRadioItem.displayName = DropdownMenuPrimitive.RadioItem.displayName;
@@ -89,7 +89,11 @@ export function LanguagePicker() {
tabIndex={0}
disabled={busy || languages.length === 0}
onKeyDown={handleTriggerKeyDown}
aria-label={t("settings.general.language.label")}
aria-label={
current
? `${t("settings.general.language.label")}: ${labelFor(current)}`
: t("settings.general.language.label")
}
aria-haspopup={"listbox"}
aria-expanded={open}
className={cn(
@@ -1,18 +1,10 @@
import { useState } from "react";
import { useTranslation } from "react-i18next";
import { ChevronDown, MonitorIcon, MoonIcon, SunMediumIcon, type LucideIcon } from "lucide-react";
import {
DropdownMenu,
DropdownMenuContent,
DropdownMenuRadioGroup,
DropdownMenuRadioItem,
DropdownMenuTrigger,
} from "@/components/DropdownMenu";
import { MonitorIcon, MoonIcon, SunMediumIcon, type LucideIcon } from "lucide-react";
import { Select } from "@/components/inputs/Select";
import { HelpText } from "@/components/typography/HelpText";
import { Label } from "@/components/typography/Label";
import { useTheme, type ThemePreference } from "@/contexts/ThemeContext";
import { useFocusVisible } from "@/hooks/useFocusVisible";
import { cn } from "@/lib/cn";
import { errorDialog, formatErrorMessage } from "@/lib/errors";
const OPTIONS: { value: ThemePreference; icon: LucideIcon; labelKey: string }[] = [
@@ -25,16 +17,12 @@ export function ThemePicker() {
const { t } = useTranslation();
const { theme, setTheme } = useTheme();
const [busy, setBusy] = useState(false);
const isFocusVisible = useFocusVisible();
const current = OPTIONS.find((o) => o.value === theme) ?? OPTIONS[0];
const CurrentIcon = current.icon;
const select = async (value: string) => {
const select = async (value: ThemePreference) => {
if (busy || value === theme) return;
setBusy(true);
try {
await setTheme(value as ThemePreference);
await setTheme(value);
} catch (e) {
await errorDialog({
Title: t("settings.error.saveTitle"),
@@ -52,57 +40,17 @@ export function ThemePicker() {
<HelpText margin={false}>{t("settings.general.theme.help")}</HelpText>
</div>
<div className={"shrink-0"}>
<DropdownMenu>
<DropdownMenuTrigger asChild>
<button
type={"button"}
tabIndex={0}
disabled={busy}
aria-label={t("settings.general.theme.label")}
className={cn(
"inline-flex h-[40px] min-w-[160px] items-center gap-2 px-3",
"rounded-md border bg-white dark:bg-nb-gray-900",
"border-neutral-200 dark:border-nb-gray-700",
"cursor-default text-xs font-semibold text-nb-gray-100 outline-none",
"hover:border-nb-gray-700 data-[state=open]:border-nb-gray-700 dark:hover:border-nb-gray-600 dark:data-[state=open]:border-nb-gray-600",
isFocusVisible &&
"focus-visible:ring-2 focus-visible:ring-nb-gray-50/60 focus-visible:ring-offset-2 focus-visible:ring-offset-nb-gray-940",
"disabled:opacity-50",
)}
>
<CurrentIcon
size={16}
aria-hidden={"true"}
className={"shrink-0 text-nb-gray-200"}
/>
<span className={"flex-1 truncate text-left"}>
{t(current.labelKey)}
</span>
<ChevronDown
size={12}
aria-hidden={"true"}
className={"shrink-0 text-nb-gray-400"}
/>
</button>
</DropdownMenuTrigger>
<DropdownMenuContent
align={"end"}
className={"w-[var(--radix-dropdown-menu-trigger-width)]"}
>
<DropdownMenuRadioGroup value={theme} onValueChange={(v) => void select(v)}>
{OPTIONS.map(({ value, icon: Icon, labelKey }) => (
<DropdownMenuRadioItem key={value} value={value}>
<Icon
size={14}
aria-hidden={"true"}
className={"mr-2 shrink-0 text-nb-gray-300"}
/>
{t(labelKey)}
</DropdownMenuRadioItem>
))}
</DropdownMenuRadioGroup>
</DropdownMenuContent>
</DropdownMenu>
<Select
value={theme}
options={OPTIONS.map(({ value, icon, labelKey }) => ({
value,
icon,
label: t(labelKey),
}))}
onChange={(v) => void select(v)}
ariaLabel={t("settings.general.theme.label")}
disabled={busy}
/>
</div>
</div>
);
@@ -147,7 +147,8 @@ export const Button = forwardRef<HTMLButtonElement, ButtonProps>(function Button
ref={ref}
type={type}
tabIndex={0}
disabled={disabled || loading}
disabled={disabled}
aria-disabled={loading || undefined}
aria-busy={loading || undefined}
className={cn(
buttonVariants({
@@ -156,10 +157,15 @@ export const Button = forwardRef<HTMLButtonElement, ButtonProps>(function Button
border: border ? 1 : 0,
size,
}),
loading && "pointer-events-none",
className,
)}
onClick={(e) => {
if (stopPropagation) e.stopPropagation();
if (loading) {
e.preventDefault();
return;
}
if (copy !== undefined) {
void navigator.clipboard
.writeText(copy)
@@ -14,6 +14,7 @@ type ConfirmModalProps = {
cancelLabel?: string;
danger?: boolean;
busy?: boolean;
cancellable?: boolean;
onConfirm: () => void;
onCancel: () => void;
};
@@ -26,11 +27,13 @@ export const ConfirmModal = ({
cancelLabel,
danger = false,
busy = false,
cancellable,
onConfirm,
onCancel,
}: ConfirmModalProps) => {
const { t } = useTranslation();
const resolvedCancel = cancelLabel ?? t("common.cancel");
const canCancel = cancellable ?? !busy;
const srTitle = typeof title === "string" ? title : undefined;
const srDescription = typeof description === "string" ? description : undefined;
@@ -39,7 +42,7 @@ export const ConfirmModal = ({
<Dialog.Root
open={open}
onOpenChange={(next) => {
if (!next && !busy) onCancel();
if (!next && canCancel) onCancel();
}}
>
<Dialog.Content
@@ -62,7 +65,7 @@ export const ConfirmModal = ({
<Button
variant={"secondary"}
size={"sm"}
disabled={busy}
disabled={!canCancel}
onClick={onCancel}
>
{resolvedCancel}
@@ -71,7 +74,7 @@ export const ConfirmModal = ({
autoFocus
variant={danger ? "danger" : "primary"}
size={"sm"}
disabled={busy}
loading={busy}
onClick={onConfirm}
>
{confirmLabel}
@@ -0,0 +1,99 @@
import type { LucideIcon } from "lucide-react";
import { ChevronDown } from "lucide-react";
import {
DropdownMenu,
DropdownMenuContent,
DropdownMenuRadioGroup,
DropdownMenuRadioItem,
DropdownMenuTrigger,
} from "@/components/DropdownMenu";
import { useFocusVisible } from "@/hooks/useFocusVisible";
import { cn } from "@/lib/cn";
export type SelectOption<T extends string> = {
value: T;
label: string;
icon?: LucideIcon;
};
type SelectProps<T extends string> = {
value: T;
options: SelectOption<T>[];
onChange: (value: T) => void;
ariaLabel: string;
disabled?: boolean;
className?: string;
};
export function Select<T extends string>({
value,
options,
onChange,
ariaLabel,
disabled,
className,
}: SelectProps<T>) {
const isFocusVisible = useFocusVisible();
const current = options.find((o) => o.value === value) ?? options[0];
const CurrentIcon = current?.icon;
return (
<DropdownMenu>
<DropdownMenuTrigger asChild>
<button
type={"button"}
tabIndex={0}
disabled={disabled}
aria-label={current ? `${ariaLabel}: ${current.label}` : ariaLabel}
className={cn(
"inline-flex h-[40px] min-w-[160px] items-center gap-2 px-3",
"rounded-md border bg-white dark:bg-nb-gray-900",
"border-neutral-200 dark:border-nb-gray-700",
"cursor-default text-xs font-semibold text-nb-gray-100 outline-none",
"hover:border-nb-gray-700 data-[state=open]:border-nb-gray-700 dark:hover:border-nb-gray-600 dark:data-[state=open]:border-nb-gray-600",
isFocusVisible &&
"focus-visible:ring-2 focus-visible:ring-nb-gray-50/60 focus-visible:ring-offset-2 focus-visible:ring-offset-nb-gray-940",
"disabled:opacity-50",
className,
)}
>
{CurrentIcon && (
<CurrentIcon
size={16}
aria-hidden={"true"}
className={"shrink-0 text-nb-gray-200"}
/>
)}
<span className={"flex-1 truncate text-left"}>{current?.label ?? "—"}</span>
<ChevronDown
size={12}
aria-hidden={"true"}
className={"shrink-0 text-nb-gray-400"}
/>
</button>
</DropdownMenuTrigger>
<DropdownMenuContent
align={"start"}
sideOffset={6}
className={
"w-[var(--radix-dropdown-menu-trigger-width)] border-nb-gray-850 bg-nb-gray-920"
}
>
<DropdownMenuRadioGroup value={value} onValueChange={(v) => onChange(v as T)}>
{options.map(({ value: optionValue, label, icon: Icon }) => (
<DropdownMenuRadioItem key={optionValue} value={optionValue}>
{Icon && (
<Icon
size={14}
aria-hidden={"true"}
className={"shrink-0 text-nb-gray-300"}
/>
)}
<span className={"min-w-0 flex-1 truncate"}>{label}</span>
</DropdownMenuRadioItem>
))}
</DropdownMenuRadioGroup>
</DropdownMenuContent>
</DropdownMenu>
);
}
@@ -8,6 +8,28 @@ import {
useState,
} from "react";
import { ConfirmModal } from "@/components/dialog/ConfirmModal";
import i18next from "@/lib/i18n";
// Nothing on the daemon path carries a deadline, so a hung call would leave the
// modal spinning with no way out. Cancel comes back once the wait stops looking
// normal, and the wait is abandoned entirely at the deadline.
const CANCELLABLE_AFTER_MS = 15_000;
const TIMEOUT_MS = 30_000;
const withTimeout = async (action: () => Promise<unknown>) => {
let timer: ReturnType<typeof setTimeout> | undefined;
const expiry = new Promise<never>((_, reject) => {
timer = setTimeout(
() => reject(new Error(i18next.t("error.daemon_unreachable"))),
TIMEOUT_MS,
);
});
try {
await Promise.race([action(), expiry]);
} finally {
clearTimeout(timer);
}
};
export type ConfirmOptions = {
title: ReactNode;
@@ -15,6 +37,7 @@ export type ConfirmOptions = {
confirmLabel: string;
cancelLabel?: string;
danger?: boolean;
onConfirm?: () => Promise<unknown>;
};
type DialogContextValue = {
@@ -23,23 +46,50 @@ type DialogContextValue = {
const DialogContext = createContext<DialogContextValue | null>(null);
type Settler = { resolve: (result: boolean) => void; reject: (reason: unknown) => void };
export function DialogProvider({ children }: Readonly<{ children: ReactNode }>) {
const [open, setOpen] = useState(false);
const [busy, setBusy] = useState(false);
const [stalled, setStalled] = useState(false);
const [options, setOptions] = useState<ConfirmOptions | null>(null);
const resolverRef = useRef<((result: boolean) => void) | null>(null);
const resolverRef = useRef<Settler | null>(null);
const confirm = useCallback((opts: ConfirmOptions) => {
setOptions(opts);
setOpen(true);
return new Promise<boolean>((resolve) => {
resolverRef.current = resolve;
return new Promise<boolean>((resolve, reject) => {
resolverRef.current = { resolve, reject };
});
}, []);
const settle = (result: boolean) => {
resolverRef.current?.(result);
const take = (expected?: Settler | null) => {
const settler = resolverRef.current;
if (expected && settler !== expected) return null;
resolverRef.current = null;
setBusy(false);
setStalled(false);
setOpen(false);
return settler;
};
const handleConfirm = async () => {
const action = options?.onConfirm;
if (!action) {
take()?.resolve(true);
return;
}
const dispatched = resolverRef.current;
setBusy(true);
const stallTimer = setTimeout(() => setStalled(true), CANCELLABLE_AFTER_MS);
try {
await withTimeout(action);
take(dispatched)?.resolve(true);
} catch (e) {
take(dispatched)?.reject(e);
} finally {
clearTimeout(stallTimer);
}
};
const value = useMemo<DialogContextValue>(() => ({ confirm }), [confirm]);
@@ -54,8 +104,10 @@ export function DialogProvider({ children }: Readonly<{ children: ReactNode }>)
confirmLabel={options?.confirmLabel ?? ""}
cancelLabel={options?.cancelLabel}
danger={options?.danger}
onConfirm={() => settle(true)}
onCancel={() => settle(false)}
busy={busy}
cancellable={!busy || stalled}
onConfirm={() => void handleConfirm()}
onCancel={() => take()?.resolve(false)}
/>
</DialogContext.Provider>
);
@@ -78,7 +78,7 @@ export function ProfilesTab() {
return items;
}, [profiles, activeProfileId]);
const guarded = async (title: string, fn: () => Promise<void>) => {
const guarded = async (title: string, fn: () => Promise<unknown>) => {
if (busy) return;
setBusy(true);
try {
@@ -115,14 +115,15 @@ export function ProfilesTab() {
const handleDelete = async (id: string, name: string) => {
if (id === DEFAULT_PROFILE_ID) return;
const ok = await confirm({
title: t("profile.delete.title", { name }),
description: t("profile.delete.message", { name }),
confirmLabel: t("common.delete"),
danger: true,
});
if (!ok) return;
void guarded(i18next.t("profile.error.deleteTitle"), () => removeProfile(id));
await guarded(i18next.t("profile.error.deleteTitle"), () =>
confirm({
title: t("profile.delete.title", { name }),
description: t("profile.delete.message", { name }),
confirmLabel: t("common.delete"),
danger: true,
onConfirm: () => removeProfile(id),
}),
);
};
const handleCreate = async (name: string, managementUrl: string) => {
@@ -1,5 +1,21 @@
import { useCallback, useEffect, useRef, useState } from "react";
import { createRoot } from "react-dom/client";
import netbirdLogo from "@/assets/logos/netbird.svg";
const scratch = new Uint32Array(1);
function random() {
crypto.getRandomValues(scratch);
return scratch[0] / 2 ** 32;
}
type Mask = {
cols: number;
rows: number;
cells: Uint8Array;
seeds: Uint8Array;
glow: Float32Array;
};
export function useAccentTrigger() {
const clicksRef = useRef(0);
@@ -50,24 +66,45 @@ function Accent({ onDone }: Readonly<{ onDone: () => void }>) {
const ctx = canvas.getContext("2d");
if (!ctx) return;
const chars = "DRIBTENMAET".split("").reverse().join("");
let disposed = false;
let mask: Mask | null = null;
const dpr = window.devicePixelRatio || 1;
let columns = 0;
let drops: number[] = [];
let latestBuild = 0;
let rebuild: ReturnType<typeof setTimeout> | undefined;
const resize = () => {
canvas.width = window.innerWidth * dpr;
canvas.height = window.innerHeight * dpr;
canvas.style.width = `${window.innerWidth}px`;
canvas.style.height = `${window.innerHeight}px`;
ctx.setTransform(dpr, 0, 0, dpr, 0, 0);
const next = Math.floor(window.innerWidth / 15);
if (next !== columns) {
columns = next;
drops = Array.from({ length: columns }, () => random() * -60);
mask = null;
}
globalThis.clearTimeout(rebuild);
rebuild = globalThis.setTimeout(() => {
const build = ++latestBuild;
void buildMask().then((m) => {
if (!disposed && build === latestBuild) mask = m;
});
}, 100);
};
resize();
window.addEventListener("resize", resize);
const chars = "TEAMNETBIRD";
const fontSize = 16;
const columns = Math.floor(window.innerWidth / fontSize);
const drops = Array.from({ length: columns }, () => Math.random() * -50);
let raf = 0;
let last = 0;
let frame = 0;
const draw = (t: number) => {
if (t - last > 50) {
last = t;
@@ -77,18 +114,26 @@ function Accent({ onDone }: Readonly<{ onDone: () => void }>) {
ctx.fillRect(0, 0, window.innerWidth, window.innerHeight);
ctx.globalCompositeOperation = "source-over";
ctx.font = `${fontSize}px ui-monospace, monospace`;
ctx.fillStyle = "#f68330";
ctx.font = "15px ui-monospace, monospace";
ctx.textBaseline = "top";
ctx.shadowBlur = 0;
ctx.fillStyle = "rgba(246, 131, 48, 0.5)";
for (let i = 0; i < drops.length; i++) {
const ch = chars[Math.floor(Math.random() * chars.length)];
const y = drops[i] * fontSize;
ctx.fillText(ch, i * fontSize, y);
if (y > window.innerHeight && Math.random() > 0.975) {
drops[i] = 0;
const ch = chars[Math.floor(random() * chars.length)];
const y = drops[i] * 15;
ctx.fillText(ch, i * 15, y);
igniteTrail(mask, i, Math.floor(drops[i]));
if (y > window.innerHeight && random() > 0.86) {
drops[i] = random() * -12;
}
drops[i]++;
}
drawGlow(ctx, mask, frame, chars);
frame++;
}
raf = requestAnimationFrame(draw);
};
@@ -100,8 +145,10 @@ function Accent({ onDone }: Readonly<{ onDone: () => void }>) {
}, 9000);
return () => {
disposed = true;
cancelAnimationFrame(raf);
globalThis.clearTimeout(timeout);
globalThis.clearTimeout(rebuild);
window.removeEventListener("resize", resize);
};
}, [onDone]);
@@ -114,3 +161,103 @@ function Accent({ onDone }: Readonly<{ onDone: () => void }>) {
</div>
);
}
function igniteTrail(mask: Mask | null, col: number, row: number) {
if (!mask || col < 0 || col >= mask.cols) return;
for (let k = 0; k < 6; k++) {
const r = row - k;
if (r < 0 || r >= mask.rows) continue;
const idx = r * mask.cols + col;
if (mask.cells[idx] === 0) continue;
const heat = 1 - k / 6;
if (heat > mask.glow[idx]) mask.glow[idx] = heat;
}
}
function drawGlow(ctx: CanvasRenderingContext2D, mask: Mask | null, frame: number, chars: string) {
if (!mask) return;
for (let idx = 0; idx < mask.cells.length; idx++) {
const heat = fade(mask, idx);
if (heat === 0) continue;
const seed = mask.seeds[idx];
const core = mask.cells[idx] === 2;
ctx.shadowColor = core ? "#f05252" : "#f68330";
ctx.shadowBlur = 10 * heat;
ctx.fillStyle = core ? `rgba(255, 226, 210, ${heat})` : `rgba(255, 255, 255, ${heat})`;
ctx.fillText(
chars[(seed + Math.floor(frame / (3 + (seed % 5)))) % chars.length],
(idx % mask.cols) * 15,
Math.floor(idx / mask.cols) * 15,
);
}
ctx.shadowBlur = 0;
}
function fade(mask: Mask, idx: number) {
if (mask.cells[idx] === 0) return 0;
const heat = mask.glow[idx];
if (heat <= 0.02) {
mask.glow[idx] = 0;
return 0;
}
mask.glow[idx] = heat * 0.94;
return heat;
}
function loadLogo() {
return new Promise<HTMLImageElement>((resolve, reject) => {
const img = new Image();
img.onload = () => resolve(img);
img.onerror = reject;
img.src = netbirdLogo;
});
}
async function buildMask(): Promise<Mask | null> {
const cols = Math.floor(window.innerWidth / 15);
const rows = Math.ceil(window.innerHeight / 15);
if (cols <= 0 || rows <= 0) return null;
let img: HTMLImageElement;
try {
img = await loadLogo();
} catch {
return null;
}
const off = document.createElement("canvas");
off.width = cols;
off.height = rows;
const offCtx = off.getContext("2d", { willReadFrequently: true });
if (!offCtx) return null;
const aspect = (img.naturalWidth || 31) / (img.naturalHeight || 23);
let w = cols * 0.8;
let h = w / aspect;
if (h > rows * 0.8) {
h = rows * 0.8;
w = h * aspect;
}
offCtx.imageSmoothingEnabled = false;
offCtx.drawImage(img, (cols - w) / 2, (rows - h) / 2, w, h);
const { data } = offCtx.getImageData(0, 0, cols, rows);
const cells = new Uint8Array(cols * rows);
const seeds = new Uint8Array(cols * rows);
const glow = new Float32Array(cols * rows);
for (let i = 0; i < cells.length; i++) {
seeds[i] = Math.floor(random() * 251);
const alpha = data[i * 4 + 3];
if (alpha < 64) continue;
const r = data[i * 4];
const g = data[i * 4 + 1];
const b = data[i * 4 + 2];
cells[i] = r > 180 && g < 130 && b < 130 && g <= b + 24 ? 2 : 1;
}
return { cols, rows, cells, seeds, glow };
}
@@ -1,6 +1,15 @@
import { useId, type ReactNode } from "react";
import { Trans, useTranslation } from "react-i18next";
import { ChevronDown, CircleCheckBig, FolderOpen, Info, Loader2 } from "lucide-react";
import {
CircleCheckBig,
FolderOpen,
Info,
Loader2,
Shield,
ShieldCheck,
ShieldOff,
type LucideIcon,
} from "lucide-react";
import { Browser } from "@wailsio/runtime";
import { Debug as DebugSvc } from "@bindings/services";
import type { DebugBundleResult } from "@bindings/services/models.js";
@@ -8,20 +17,13 @@ import { Button } from "@/components/buttons/Button";
import { DialogActions } from "@/components/dialog/DialogActions";
import { DialogDescription } from "@/components/dialog/DialogDescription";
import { DialogHeading } from "@/components/dialog/DialogHeading";
import {
DropdownMenu,
DropdownMenuContent,
DropdownMenuRadioGroup,
DropdownMenuRadioItem,
DropdownMenuTrigger,
} from "@/components/DropdownMenu";
import FancyToggleSwitch from "@/components/switches/FancyToggleSwitch";
import HelpText from "@/components/typography/HelpText.tsx";
import { Input } from "@/components/inputs/Input";
import { Label } from "@/components/typography/Label";
import { Select } from "@/components/inputs/Select";
import { SquareIcon } from "@/components/SquareIcon";
import { Tooltip } from "@/components/Tooltip";
import { cn } from "@/lib/cn";
import { formatRemaining } from "@/lib/formatters";
import type { AnonymizeLevel, DebugStage } from "@/contexts/DebugBundleContext";
import { useDebugBundleContext } from "@/contexts/DebugBundleContext";
@@ -29,6 +31,12 @@ import { SectionGroup, SettingsBottomBar } from "@/modules/settings/SettingsSect
const SUPPORT_DOCS_URL = "https://docs.netbird.io/help/report-bug-issues";
const ANONYMIZE_LEVELS: { value: AnonymizeLevel; icon: LucideIcon }[] = [
{ value: "none", icon: ShieldOff },
{ value: "default", icon: Shield },
{ value: "strict", icon: ShieldCheck },
];
export function SettingsTroubleshooting() {
const { t } = useTranslation();
const durationId = useId();
@@ -89,44 +97,16 @@ export function SettingsTroubleshooting() {
</HelpText>
</div>
<div className={"shrink-0"}>
<DropdownMenu>
<DropdownMenuTrigger asChild>
<button
type={"button"}
aria-label={t("settings.troubleshooting.anonymize.label")}
className={cn(
"inline-flex h-[40px] min-w-[160px] items-center justify-between gap-2 px-3",
"rounded-md border bg-white dark:bg-nb-gray-900",
"border-neutral-200 dark:border-nb-gray-700",
"cursor-default text-xs font-semibold text-nb-gray-100 outline-none",
"hover:border-nb-gray-700 data-[state=open]:border-nb-gray-700 dark:hover:border-nb-gray-600 dark:data-[state=open]:border-nb-gray-600",
)}
>
{t(`settings.troubleshooting.anonymize.${anonymizeLevel}`)}
<ChevronDown
size={16}
aria-hidden={"true"}
className={"shrink-0 text-nb-gray-200"}
/>
</button>
</DropdownMenuTrigger>
<DropdownMenuContent align={"end"} className={"min-w-[160px]"}>
<DropdownMenuRadioGroup
value={anonymizeLevel}
onValueChange={(v) => setAnonymizeLevel(v as AnonymizeLevel)}
>
<DropdownMenuRadioItem value={"none"}>
{t("settings.troubleshooting.anonymize.none")}
</DropdownMenuRadioItem>
<DropdownMenuRadioItem value={"default"}>
{t("settings.troubleshooting.anonymize.default")}
</DropdownMenuRadioItem>
<DropdownMenuRadioItem value={"strict"}>
{t("settings.troubleshooting.anonymize.strict")}
</DropdownMenuRadioItem>
</DropdownMenuRadioGroup>
</DropdownMenuContent>
</DropdownMenu>
<Select
value={anonymizeLevel}
options={ANONYMIZE_LEVELS.map(({ value, icon }) => ({
value,
icon,
label: t(`settings.troubleshooting.anonymize.${value}`),
}))}
onChange={setAnonymizeLevel}
ariaLabel={t("settings.troubleshooting.anonymize.label")}
/>
</div>
</div>
<FancyToggleSwitch
+6 -2
View File
@@ -120,7 +120,7 @@ func execute(cmd *cobra.Command, _ []string) error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
err = shutdownServers(ctx, servers.relaySrv, servers.healthcheck, servers.stunServer, servers.mgmtSrv, servers.metricsServer)
err = shutdownServers(ctx, servers.relaySrv, servers.healthcheck, servers.stunServer, servers.mgmtSrv, servers.signalSrv, servers.metricsServer)
wg.Wait()
return err
}
@@ -399,7 +399,7 @@ func startServers(wg *sync.WaitGroup, srv *relayServer.Server, httpHealthcheck *
}
}
func shutdownServers(ctx context.Context, srv *relayServer.Server, httpHealthcheck *healthcheck.Server, stunServer *stun.Server, mgmtSrv mgmtServer.Server, metricsServer *sharedMetrics.Metrics) error {
func shutdownServers(ctx context.Context, srv *relayServer.Server, httpHealthcheck *healthcheck.Server, stunServer *stun.Server, mgmtSrv mgmtServer.Server, signalSrv *signalServer.Server, metricsServer *sharedMetrics.Metrics) error {
var errs error
if err := httpHealthcheck.Shutdown(ctx); err != nil {
@@ -425,6 +425,10 @@ func shutdownServers(ctx context.Context, srv *relayServer.Server, httpHealthche
}
}
if signalSrv != nil {
signalSrv.Stop()
}
if metricsServer != nil {
log.Infof("shutting down metrics server")
if err := metricsServer.Shutdown(ctx); err != nil {
+5
View File
@@ -26,3 +26,8 @@ window. Restarting management does not extend a previously assigned deadline.
Registrations with existing services, including services using subdomains, are
retained for operator review. Management logs their account and domain IDs so
an operator can identify and resolve those dependencies before cleanup.
Manual deletion is also refused while any service uses the domain or a subdomain,
including disabled services. Delete those services or move them to another domain
before removing the registration. A refused deletion returns HTTP 412 and leaves
the domain and its services unchanged; no deletion activity event is recorded.
+5 -2
View File
@@ -40,6 +40,7 @@ require (
github.com/aws/aws-sdk-go-v2/credentials v1.18.10
github.com/aws/aws-sdk-go-v2/service/s3 v1.87.3
github.com/c-robinson/iplib v1.0.3
github.com/caarlos0/env/v11 v11.4.1
github.com/caddyserver/certmagic v0.21.3
github.com/coder/websocket v1.8.14
github.com/coreos/go-iptables v0.7.0
@@ -67,6 +68,7 @@ require (
github.com/google/gopacket v1.1.19
github.com/google/nftables v0.3.0
github.com/gopacket/gopacket v1.4.0
github.com/grafana/pyroscope-go v1.4.2
github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.0.2-0.20240212192251-757544f21357
github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3
github.com/hashicorp/go-multierror v1.1.1
@@ -236,6 +238,7 @@ require (
github.com/googleapis/gax-go/v2 v2.21.0 // indirect
github.com/goreleaser/chglog v0.7.4 // indirect
github.com/gorilla/handlers v1.5.2 // indirect
github.com/grafana/pyroscope-go/godeltaprof v0.1.11 // indirect
github.com/hashicorp/errwrap v1.1.0 // indirect
github.com/hashicorp/go-cleanhttp v0.5.2 // indirect
github.com/hashicorp/go-retryablehttp v0.7.8 // indirect
@@ -259,7 +262,7 @@ require (
github.com/josharian/intern v1.0.0 // indirect
github.com/kelseyhightower/envconfig v1.4.0 // indirect
github.com/kevinburke/ssh_config v1.4.0 // indirect
github.com/klauspost/compress v1.18.3 // indirect
github.com/klauspost/compress v1.18.7 // indirect
github.com/klauspost/cpuid/v2 v2.3.0 // indirect
github.com/koron/go-ssdp v0.0.4 // indirect
github.com/kr/fs v0.1.0 // indirect
@@ -365,7 +368,7 @@ replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-2
replace github.com/mailru/easyjson => github.com/netbirdio/easyjson v0.9.0
replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260902163841-4a71f7b1d9e1
replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260928122122-d1db5fa35318
tool (
github.com/goreleaser/chglog/cmd/chglog
+10 -4
View File
@@ -106,6 +106,8 @@ github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
github.com/c-robinson/iplib v1.0.3 h1:NG0UF0GoEsrC1/vyfX1Lx2Ss7CySWl3KqqXh3q4DdPU=
github.com/c-robinson/iplib v1.0.3/go.mod h1:i3LuuFL1hRT5gFpBRnEydzw8R6yhGkF4szNDIbF8pgo=
github.com/caarlos0/env/v11 v11.4.1 h1:fYwH0sWEsBSMPG7t4e/PEfTFzrWrpjyygXyUnWiSwEw=
github.com/caarlos0/env/v11 v11.4.1/go.mod h1:qupehSf/Y0TUTsxKywqRt/vJjN5nz6vauiYEUUr8P4U=
github.com/caddyserver/certmagic v0.21.3 h1:pqRRry3yuB4CWBVq9+cUqu+Y6E2z8TswbhNx1AZeYm0=
github.com/caddyserver/certmagic v0.21.3/go.mod h1:Zq6pklO9nVRl3DIFUw9gVUfXKdpc/0qwTUAQMBlfgtI=
github.com/caddyserver/zerossl v0.1.3 h1:onS+pxp3M8HnHpN5MMbOMyNjmTheJyWRaZYwn+YTAyA=
@@ -327,6 +329,10 @@ github.com/gorilla/handlers v1.5.2 h1:cLTUSsNkgcwhgRqvCNmdbRWG0A3N4F+M2nWKdScwyE
github.com/gorilla/handlers v1.5.2/go.mod h1:dX+xVpaxdSw+q0Qek8SSsl3dfMk3jNddUkMzo0GtH0w=
github.com/gorilla/mux v1.8.1 h1:TuBL49tXwgrFYWhqrNgrUNEY92u81SPhu7sTdzQEiWY=
github.com/gorilla/mux v1.8.1/go.mod h1:AKf9I4AEqPTmMytcMc0KkNouC66V3BtZ4qD5fmWSiMQ=
github.com/grafana/pyroscope-go v1.4.2 h1:0LW5HrUJXgGr9zF5gITP/HaFXN9/LsMiwlgVJAK75l0=
github.com/grafana/pyroscope-go v1.4.2/go.mod h1:Ej13Jr05rRJrjWvrrFhfh6gGYXtfibuukOs3Tl3Y7QQ=
github.com/grafana/pyroscope-go/godeltaprof v0.1.11 h1:el5LYpXissAiCKZ5/6yjlr6mhYVV6Cp5lahTocxraXM=
github.com/grafana/pyroscope-go/godeltaprof v0.1.11/go.mod h1:jl1V8M4cWsXciROCPIDDG7CtjSjT/ECbp6eLVuMxYRI=
github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.0.2-0.20240212192251-757544f21357 h1:Fkzd8ktnpOR9h47SXHe2AYPwelXLH2GjGsjlAloiWfo=
github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.0.2-0.20240212192251-757544f21357/go.mod h1:w9Y7gY31krpLmrVU5ZPG9H7l9fZuRu5/3R3S3FMtVQ4=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 h1:5ZPtiqj0JL5oKWmcsq4VMaAW5ukBEgSGXEN89zeH1Jo=
@@ -411,8 +417,8 @@ github.com/kevinburke/ssh_config v1.4.0 h1:6xxtP5bZ2E4NF5tuQulISpTO2z8XbtH8cg1PW
github.com/kevinburke/ssh_config v1.4.0/go.mod h1:q2RIzfka+BXARoNexmF9gkxEX7DmvbW9P4hIVx2Kg4M=
github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8=
github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
github.com/klauspost/compress v1.18.3 h1:9PJRvfbmTabkOX8moIpXPbMMbYN60bWImDDU7L+/6zw=
github.com/klauspost/compress v1.18.3/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4=
github.com/klauspost/compress v1.18.7 h1:aUyZsS4kH3QTKurYhAOwAHxllVPnOthb3vPfnF1Ehjw=
github.com/klauspost/compress v1.18.7/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
github.com/klauspost/cpuid/v2 v2.0.12/go.mod h1:g2LTdtYhdyuGPqyWyv7qRAmj1WBqxuObKfj5c0PQa7c=
github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y=
github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
@@ -519,8 +525,8 @@ github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502 h1:3tHlFmhTdX9ax
github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502/go.mod h1:CIMRFEJVL+0DS1a3Nx06NaMn4Dz63Ng6O7dl0qH0zVM=
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45 h1:ujgviVYmx243Ksy7NdSwrdGPSRNE3pb8kEDSpH0QuAQ=
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45/go.mod h1:5/sjFmLb8O96B5737VCqhHyGRzNFIaN/Bu7ZodXc3qQ=
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260902163841-4a71f7b1d9e1 h1:n5aXV/U6I9bLc+yWN088TyVR4OfF64Gy+L6Hrffc+n4=
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260902163841-4a71f7b1d9e1/go.mod h1:/6QR46/nhGCSADHbS++XtDb9dkTnenTHlGskTPRo9S0=
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260928122122-d1db5fa35318 h1:Qv2jYeucRkuqkkM9aAVFcO9avmSfPEoB+gxKKKuI4vA=
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260928122122-d1db5fa35318/go.mod h1:62UsqQRqanuCbFi6xElIw3/2X+yGF+TKlyiWImA5taU=
github.com/netbirdio/wireguard-go v0.0.0-20260914123147-8bf8fa968f1a h1:Nt8BgkTkI56LGBPPEBywM406MVKJDqeDIVdgsZyYs80=
github.com/netbirdio/wireguard-go v0.0.0-20260914123147-8bf8fa968f1a/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
github.com/nxadm/tail v1.4.4/go.mod h1:kenIhsEOeOJmVchQTgglprH7qJGnHDVpk1VPCcaMI8A=
+4
View File
@@ -737,6 +737,10 @@ func (p *Provider) UpdateUserPassword(ctx context.Context, userID string, oldPas
return fmt.Errorf("failed to update password: %w", err)
}
if err := p.storage.DeleteAuthSession(ctx, user.UserID, server.LocalConnector); err != nil && !errors.Is(err, storage.ErrNotFound) {
p.logger.Error("failed to revoke local session after password change", "error", err)
}
return nil
}
@@ -6,11 +6,19 @@ set -o pipefail
# NetBird Enterprise — Getting Started
# Single-node bootstrap for a self-hosted NetBird Enterprise stack with the
# embedded identity provider. Owner is created via first-login flow.
# Add features to an existing install with --enable-proxy or --enable-traffic-events.
SED_STRIP_PADDING='s/=//g'
NETBIRD_EULA_URL="https://netbird.io/self-hosted-EULA"
STACK_FILES=(.env docker-compose.yml config.yaml)
# Host directory of a custom TLS certificate mounted at /certs, see
# https://docs.netbird.io/selfhosted/enterprise/getting-started#appendix-using-a-custom-tls-certificate
CUSTOM_TLS_CERTS=""
PROXY_TOKEN_ID=""
# Static IP for Traefik inside the compose bridge network. The management
# server trusts X-Forwarded-* headers from this address only.
TRAEFIK_IP="172.30.0.10"
@@ -44,6 +52,28 @@ check_openssl() {
fi
}
die() {
echo "$1" > /dev/stderr
exit 1
}
# env_get KEY [DEFAULT] prints KEY's value from .env, or DEFAULT if unset.
env_get() {
local value
value=$(sed -n "s/^$1=//p" .env | tail -n 1)
echo "${value:-$2}"
}
# merge_env upserts the KEY=VALUE lines from stdin into .env, in place.
merge_env() {
local merged
merged=$(awk -F= 'NR == FNR { v[$1] = $0; o[++n] = $1; next }
$1 in v { print v[$1]; delete v[$1]; next }
{ print }
END { for (i = 1; i <= n; i++) if (o[i] in v) print v[o[i]] }' - .env)
printf '%s\n' "$merged" > .env
}
rand_secret() {
openssl rand -base64 32 | sed "$SED_STRIP_PADDING"
}
@@ -171,6 +201,15 @@ read_yes_no() {
esac
}
read_crowdsec_option() {
echo ""
echo "CrowdSec:"
echo " Checks client IPs against a community threat intelligence database and"
echo " blocks known malicious sources before they reach services exposed through"
echo " the proxy. Adds a CrowdSec container to the stack."
NETBIRD_CROWDSEC=$(read_yes_no "Enable CrowdSec" "n")
}
# Gate the install on explicit acceptance of the NetBird On-Premise EULA.
require_eula_acceptance() {
cat > /dev/stderr <<EOF
@@ -314,15 +353,100 @@ report_license_verdict() {
return 0
}
# up_all_but_proxy skips the proxy, which needs a token from the running server.
up_all_but_proxy() {
local services
services=$($DOCKER_COMPOSE_COMMAND config --services | grep -vx proxy)
# shellcheck disable=SC2086
$DOCKER_COMPOSE_COMMAND up -d $services
}
wait_crowdsec() {
$DOCKER_COMPOSE_COMMAND up -d crowdsec || return 1
for _ in {1..60}; do
$DOCKER_COMPOSE_COMMAND exec -T crowdsec cscli lapi status &> /dev/null && return 0
sleep 2
done
return 1
}
admin_token() {
$DOCKER_COMPOSE_COMMAND run --rm --no-deps -T netbird-server admin token "$@" --config /etc/netbird/config.yaml
}
# revoke_proxy_token revokes the token this run minted if the proxy never started.
# On failure the ID is kept, so rollback retries it.
revoke_proxy_token() {
[[ -n "$PROXY_TOKEN_ID" ]] || return 0
if admin_token revoke "$PROXY_TOKEN_ID" > /dev/null; then
PROXY_TOKEN_ID=""
return 0
fi
echo "Could not revoke the unused proxy token ${PROXY_TOKEN_ID}. Revoke it with:" > /dev/stderr
echo " $DOCKER_COMPOSE_COMMAND run --rm netbird-server admin token revoke ${PROXY_TOKEN_ID} --config /etc/netbird/config.yaml" > /dev/stderr
}
# start_proxy mints the proxy token and CrowdSec bouncer key, then starts the proxy.
start_proxy() {
local out token key
echo "Creating the proxy access token ..."
out=$(admin_token create --name default-proxy) || true
token=$(awk '/^Token:/ {print $2}' <<< "$out")
PROXY_TOKEN_ID=$(awk '/^Token ID:/ {print $3}' <<< "$out")
[[ -n "$token" ]] || die "Could not create the proxy access token. Check the netbird-server logs, then re-run with --enable-proxy."
if [[ "$NETBIRD_CROWDSEC" == "yes" ]]; then
echo "Registering the CrowdSec bouncer ..."
if wait_crowdsec; then
# "add" fails if an earlier attempt already registered the bouncer.
$DOCKER_COMPOSE_COMMAND exec -T crowdsec cscli bouncers delete netbird-proxy &> /dev/null || true
key=$($DOCKER_COMPOSE_COMMAND exec -T crowdsec cscli bouncers add netbird-proxy -o raw) || true
fi
if [[ -z "$key" ]]; then
revoke_proxy_token
die "Could not register the CrowdSec bouncer. Check the crowdsec logs, then re-run with --enable-proxy."
fi
fi
# A stored token marks the proxy as set up, so it is cleared again on failure.
{
echo "NETBIRD_PROXY_TOKEN=${token}"
if [[ -n "$key" ]]; then echo "NETBIRD_CROWDSEC_BOUNCER_KEY=${key}"; fi
} | merge_env
if ! $DOCKER_COMPOSE_COMMAND up -d proxy; then
revoke_proxy_token
echo "NETBIRD_PROXY_TOKEN=" | merge_env
die "Could not start the proxy. Check the proxy logs, then re-run with --enable-proxy."
fi
PROXY_TOKEN_ID=""
}
print_proxy_notes() {
echo ""
echo "NetBird Proxy:"
echo " Every domain other than ${NETBIRD_DOMAIN} is passed through to the proxy,"
echo " which issues its own TLS certificates. Point proxy domains at this host:"
echo ""
echo " *.${NETBIRD_DOMAIN} CNAME ${NETBIRD_DOMAIN}"
echo ""
echo " Open 51820/udp (optional) for peer-to-peer proxy connections."
if [[ "$NETBIRD_CROWDSEC" == "yes" ]]; then
echo " CrowdSec is running. Enable it per service in the dashboard under Access Control."
fi
}
init_environment() {
check_openssl
DOCKER_COMPOSE_COMMAND=$(check_docker_compose)
if [[ -f .env ]] || [[ -f docker-compose.yml ]] || [[ -f config.yaml ]]; then
echo "Generated files already exist in $(pwd)."
echo "To add the proxy or traffic events to this installation, re-run with"
echo "--enable-proxy or --enable-traffic-events."
echo ""
echo "If you want to reinitialize the environment, please remove them first:"
echo " $DOCKER_COMPOSE_COMMAND down --volumes # removes all containers and volumes"
echo " rm -f .env docker-compose.yml config.yaml"
echo " rm -rf .env docker-compose.yml config.yaml traefik"
echo "Be aware this will remove all data from the database."
exit 1
fi
@@ -341,6 +465,16 @@ init_environment() {
echo " See https://docs.netbird.io/manage/activity/traffic-events-logging"
NETBIRD_TRAFFIC_FLOW=$(read_yes_no "Enable traffic flow" "n")
echo ""
echo "NetBird Proxy:"
echo " Exposes selected resources from your NetBird network to the internet."
echo " You choose which resources are exposed from the dashboard."
NETBIRD_PROXY=$(read_yes_no "Enable the NetBird Proxy" "n")
NETBIRD_CROWDSEC="no"
if [[ "$NETBIRD_PROXY" == "yes" ]]; then
read_crowdsec_option
fi
echo ""
NETBIRD_DOMAIN=$(read_nb_domain)
@@ -364,6 +498,8 @@ init_environment() {
echo ""
echo "Selected:"
echo " Traffic flow: ${NETBIRD_TRAFFIC_FLOW}"
echo " Proxy: ${NETBIRD_PROXY}"
echo " CrowdSec: ${NETBIRD_CROWDSEC}"
echo " Domain: ${NETBIRD_DOMAIN}"
echo " ACME email: ${NETBIRD_LETSENCRYPT_EMAIL}"
echo ""
@@ -371,9 +507,9 @@ init_environment() {
install -m 600 /dev/null .env
render_env >> .env
render_docker_compose > docker-compose.yml
if [[ -z "${NETBIRD_LICENSE_SERVER_BASE_URL:-}" ]]; then
sed -i.bak '/NETBIRD_LICENSE_SERVER_BASE_URL/d' docker-compose.yml && rm -f docker-compose.yml.bak
mkdir -p traefik
if [[ "$NETBIRD_PROXY" == "yes" ]]; then
render_traefik_proxy > traefik/proxy.yaml
fi
install -m 600 /dev/null config.yaml
render_config_yaml >> config.yaml
@@ -390,11 +526,16 @@ init_environment() {
echo ""
echo "Starting remaining services ..."
$DOCKER_COMPOSE_COMMAND up -d
up_all_but_proxy
echo ""
wait_for_license_verdict
if [[ "$NETBIRD_PROXY" == "yes" ]]; then
echo ""
start_proxy
fi
echo ""
echo "Done."
echo ""
@@ -402,6 +543,9 @@ init_environment() {
echo ""
echo "Open the dashboard in a browser to complete the first-login owner setup."
echo "All configuration and secrets are stored (mode 600) in $(pwd)/.env"
if [[ "$NETBIRD_PROXY" == "yes" ]]; then
print_proxy_notes
fi
echo ""
echo "Tail logs:"
echo " cd $(pwd) && $DOCKER_COMPOSE_COMMAND logs -f netbird-server traefik"
@@ -413,6 +557,148 @@ init_environment() {
fi
}
# service_block NAME prints a service's definition from the compose file on stdin.
service_block() {
local name="$1"
awk -v s=" ${name}:" '$0 == s { p = 1; print; next } p && (/^[^ ]/ || /^ [^ ]/) { exit } p'
}
# enable_features adds the proxy and/or traffic events to the install in the
# current directory, restoring the backed-up files if any step fails.
enable_features() {
local want_proxy="$1" want_flow="$2" f compose
DOCKER_COMPOSE_COMMAND=$(check_docker_compose)
for f in "${STACK_FILES[@]}"; do
[[ -f "$f" ]] || die "$f not found in $(pwd). Run this from an existing installation directory."
done
grep -q '^# Generated by getting-started-enterprise.sh' .env || die ".env was not generated by getting-started-enterprise.sh."
# Installs from before the move to Traefik run Caddy and can't be re-rendered.
[[ -n "$(env_get NETBIRD_TRAEFIK_IP)" ]] || die "This installation predates the Traefik layout and can't be updated in place."
NETBIRD_DOMAIN=$(env_get NETBIRD_DOMAIN)
NETBIRD_LICENSE_SERVER_BASE_URL=$(env_get NETBIRD_LICENSE_SERVER_BASE_URL)
NETBIRD_TRAFFIC_FLOW=$(env_get NETBIRD_TRAFFIC_FLOW_ENABLED no)
NETBIRD_PROXY=$(env_get NETBIRD_PROXY_ENABLED no)
NETBIRD_CROWDSEC=$(env_get NETBIRD_CROWDSEC_ENABLED no)
# A custom certificate counts only once Traefik mounts it and ACME is already gone,
# so the re-render never removes a working Let's Encrypt setup.
local traefik_block
traefik_block=$(service_block traefik < docker-compose.yml)
if ! grep -q certificatesresolvers <<< "$traefik_block"; then
CUSTOM_TLS_CERTS=$(awk '/:\/certs:ro$/ { sub(/^ *- /, ""); sub(/:\/certs:ro$/, ""); print; exit }' <<< "$traefik_block")
fi
if [[ "$want_flow" == "yes" && "$NETBIRD_TRAFFIC_FLOW" == "yes" ]]; then
echo "Traffic events are already enabled."
want_flow="no"
fi
# No token means an earlier proxy setup failed, so let it run again.
if [[ "$want_proxy" == "yes" && -n "$(env_get NETBIRD_PROXY_TOKEN)" ]]; then
echo "The NetBird Proxy is already enabled."
want_proxy="no"
fi
if [[ "$want_flow" == "no" && "$want_proxy" == "no" ]]; then
exit 0
fi
if [[ "$want_flow" == "yes" ]]; then
NETBIRD_TRAFFIC_FLOW="yes"
fi
# A retry keeps the CrowdSec choice made at install time.
if [[ "$want_proxy" == "yes" && "$NETBIRD_PROXY" != "yes" ]]; then
read_crowdsec_option
fi
if [[ "$want_proxy" == "yes" ]]; then
NETBIRD_PROXY="yes"
fi
compose=$(render_docker_compose)
echo ""
echo "Changes to docker-compose.yml:"
printf '%s\n' "$compose" | diff -u docker-compose.yml - || true
# Lines the new file drops are most likely local edits, so don't default to applying.
local lost s restarts="" apply="y"
lost=$(printf '%s\n' "$compose" | awk 'NR == FNR { keep[$0]; next } !($0 in keep)' - docker-compose.yml)
if [[ -n "$lost" ]]; then
echo ""
echo "These lines are not in the new docker-compose.yml and will be lost:"
printf '%s\n' "$lost"
apply="n"
fi
for s in $($DOCKER_COMPOSE_COMMAND config --services); do
if [[ "$(service_block "$s" < docker-compose.yml)" != "$(printf '%s\n' "$compose" | service_block "$s")" ]] \
|| [[ "$s" == "netbird-server" && "$want_flow" == "yes" ]]; then
restarts+=" $s"
fi
done
echo ""
if [[ -n "$restarts" ]]; then
echo "These services will restart:${restarts}"
fi
if [[ "$(read_yes_no "Apply these changes?" "$apply")" != "yes" ]]; then
echo "Aborted."
exit 0
fi
BACKUP_SUFFIX=".bak.$(date -u +%Y%m%d%H%M%S)"
for f in "${STACK_FILES[@]}" traefik/proxy.yaml; do
if [[ -f "$f" ]]; then cp -p "$f" "$f$BACKUP_SUFFIX"; fi
done
trap rollback EXIT
printf '%s\n' "$compose" > docker-compose.yml
{
echo "NETBIRD_TRAFFIC_FLOW_ENABLED=${NETBIRD_TRAFFIC_FLOW}"
echo "NETBIRD_PROXY_ENABLED=${NETBIRD_PROXY}"
echo "NETBIRD_CROWDSEC_ENABLED=${NETBIRD_CROWDSEC}"
if [[ "$want_flow" == "yes" ]]; then render_env_flow; fi
if [[ "$want_proxy" == "yes" ]]; then render_env_proxy; fi
} | merge_env
if [[ "$want_flow" == "yes" ]] && ! grep -q '^ trafficFlow:' config.yaml; then
render_config_flow >> config.yaml
fi
mkdir -p traefik
if [[ "$want_proxy" == "yes" ]]; then
render_traefik_proxy > traefik/proxy.yaml
fi
up_all_but_proxy
if [[ "$want_flow" == "yes" ]]; then
# Compose does not notice changes to the bind-mounted config.yaml.
$DOCKER_COMPOSE_COMMAND restart netbird-server
fi
if [[ "$want_proxy" == "yes" ]]; then
start_proxy
fi
trap - EXIT
echo ""
echo "Done. The previous files are kept with the ${BACKUP_SUFFIX} suffix."
if [[ "$want_flow" == "yes" ]]; then
echo ""
echo "Traffic events still have to be turned on from the dashboard settings."
echo " See https://docs.netbird.io/manage/activity/traffic-events-logging"
fi
if [[ "$want_proxy" == "yes" ]]; then
print_proxy_notes
fi
}
rollback() {
local f
echo "" > /dev/stderr
echo "Enabling failed. Restoring the previous configuration ..." > /dev/stderr
revoke_proxy_token
# Files without a backup were created by this run.
for f in "${STACK_FILES[@]}" traefik/proxy.yaml; do
if [[ -f "$f$BACKUP_SUFFIX" ]]; then cp -p "$f$BACKUP_SUFFIX" "$f"; else rm -f "$f"; fi
done
$DOCKER_COMPOSE_COMMAND up -d --remove-orphans
$DOCKER_COMPOSE_COMMAND restart netbird-server
}
# ------------------------------------------------------------------
# Renderers
# ------------------------------------------------------------------
@@ -427,8 +713,10 @@ NETBIRD_EULA_ACCEPTED=yes
NETBIRD_EULA_ACCEPTED_AT=${NETBIRD_EULA_ACCEPTED_AT}
NETBIRD_EULA_URL=${NETBIRD_EULA_URL}
# Features (set by the script; don't edit without re-running)
# Features (change with --enable-proxy or --enable-traffic-events, not by hand)
NETBIRD_TRAFFIC_FLOW_ENABLED=${NETBIRD_TRAFFIC_FLOW}
NETBIRD_PROXY_ENABLED=${NETBIRD_PROXY}
NETBIRD_CROWDSEC_ENABLED=${NETBIRD_CROWDSEC}
# Domain
NETBIRD_DOMAIN=${NETBIRD_DOMAIN}
@@ -444,10 +732,11 @@ NETBIRD_SERVER_TAG=${NETBIRD_SERVER_TAG:-latest}
EOF
if [[ "$NETBIRD_TRAFFIC_FLOW" == "yes" ]]; then
cat <<EOF
NETBIRD_ENRICHER_TAG=${NETBIRD_ENRICHER_TAG:-latest}
NETBIRD_RECEIVER_TAG=${NETBIRD_RECEIVER_TAG:-latest}
EOF
render_env_flow
fi
if [[ "$NETBIRD_PROXY" == "yes" ]]; then
printf '\n# NetBird Proxy (token and bouncer key are filled in after startup)\n'
render_env_proxy
fi
cat <<EOF
@@ -483,15 +772,35 @@ NETBIRD_AUTH_SUPPORTED_SCOPES=${NETBIRD_AUTH_SUPPORTED_SCOPES:-openid profile em
EOF
}
render_docker_compose() {
render_compose_header
render_compose_common
render_compose_server
if [[ "$NETBIRD_TRAFFIC_FLOW" == "yes" ]]; then
render_compose_flow
render_env_flow() {
echo "NETBIRD_ENRICHER_TAG=${NETBIRD_ENRICHER_TAG:-latest}"
echo "NETBIRD_RECEIVER_TAG=${NETBIRD_RECEIVER_TAG:-latest}"
}
render_env_proxy() {
echo "NETBIRD_PROXY_TAG=${NETBIRD_PROXY_TAG:-latest}"
echo "NETBIRD_PROXY_TOKEN="
if [[ "$NETBIRD_CROWDSEC" == "yes" ]]; then
echo "NETBIRD_CROWDSEC_TAG=${NETBIRD_CROWDSEC_TAG:-v1.7.7}"
echo "NETBIRD_CROWDSEC_BOUNCER_KEY="
fi
render_compose_postgres
render_compose_footer
}
render_docker_compose() {
{
render_compose_header
render_compose_common
render_compose_server
if [[ "$NETBIRD_TRAFFIC_FLOW" == "yes" ]]; then
render_compose_flow
fi
if [[ "$NETBIRD_PROXY" == "yes" ]]; then
render_compose_proxy
fi
render_compose_postgres
render_compose_footer
} | if [[ -n "${NETBIRD_LICENSE_SERVER_BASE_URL:-}" ]]; then cat; else sed '/NETBIRD_LICENSE_SERVER_BASE_URL/d'; fi \
| if [[ -n "$CUSTOM_TLS_CERTS" ]]; then sed -e '/certificatesresolvers/d' -e '/certresolver/d'; else cat; fi
}
render_compose_header() {
@@ -544,12 +853,20 @@ render_compose_common() {
- "--certificatesresolvers.letsencrypt.acme.email=${NETBIRD_LETSENCRYPT_EMAIL}"
- "--certificatesresolvers.letsencrypt.acme.storage=/letsencrypt/acme.json"
- "--certificatesresolvers.letsencrypt.acme.tlschallenge=true"
# Dynamic config in ./traefik: the proxy transport and an optional custom certificate
- "--providers.file.directory=/etc/traefik/dynamic"
ports:
- '443:443'
- '80:80'
volumes:
- /var/run/docker.sock:/var/run/docker.sock:ro
- netbird_traefik_letsencrypt:/letsencrypt
- ./traefik:/etc/traefik/dynamic:ro
EOF
if [[ -n "$CUSTOM_TLS_CERTS" ]]; then
echo " - ${CUSTOM_TLS_CERTS}:/certs:ro"
fi
cat <<'EOF'
labels:
- traefik.enable=true
# Shared security headers, referenced by every NetBird router below. A
@@ -719,6 +1036,69 @@ render_compose_flow() {
EOF
}
render_compose_proxy() {
cat <<'EOF'
# Traefik passes TLS for every other domain through to the proxy, which issues
# its own certificates. PROXY protocol v2 preserves the client IP.
proxy:
<<: *default
image: ghcr.io/netbirdio/reverse-proxy:${NETBIRD_PROXY_TAG}
container_name: netbird-proxy
networks: [netbird]
depends_on:
netbird-server:
condition: service_started
ports:
- '51820:51820/udp'
volumes:
- netbird_proxy_certs:/certs
labels:
- traefik.enable=true
- traefik.tcp.routers.proxy-passthrough.rule=HostSNI(`*`)
- traefik.tcp.routers.proxy-passthrough.entrypoints=websecure
- traefik.tcp.routers.proxy-passthrough.tls.passthrough=true
- traefik.tcp.routers.proxy-passthrough.service=proxy-tls
- traefik.tcp.routers.proxy-passthrough.priority=1
- traefik.tcp.services.proxy-tls.loadbalancer.server.port=8443
- traefik.tcp.services.proxy-tls.loadbalancer.serverstransport=pp-v2@file
environment:
# Plaintext gRPC over the internal network, not the public address.
- NB_PROXY_MANAGEMENT_ADDRESS=http://netbird-server:80
- NB_PROXY_ALLOW_INSECURE=true
- NB_PROXY_TOKEN=${NETBIRD_PROXY_TOKEN}
- NB_PROXY_DOMAIN=${NETBIRD_DOMAIN}
- NB_PROXY_ADDRESS=:8443
- NB_PROXY_CERTIFICATE_DIRECTORY=/certs
- NB_PROXY_ACME_CERTIFICATES=true
- NB_PROXY_FORWARDED_PROTO=https
- NB_PROXY_PROXY_PROTOCOL=true
- NB_PROXY_TRUSTED_PROXIES=${NETBIRD_TRAEFIK_IP}
EOF
if [[ "$NETBIRD_CROWDSEC" == "yes" ]]; then
cat <<'EOF'
- NB_PROXY_CROWDSEC_API_URL=http://crowdsec:8080
- NB_PROXY_CROWDSEC_API_KEY=${NETBIRD_CROWDSEC_BOUNCER_KEY}
crowdsec:
<<: *default
image: crowdsecurity/crowdsec:${NETBIRD_CROWDSEC_TAG}
container_name: netbird-crowdsec
networks: [netbird]
environment:
- COLLECTIONS=crowdsecurity/linux
volumes:
- netbird_crowdsec_config:/etc/crowdsec
- netbird_crowdsec_data:/var/lib/crowdsec/data
healthcheck:
test: ["CMD", "cscli", "lapi", "status"]
interval: 10s
timeout: 5s
retries: 15
EOF
fi
echo ""
}
render_compose_postgres() {
cat <<'EOF'
postgres:
@@ -752,6 +1132,12 @@ EOF
netbird_enricher:
EOF
fi
if [[ "$NETBIRD_PROXY" == "yes" ]]; then
echo " netbird_proxy_certs:"
fi
if [[ "$NETBIRD_CROWDSEC" == "yes" ]]; then
printf ' %s:\n' netbird_crowdsec_config netbird_crowdsec_data
fi
cat <<'EOF'
netbird_postgres:
netbird_traefik_letsencrypt:
@@ -827,14 +1213,61 @@ server:
EOF
if [[ "$NETBIRD_TRAFFIC_FLOW" == "yes" ]]; then
cat <<EOF
render_config_flow
fi
}
render_config_flow() {
cat <<EOF
trafficFlow:
enabled: true
address: "https://${NETBIRD_DOMAIN}:443"
interval: "60s"
EOF
}
render_traefik_proxy() {
cat <<'EOF'
tcp:
serversTransports:
pp-v2:
proxyProtocol:
version: 2
EOF
}
usage() {
cat <<EOF
Usage: $0 [--enable-proxy] [--enable-traffic-events]
Without flags, bootstraps a new NetBird Enterprise stack in the current
directory. With flags, turns the feature on for the existing installation in
the current directory.
--enable-proxy add the NetBird Proxy, optionally with CrowdSec
--enable-traffic-events add traffic events logging (NATS, receiver, enricher)
-h, --help show this help
EOF
}
main() {
local enable_proxy="no" enable_flow="no"
while [[ $# -gt 0 ]]; do
case "$1" in
--enable-proxy) enable_proxy="yes" ;;
--enable-traffic-events) enable_flow="yes" ;;
-h | --help) usage; exit 0 ;;
*) usage > /dev/stderr; exit 1 ;;
esac
shift
done
if [[ "$enable_proxy" == "no" && "$enable_flow" == "no" ]]; then
init_environment
else
enable_features "$enable_proxy" "$enable_flow"
fi
}
init_environment
main "$@"
+18 -4
View File
@@ -220,6 +220,21 @@ detect_exposed_address() {
yq eval '.server.exposedAddress // ""' "$CONFIG_YAML_HOST"
}
detect_relay_auth_secret() {
local secret=""
local external_relay_count
external_relay_count=$(yq eval '(.server.relays.addresses // []) | length' "$CONFIG_YAML_HOST")
if (( external_relay_count > 0 )); then
secret=$(yq eval '.server.relays.secret // ""' "$CONFIG_YAML_HOST")
fi
if [[ -z "$secret" ]] || [[ "$secret" == "null" ]]; then
secret=$(yq eval '.server.authSecret // ""' "$CONFIG_YAML_HOST")
fi
printf '%s' "$secret"
}
# The engine is a config.yaml-only setting — there is no env override for it
# (combined/cmd/root.go reads it from YAML and derives the env vars), so
# config.yaml is authoritative. Absent means the sqlite default.
@@ -945,11 +960,10 @@ init_migration() {
if [[ "$MIGRATE_POSTGRES" == "yes" ]] || [[ "$EXISTING_POSTGRES" == "yes" ]]; then
ENABLE_FLOW=$(read_yes_no "Step 3: Enable traffic flow? (requires Postgres)" "n")
if [[ "$ENABLE_FLOW" == "yes" ]]; then
# Auth secret MUST match server.authSecret from config.yaml
NB_FLOW_AUTH_SECRET=$(yq eval '.server.authSecret // ""' "$CONFIG_YAML_HOST")
NB_FLOW_AUTH_SECRET=$(detect_relay_auth_secret)
if [[ -z "$NB_FLOW_AUTH_SECRET" ]] || [[ "$NB_FLOW_AUTH_SECRET" == "null" ]]; then
echo "Could not read server.authSecret from $CONFIG_YAML_HOST." > /dev/stderr
echo "Flow receiver auth must match the combined server's authSecret." > /dev/stderr
echo "Could not resolve the Relay auth secret from $CONFIG_YAML_HOST." > /dev/stderr
echo "Set server.relays.secret for external Relays or server.authSecret for the local Relay." > /dev/stderr
exit 1
fi
@@ -9,8 +9,8 @@ import (
"strings"
networkmap_sqlite "github.com/netbirdio/netbird/management/internals/network_map_db/sqlite"
nbdb "github.com/netbirdio/netbird/management/internals/shared/db"
gormstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
log "github.com/sirupsen/logrus"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
@@ -28,7 +28,11 @@ func createSqliteTestStore(baseData string) (*networkmap_sqlite.SqliteStore, fun
if err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
_, err = gormstore.NewSqlStore(context.TODO(), db, types.SqliteStoreEngine, nil, false)
conn, err := nbdb.NewConn(context.TODO(), db, nbdb.SqliteStoreEngine, nil)
if err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
_, err = gormstore.NewSqlStore(context.TODO(), conn, nil, false)
if err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
+23 -4
View File
@@ -10,6 +10,7 @@ import (
"io"
"strings"
"text/tabwriter"
"unicode"
"github.com/spf13/cobra"
@@ -68,8 +69,8 @@ func runDisconnectAll(ctx context.Context, s store.Store, out io.Writer, in io.R
toDisconnect := 0
w := tabwriter.NewWriter(out, 0, 0, 2, ' ', 0)
_, _ = fmt.Fprintln(w, "ID\tCLUSTER\tIP\tACCOUNT\tSTATUS\tLAST SEEN")
_, _ = fmt.Fprintln(w, "--\t-------\t--\t-------\t------\t---------")
_, _ = fmt.Fprintln(w, "ID\tCLUSTER\tIP\tVERSION\tACCOUNT\tSTATUS\tLAST SEEN")
_, _ = fmt.Fprintln(w, "--\t-------\t--\t-------\t-------\t------\t---------")
for _, p := range proxies {
if p.Status != rpproxy.StatusDisconnected {
@@ -80,11 +81,16 @@ func runDisconnectAll(ctx context.Context, s store.Store, out io.Writer, in io.R
if p.AccountID != nil {
account = *p.AccountID
}
version := "-"
if p.Version != "" {
version = sanitizeReportedValue(p.Version)
}
_, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\t%s\n",
p.ID,
_, _ = fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\t%s\t%s\n",
sanitizeReportedValue(p.ID),
p.ClusterAddress,
p.IPAddress,
version,
account,
p.Status,
p.LastSeen.Format("2006-01-02 15:04:05"),
@@ -139,3 +145,16 @@ func confirmDisconnectAll(out io.Writer, in io.Reader) (bool, error) {
return strings.EqualFold(strings.TrimSpace(scanner.Text()), disconnectAllConfirmation), nil
}
// sanitizeReportedValue replaces non-printable characters in a value the proxy
// reports about itself. Both the id and the version arrive unvalidated over
// gRPC, so a tab would forge a column, a carriage return or ANSI escape would
// redraw the operator's terminal, and U+202E would reverse the rest of the line.
func sanitizeReportedValue(s string) string {
return strings.Map(func(r rune) rune {
if unicode.IsPrint(r) {
return r
}
return '\uFFFD'
}, s)
}
+39
View File
@@ -35,6 +35,7 @@ func seedProxies(t *testing.T, ctx context.Context, s store.Store) {
SessionID: "session-1",
ClusterAddress: "cluster-a.example.com",
IPAddress: "10.0.0.1",
Version: "0.60.0",
LastSeen: time.Now(),
Status: rpproxy.StatusConnected,
},
@@ -89,6 +90,7 @@ func TestRunDisconnectAllWithConfirmation(t *testing.T) {
require.Contains(t, output, "proxy-2")
require.Contains(t, output, "proxy-3")
require.Contains(t, output, "cluster-a.example.com")
require.Contains(t, output, "0.60.0")
require.Contains(t, output, "account-1")
require.Contains(t, output, "Type \"disconnect all proxies\" to continue")
require.Contains(t, output, "Force-marked 2 of 3 reverse proxy instance(s) as disconnected.")
@@ -178,3 +180,40 @@ func TestRunDisconnectAllEmpty(t *testing.T) {
require.NoError(t, runDisconnectAll(ctx, s, &out, strings.NewReader(""), false, false))
require.Contains(t, out.String(), "No reverse proxy instances found.")
}
func TestRunDisconnectAllEscapesProxyReportedFields(t *testing.T) {
ctx := context.Background()
s := newTestStore(t)
// A proxy reports its own id and version on connect, so both reach this
// listing unvalidated. Carriage returns, tabs and ANSI escapes would let
// a malicious proxy redraw the table or forge a row on the operator's
// terminal; U+202E would reverse the rendering of the rest of the line.
require.NoError(t, s.SaveProxy(ctx, &rpproxy.Proxy{
ID: "proxy-\r\x1b[2Kevil",
SessionID: "session-1",
ClusterAddress: "cluster-a.example.com",
IPAddress: "10.0.0.1",
Version: "0.60.0\tfake\rcolumn\u202e",
LastSeen: time.Now(),
Status: rpproxy.StatusConnected,
}))
var out bytes.Buffer
require.NoError(t, runDisconnectAll(ctx, s, &out, strings.NewReader(disconnectAllConfirmation+"\n"), true, false))
output := out.String()
for _, forbidden := range []string{"\r", "\x1b", "\u202e"} {
require.NotContains(t, output, forbidden, "listing must not carry proxy-reported control characters")
}
// The table has one data row; a smuggled tab would add a phantom column.
var dataRow string
for _, line := range strings.Split(output, "\n") {
if strings.Contains(line, "evil") {
dataRow = line
}
}
require.NotEmpty(t, dataRow, "listing should still show the proxy row")
require.NotContains(t, dataRow, "\t", "tabwriter output should not carry a smuggled column separator")
require.Contains(t, dataRow, "0.60.0", "the printable part of the version should survive")
}
@@ -245,7 +245,7 @@ func computeMode(t *testing.T, ctx context.Context, mode Mode, nmData *networkma
peerGroups := maps.Keys(nmData.GetPeerGroups(peerID))
resp := mgmtgrpc.ToComponentSyncResponse(ctx, nil, nil, nil, peer, nil, nil, components, nil,
dnsDomain, nil, nmData.AccountSettings, nil, peerGroups, dnsFwdPort)
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain)
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain, false)
require.NoError(t, err, "expand envelope")
return res.NetworkMap
default:
@@ -9,6 +9,7 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/management/server/geolocation"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
@@ -18,14 +19,16 @@ import (
)
type managerImpl struct {
repo accesslogs.Repository
store store.Store
permissionsManager permissions.Manager
geo geolocation.Geolocation
cleanupCancel context.CancelFunc
}
func NewManager(store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager {
func NewManager(repo accesslogs.Repository, store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager {
return &managerImpl{
repo: repo,
store: store,
permissionsManager: permissionsManager,
geo: geo,
@@ -54,7 +57,7 @@ func (m *managerImpl) SaveAccessLog(ctx context.Context, logEntry *accesslogs.Ac
}
}
if err := m.store.CreateAccessLog(ctx, logEntry); err != nil {
if err := m.repo.Create(ctx, logEntry); err != nil {
log.WithContext(ctx).WithFields(log.Fields{
"service_id": logEntry.ServiceID,
"method": logEntry.Method,
@@ -82,7 +85,7 @@ func (m *managerImpl) GetAllAccessLogs(ctx context.Context, accountID, userID st
log.WithContext(ctx).Warnf("failed to resolve user filters: %v", err)
}
logs, totalCount, err := m.store.GetAccountAccessLogs(ctx, store.LockingStrengthNone, accountID, *filter)
logs, totalCount, err := m.repo.ListByAccount(ctx, db.LockingStrengthNone, accountID, *filter)
if err != nil {
return nil, 0, err
}
@@ -98,7 +101,7 @@ func (m *managerImpl) CleanupOldAccessLogs(ctx context.Context, retentionDays in
}
cutoffTime := time.Now().AddDate(0, 0, -retentionDays)
deletedCount, err := m.store.DeleteOldAccessLogs(ctx, cutoffTime)
deletedCount, err := m.repo.DeleteOlderThan(ctx, cutoffTime)
if err != nil {
log.WithContext(ctx).Errorf("failed to cleanup old access logs: %v", err)
return 0, err
@@ -5,27 +5,27 @@ import (
"testing"
"time"
"go.uber.org/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
)
func TestCleanupOldAccessLogs(t *testing.T) {
tests := []struct {
name string
retentionDays int
setupMock func(*store.MockStore)
setupMock func(*accesslogs.MockRepository)
expectedCount int64
expectedError bool
}{
{
name: "cleanup logs older than retention period",
retentionDays: 30,
setupMock: func(mockStore *store.MockStore) {
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
setupMock: func(mockRepo *accesslogs.MockRepository) {
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) {
expectedCutoff := time.Now().AddDate(0, 0, -30)
timeDiff := olderThan.Sub(expectedCutoff)
@@ -41,9 +41,9 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "no logs to cleanup",
retentionDays: 30,
setupMock: func(mockStore *store.MockStore) {
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
setupMock: func(mockRepo *accesslogs.MockRepository) {
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(0), nil)
},
expectedCount: 0,
@@ -52,8 +52,8 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "zero retention days skips cleanup",
retentionDays: 0,
setupMock: func(mockStore *store.MockStore) {
// No expectations - DeleteOldAccessLogs should not be called
setupMock: func(mockRepo *accesslogs.MockRepository) {
// No expectations - DeleteOlderThan should not be called
},
expectedCount: 0,
expectedError: false,
@@ -61,8 +61,8 @@ func TestCleanupOldAccessLogs(t *testing.T) {
{
name: "negative retention days skips cleanup",
retentionDays: -10,
setupMock: func(mockStore *store.MockStore) {
// No expectations - DeleteOldAccessLogs should not be called
setupMock: func(mockRepo *accesslogs.MockRepository) {
// No expectations - DeleteOlderThan should not be called
},
expectedCount: 0,
expectedError: false,
@@ -74,11 +74,11 @@ func TestCleanupOldAccessLogs(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
tt.setupMock(mockStore)
mockRepo := accesslogs.NewMockRepository(ctrl)
tt.setupMock(mockRepo)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -98,10 +98,10 @@ func TestCleanupWithExactBoundary(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) {
expectedCutoff := time.Now().AddDate(0, 0, -30)
timeDiff := olderThan.Sub(expectedCutoff)
@@ -110,7 +110,7 @@ func TestCleanupWithExactBoundary(t *testing.T) {
})
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -125,11 +125,11 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
// No expectations - cleanup should not run
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -139,22 +139,22 @@ func TestStartPeriodicCleanup(t *testing.T) {
time.Sleep(100 * time.Millisecond)
// If DeleteOldAccessLogs was called, the test will fail due to unexpected call
// If DeleteOlderThan was called, the test will fail due to unexpected call
})
t.Run("periodic cleanup runs immediately on start", func(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(2), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -171,15 +171,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(1), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -198,15 +198,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(0), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -223,15 +223,15 @@ func TestStartPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(3), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx, cancel := context.WithCancel(context.Background())
@@ -249,15 +249,15 @@ func TestStopPeriodicCleanup(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
mockRepo := accesslogs.NewMockRepository(ctrl)
mockStore.EXPECT().
DeleteOldAccessLogs(gomock.Any(), gomock.Any()).
mockRepo.EXPECT().
DeleteOlderThan(gomock.Any(), gomock.Any()).
Return(int64(1), nil).
Times(1)
manager := &managerImpl{
store: mockStore,
repo: mockRepo,
}
ctx := context.Background()
@@ -0,0 +1,135 @@
package manager
import (
"context"
"strings"
"time"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/shared/management/status"
)
type sqlRepository struct {
conn *db.Conn
db *gorm.DB
}
// NewRepository returns the access log repository backed by conn.
func NewRepository(conn *db.Conn) accesslogs.Repository {
return &sqlRepository{conn: conn, db: conn.DB(nil)}
}
func (r *sqlRepository) WithTx(tx *db.Tx) accesslogs.Repository {
return &sqlRepository{conn: r.conn, db: r.conn.DB(tx)}
}
func (r *sqlRepository) Create(ctx context.Context, entry *accesslogs.AccessLogEntry) error {
if err := r.db.Create(entry).Error; err != nil {
log.WithContext(ctx).WithFields(log.Fields{
"service_id": entry.ServiceID,
"method": entry.Method,
"host": entry.Host,
"path": entry.Path,
}).Errorf("failed to create access log entry in store: %v", err)
return status.Errorf(status.Internal, "failed to create access log entry in store")
}
return nil
}
// ListByAccount returns one page of an account's access logs together with the
// total number of entries matching the filter.
func (r *sqlRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) {
var totalCount int64
countQuery := applyFilters(r.db.Model(&accesslogs.AccessLogEntry{}).Where("account_id = ?", accountID), filter)
if err := countQuery.Count(&totalCount).Error; err != nil {
log.WithContext(ctx).Errorf("failed to count access logs: %v", err)
return nil, 0, status.Errorf(status.Internal, "failed to count access logs")
}
query := applyFilters(r.db.Where("account_id = ?", accountID), filter)
sortOrder := strings.ToUpper(filter.GetSortOrder())
for _, column := range strings.Split(filter.GetSortColumn(), ",") {
if column = strings.TrimSpace(column); column != "" {
query = query.Order(column + " " + sortOrder)
}
}
query = query.Limit(filter.GetLimit()).Offset(filter.GetOffset())
if lockStrength != db.LockingStrengthNone {
query = query.Clauses(clause.Locking{Strength: string(lockStrength)})
}
var logs []*accesslogs.AccessLogEntry
if err := query.Find(&logs).Error; err != nil {
log.WithContext(ctx).Errorf("failed to get access logs from store: %v", err)
return nil, 0, status.Errorf(status.Internal, "failed to get access logs from store")
}
return logs, totalCount, nil
}
func (r *sqlRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) {
result := r.db.Where("timestamp < ?", olderThan).Delete(&accesslogs.AccessLogEntry{})
if result.Error != nil {
log.WithContext(ctx).Errorf("failed to delete old access logs: %v", result.Error)
return 0, status.Errorf(status.Internal, "failed to delete old access logs")
}
return result.RowsAffected, nil
}
func applyFilters(query *gorm.DB, filter accesslogs.AccessLogFilter) *gorm.DB {
if filter.Search != nil {
searchPattern := "%" + *filter.Search + "%"
query = query.Where(
"id LIKE ? OR location_connection_ip LIKE ? OR host LIKE ? OR path LIKE ? OR CONCAT(host, path) LIKE ? OR user_id IN (SELECT id FROM users WHERE email LIKE ? OR name LIKE ?)",
searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern, searchPattern,
)
}
if filter.SourceIP != nil {
query = query.Where("location_connection_ip = ?", *filter.SourceIP)
}
if filter.Host != nil {
query = query.Where("host = ?", *filter.Host)
}
if filter.Path != nil {
query = query.Where("path LIKE ?", "%"+*filter.Path+"%")
}
if filter.UserID != nil {
query = query.Where("user_id = ?", *filter.UserID)
}
if filter.Method != nil {
query = query.Where("method = ?", *filter.Method)
}
if filter.Status != nil {
switch *filter.Status {
case "success":
query = query.Where("(status_code >= ? AND status_code < ?)", 200, 400)
case "failed":
query = query.Where("((status_code >= ? AND status_code < ?) OR status_code >= ?)", 100, 200, 400)
}
}
if filter.StatusCode != nil {
query = query.Where("status_code = ?", *filter.StatusCode)
}
if filter.StartDate != nil {
query = query.Where("timestamp >= ?", *filter.StartDate)
}
if filter.EndDate != nil {
query = query.Where("timestamp <= ?", *filter.EndDate)
}
return query
}
@@ -0,0 +1,125 @@
package manager
import (
"context"
"errors"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/shared/db"
"github.com/netbirdio/netbird/management/internals/shared/db/dbtest"
)
func newTestRepository(t *testing.T) (accesslogs.Repository, *db.Conn) {
conn := dbtest.NewConn(t, &accesslogs.AccessLogEntry{})
return NewRepository(conn), conn
}
func newEntry(id, accountID, method string, age time.Duration) *accesslogs.AccessLogEntry {
return &accesslogs.AccessLogEntry{
ID: id,
AccountID: accountID,
Method: method,
Host: "app.example.com",
Path: "/",
StatusCode: 200,
Timestamp: time.Now().Add(-age),
}
}
func TestSqlRepository_ListByAccount(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
for _, entry := range []*accesslogs.AccessLogEntry{
newEntry("a1", "acc-a", "GET", 3*time.Hour),
newEntry("a2", "acc-a", "POST", 2*time.Hour),
newEntry("a3", "acc-a", "GET", time.Hour),
newEntry("b1", "acc-b", "GET", time.Hour),
} {
require.NoError(t, repo.Create(ctx, entry))
}
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 2})
require.NoError(t, err)
assert.EqualValues(t, 3, total)
require.Len(t, logs, 2)
assert.Equal(t, "a3", logs[0].ID)
assert.Equal(t, "a2", logs[1].ID)
method := "GET"
logs, total, err = repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Method: &method, SortOrder: "asc"})
require.NoError(t, err)
assert.EqualValues(t, 2, total)
require.Len(t, logs, 2)
assert.Equal(t, "a1", logs[0].ID)
assert.Equal(t, "a3", logs[1].ID)
}
func TestSqlRepository_DeleteOlderThan(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
require.NoError(t, repo.Create(ctx, newEntry("old", "acc", "GET", 48*time.Hour)))
require.NoError(t, repo.Create(ctx, newEntry("new", "acc", "GET", time.Hour)))
deleted, err := repo.DeleteOlderThan(ctx, time.Now().Add(-24*time.Hour))
require.NoError(t, err)
assert.EqualValues(t, 1, deleted)
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.EqualValues(t, 1, total)
require.Len(t, logs, 1)
assert.Equal(t, "new", logs[0].ID)
}
func TestSqlRepository_CreateInsideTransactionRollsBack(t *testing.T) {
repo, conn := newTestRepository(t)
ctx := context.Background()
failure := errors.New("abort")
err := conn.RunInTx(ctx, func(tx *db.Tx) error {
txRepo := repo.WithTx(tx)
require.NoError(t, txRepo.Create(ctx, newEntry("tx", "acc", "GET", 0)))
_, total, err := txRepo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.EqualValues(t, 1, total)
return failure
})
require.ErrorIs(t, err, failure)
_, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Zero(t, total)
}
func TestSqlRepository_ListByAccount_StatusFilter(t *testing.T) {
repo, _ := newTestRepository(t)
ctx := context.Background()
statusCodes := map[string]int{"l4": 0, "info": 101, "ok": 200, "notfound": 404}
for id, code := range statusCodes {
entry := newEntry(id, "acc", "GET", time.Hour)
entry.StatusCode = code
require.NoError(t, repo.Create(ctx, entry))
}
foreign := newEntry("foreign", "other", "GET", time.Hour)
foreign.StatusCode = 500
require.NoError(t, repo.Create(ctx, foreign))
listIDs := func(status string) []string {
logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Status: &status, SortBy: "status_code", SortOrder: "asc"})
require.NoError(t, err)
require.EqualValues(t, len(logs), total)
ids := make([]string, 0, len(logs))
for _, entry := range logs {
ids = append(ids, entry.ID)
}
return ids
}
assert.Equal(t, []string{"info", "notfound"}, listIDs("failed"))
assert.Equal(t, []string{"ok"}, listIDs("success"))
}
@@ -0,0 +1,18 @@
package accesslogs
import (
"context"
"time"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
//go:generate go tool mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod
// Repository persists reverse proxy access log entries.
type Repository interface {
WithTx(tx *db.Tx) Repository
Create(ctx context.Context, entry *AccessLogEntry) error
ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error)
DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error)
}
@@ -0,0 +1,102 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./repository.go
//
// Generated by this command:
//
// mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod
//
// Package accesslogs is a generated GoMock package.
package accesslogs
import (
context "context"
reflect "reflect"
time "time"
db "github.com/netbirdio/netbird/management/internals/shared/db"
gomock "go.uber.org/mock/gomock"
)
// MockRepository is a mock of Repository interface.
type MockRepository struct {
ctrl *gomock.Controller
recorder *MockRepositoryMockRecorder
isgomock struct{}
}
// MockRepositoryMockRecorder is the mock recorder for MockRepository.
type MockRepositoryMockRecorder struct {
mock *MockRepository
}
// NewMockRepository creates a new mock instance.
func NewMockRepository(ctrl *gomock.Controller) *MockRepository {
mock := &MockRepository{ctrl: ctrl}
mock.recorder = &MockRepositoryMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockRepository) EXPECT() *MockRepositoryMockRecorder {
return m.recorder
}
// Create mocks base method.
func (m *MockRepository) Create(ctx context.Context, entry *AccessLogEntry) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Create", ctx, entry)
ret0, _ := ret[0].(error)
return ret0
}
// Create indicates an expected call of Create.
func (mr *MockRepositoryMockRecorder) Create(ctx, entry any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Create", reflect.TypeOf((*MockRepository)(nil).Create), ctx, entry)
}
// DeleteOlderThan mocks base method.
func (m *MockRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "DeleteOlderThan", ctx, olderThan)
ret0, _ := ret[0].(int64)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// DeleteOlderThan indicates an expected call of DeleteOlderThan.
func (mr *MockRepositoryMockRecorder) DeleteOlderThan(ctx, olderThan any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteOlderThan", reflect.TypeOf((*MockRepository)(nil).DeleteOlderThan), ctx, olderThan)
}
// ListByAccount mocks base method.
func (m *MockRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ListByAccount", ctx, lockStrength, accountID, filter)
ret0, _ := ret[0].([]*AccessLogEntry)
ret1, _ := ret[1].(int64)
ret2, _ := ret[2].(error)
return ret0, ret1, ret2
}
// ListByAccount indicates an expected call of ListByAccount.
func (mr *MockRepositoryMockRecorder) ListByAccount(ctx, lockStrength, accountID, filter any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListByAccount", reflect.TypeOf((*MockRepository)(nil).ListByAccount), ctx, lockStrength, accountID, filter)
}
// WithTx mocks base method.
func (m *MockRepository) WithTx(tx *db.Tx) Repository {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "WithTx", tx)
ret0, _ := ret[0].(Repository)
return ret0
}
// WithTx indicates an expected call of WithTx.
func (mr *MockRepositoryMockRecorder) WithTx(tx any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WithTx", reflect.TypeOf((*MockRepository)(nil).WithTx), tx)
}
@@ -0,0 +1,80 @@
package manager
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/gorilla/mux"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/activity"
nbcontext "github.com/netbirdio/netbird/management/server/context"
nbstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/auth"
)
func TestDeleteDomain_ServiceDependencies(t *testing.T) {
for _, tt := range []struct {
name string
domainName string
serviceHost string
accountID string
enabled bool
protected bool
}{
{"exact", "example.com", "example.com", accountA, true, true},
{"subdomain", "example.com", "deep.app.example.com", accountA, true, true},
{"disabled", "example.com", "app.example.com", accountA, false, true},
// A service is authorized by its own account's registration, so another
// account's service under this namespace is not a dependency of it.
{"other account", "example.com", "app.example.com", accountB, true, false},
{"case and trailing dot", "example.com", "APP.EXAMPLE.COM.", accountA, true, true},
{"suffix boundary", "example.com", "notexample.com", accountA, true, false},
{"literal underscore", "a_b.example.com", "app.a_b.example.com", accountA, true, true},
{"underscore wildcard", "a_b.example.com", "app.axb.example.com", accountA, true, false},
} {
t.Run(tt.name, func(t *testing.T) {
ctx := context.Background()
env := setupDomainTest(t)
events := captureDomainEvents(env)
d, err := env.store.CreateCustomDomain(ctx, accountA, tt.domainName, testCluster, true)
require.NoError(t, err)
svc := &rpservice.Service{
ID: "dependent", AccountID: tt.accountID, Domain: tt.serviceHost,
Enabled: tt.enabled, ProxyCluster: testCluster,
}
require.NoError(t, env.store.CreateService(ctx, svc))
router := mux.NewRouter()
RegisterEndpoints(router, env.manager)
deleteDomain := func() *httptest.ResponseRecorder {
req := httptest.NewRequest(http.MethodDelete, "/domains/"+d.ID, nil)
req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: accountA, UserId: accountAUser})
response := httptest.NewRecorder()
router.ServeHTTP(response, req)
return response
}
response := deleteDomain()
if tt.protected {
require.Equal(t, http.StatusPreconditionFailed, response.Code, "dependent services must block deletion: %s", response.Body.String())
assert.NotContains(t, response.Body.String(), tt.accountID, "the error must not reveal the service's account")
assert.NotNil(t, storedDomain(t, env.store, accountA, d.Domain), "the namespace must remain reserved")
assert.Empty(t, events.get(), "rejected deletion must not emit DomainDeleted")
stored, err := env.store.GetServiceByID(ctx, nbstore.LockingStrengthNone, tt.accountID, svc.ID)
require.NoError(t, err)
assert.Equal(t, svc.Enabled, stored.Enabled, "rejected deletion must preserve the service")
require.NoError(t, env.store.DeleteService(ctx, tt.accountID, svc.ID))
response = deleteDomain()
}
require.Equal(t, http.StatusNoContent, response.Code, "deletion must succeed without dependencies: %s", response.Body.String())
assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "the registration must be deleted")
captured := events.get()
require.Len(t, captured, 1, "only successful deletion may emit an event")
assert.Equal(t, activity.DomainDeleted, captured[0].Activity, "the event must describe the successful deletion")
})
}
}
@@ -357,6 +357,26 @@ func (m Manager) DeriveClusterFromDomain(ctx context.Context, accountID, domain
return "", fmt.Errorf("domain %s does not match any available proxy cluster", domain)
}
// ValidateServiceDomain holds custom domain authorization through a service write transaction.
func (m Manager) ValidateServiceDomain(ctx context.Context, tx nbstore.Store, accountID, serviceDomain, cluster string) error {
if _, ok := ExtractClusterFromFreeDomain(serviceDomain, []string{cluster}); ok {
return nil
}
name, err := nbdomain.FromString(serviceDomain)
if err != nil {
return status.Errorf(status.InvalidArgument, "invalid service domain: %v", err)
}
customDomains, err := tx.LockCustomDomains(ctx, accountID, name)
if err != nil {
return err
}
target, match := extractClusterFromCustomDomains(serviceDomain, customDomains)
if match != customDomainValidated || target != cluster {
return status.Errorf(status.PreconditionFailed, "custom domain authorization changed; retry the service operation")
}
return nil
}
func (m Manager) getClusterAllowList(ctx context.Context, accountID string) ([]string, error) {
byopAddresses, err := m.proxyManager.GetActiveClusterAddressesForAccount(ctx, accountID)
if err != nil {
@@ -99,7 +99,7 @@ func setupDomainTest(t *testing.T) *domainTestEnv {
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", nil, nil)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", testCluster, "127.0.0.1", "", nil, nil)
require.NoError(t, err)
resolver := &stubResolver{cnames: make(map[string]string)}
@@ -11,7 +11,7 @@ import (
// Manager defines the interface for proxy operations
type Manager interface {
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error)
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error)
Disconnect(ctx context.Context, proxyID, sessionID string) error
Heartbeat(ctx context.Context, p *Proxy) error
GetActiveClusterAddresses(ctx context.Context) ([]string, error)
@@ -20,6 +20,7 @@ type Manager interface {
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool
CleanupStale(ctx context.Context, inactivityDuration time.Duration) error
GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error)
CountAccountProxies(ctx context.Context, accountID string) (int64, error)
@@ -8,6 +8,7 @@ import (
"go.opentelemetry.io/otel/metric"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
nbversion "github.com/netbirdio/netbird/version"
)
// store defines the interface for proxy persistence operations
@@ -22,6 +23,7 @@ type store interface {
GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error)
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error)
CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error)
@@ -29,6 +31,8 @@ type store interface {
DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error
}
const minSessionCodeVersion = "0.81.0"
// Manager handles all proxy operations
type Manager struct {
store store
@@ -50,7 +54,7 @@ func NewManager(store store, meter metric.Meter) (*Manager, error) {
// Connect registers a new proxy connection in the database.
// capabilities may be nil for old proxies that do not report them.
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *proxy.Capabilities) (*proxy.Proxy, error) {
now := time.Now()
var caps proxy.Capabilities
if capabilities != nil {
@@ -61,6 +65,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
SessionID: sessionID,
ClusterAddress: clusterAddress,
IPAddress: ipAddress,
Version: truncateVersion(version),
AccountID: accountID,
LastSeen: now,
ConnectedAt: &now,
@@ -78,6 +83,7 @@ func (m *Manager) Connect(ctx context.Context, proxyID, sessionID, clusterAddres
"sessionID": sessionID,
"clusterAddress": clusterAddress,
"ipAddress": ipAddress,
"version": p.Version,
}).Info("proxy connected")
return p, nil
@@ -143,6 +149,22 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string)
return m.store.GetClusterSupportsPrivate(ctx, clusterAddr)
}
// ClusterSupportsSessionCode reports whether all active proxies support session codes.
func (m Manager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
versions, err := m.store.GetActiveProxyVersions(ctx, clusterAddr)
if err != nil || len(versions) == 0 {
return false
}
for _, version := range versions {
if supported, err := nbversion.MeetsMinVersion(minSessionCodeVersion, version); err != nil || !supported {
return false
}
}
return true
}
// CleanupStale removes proxies that haven't sent heartbeat in the specified duration
func (m *Manager) CleanupStale(ctx context.Context, inactivityDuration time.Duration) error {
if err := m.store.CleanupStaleProxies(ctx, inactivityDuration); err != nil {
@@ -184,3 +206,13 @@ func (m *Manager) DeleteAccountCluster(ctx context.Context, clusterAddress, acco
}
return nil
}
// truncateVersion cuts a proxy-reported version to the column width so an
// oversized value cannot fail the save and block the connect.
func truncateVersion(version string) string {
runes := []rune(version)
if len(runes) <= proxy.MaxVersionLength {
return version
}
return string(runes[:proxy.MaxVersionLength])
}
@@ -4,8 +4,10 @@ import (
"context"
"errors"
"fmt"
"strings"
"testing"
"time"
"unicode/utf8"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -20,6 +22,7 @@ type mockStore struct {
updateProxyHeartbeatFunc func(ctx context.Context, p *proxy.Proxy) error
getActiveProxyClusterAddressesFunc func(ctx context.Context) ([]string, error)
getActiveProxyClusterAddressesForAccFunc func(ctx context.Context, accountID string) ([]string, error)
getActiveProxyVersionsFunc func(ctx context.Context, clusterAddress string) ([]string, error)
cleanupStaleProxiesFunc func(ctx context.Context, d time.Duration) error
getProxyByAccountIDFunc func(ctx context.Context, accountID string) (*proxy.Proxy, error)
countProxiesByAccountIDFunc func(ctx context.Context, accountID string) (int64, error)
@@ -102,6 +105,12 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo
func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool {
return nil
}
func (m *mockStore) GetActiveProxyVersions(ctx context.Context, clusterAddress string) ([]string, error) {
if m.getActiveProxyVersionsFunc != nil {
return m.getActiveProxyVersionsFunc(ctx, clusterAddress)
}
return nil, nil
}
func newTestManager(s store) *Manager {
meter := noop.NewMeterProvider().Meter("test")
@@ -112,6 +121,34 @@ func newTestManager(s store) *Manager {
return m
}
func TestClusterSupportsSessionCode(t *testing.T) {
tests := []struct {
name string
versions []string
storeErr error
want bool
}{
{name: "all supported", versions: []string{"0.81.0", "0.81.2"}, want: true},
{name: "one old proxy", versions: []string{"0.81.0", "0.80.0"}},
{name: "missing version", versions: []string{"0.81.0", ""}},
{name: "no active proxies"},
{name: "store error", storeErr: errors.New("db error")},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
s := &mockStore{
getActiveProxyVersionsFunc: func(_ context.Context, _ string) ([]string, error) {
return tt.versions, tt.storeErr
},
}
got := newTestManager(s).ClusterSupportsSessionCode(context.Background(), "cluster.example.com")
assert.Equal(t, tt.want, got)
})
}
}
func TestConnect_WithAccountID(t *testing.T) {
accountID := "acc-123"
@@ -124,7 +161,7 @@ func TestConnect_WithAccountID(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", &accountID, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "0.60.0", &accountID, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
@@ -132,6 +169,7 @@ func TestConnect_WithAccountID(t *testing.T) {
assert.Equal(t, "session-1", savedProxy.SessionID)
assert.Equal(t, "cluster.example.com", savedProxy.ClusterAddress)
assert.Equal(t, "10.0.0.1", savedProxy.IPAddress)
assert.Equal(t, "0.60.0", savedProxy.Version, "reported proxy version should be stored")
assert.Equal(t, &accountID, savedProxy.AccountID)
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
assert.NotNil(t, savedProxy.ConnectedAt)
@@ -147,7 +185,7 @@ func TestConnect_WithoutAccountID(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", nil, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "eu.proxy.netbird.io", "10.0.0.1", "", nil, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
@@ -155,6 +193,29 @@ func TestConnect_WithoutAccountID(t *testing.T) {
assert.Equal(t, proxy.StatusConnected, savedProxy.Status)
}
func TestConnect_TruncatesOversizedVersion(t *testing.T) {
var savedProxy *proxy.Proxy
s := &mockStore{
saveProxyFunc: func(_ context.Context, p *proxy.Proxy) error {
savedProxy = p
return nil
},
}
// Multi-byte runes make sure the cut counts characters, as varchar does,
// and never splits a rune into invalid UTF-8.
version := strings.Repeat("ü", proxy.MaxVersionLength+10)
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", version, nil, nil)
require.NoError(t, err)
require.NotNil(t, savedProxy)
assert.Equal(t, proxy.MaxVersionLength, utf8.RuneCountInString(savedProxy.Version), "stored version should be cut to the column width")
assert.True(t, utf8.ValidString(savedProxy.Version), "stored version should remain valid UTF-8")
assert.True(t, strings.HasPrefix(version, savedProxy.Version), "stored version should be a prefix of the reported one")
}
func TestConnect_StoreError(t *testing.T) {
s := &mockStore{
saveProxyFunc: func(_ context.Context, _ *proxy.Proxy) error {
@@ -163,7 +224,7 @@ func TestConnect_StoreError(t *testing.T) {
}
mgr := newTestManager(s)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", nil, nil)
_, err := mgr.Connect(context.Background(), "proxy-1", "session-1", "cluster.example.com", "10.0.0.1", "", nil, nil)
assert.Error(t, err)
}
@@ -112,19 +112,33 @@ func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
}
// Connect mocks base method.
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
// ClusterSupportsSessionCode mocks base method.
func (m *MockManager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
ret := m.ctrl.Call(m, "ClusterSupportsSessionCode", ctx, clusterAddr)
ret0, _ := ret[0].(bool)
return ret0
}
// ClusterSupportsSessionCode indicates an expected call of ClusterSupportsSessionCode.
func (mr *MockManagerMockRecorder) ClusterSupportsSessionCode(ctx, clusterAddr any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsSessionCode", reflect.TypeOf((*MockManager)(nil).ClusterSupportsSessionCode), ctx, clusterAddr)
}
// Connect mocks base method.
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Connect", ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
ret0, _ := ret[0].(*Proxy)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// Connect indicates an expected call of Connect.
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call {
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, version, accountID, capabilities)
}
// CountAccountProxies mocks base method.
@@ -9,6 +9,9 @@ const (
StatusDisconnected = "disconnected"
)
// MaxVersionLength is the width of the Version column, in characters.
const MaxVersionLength = 255
// Capabilities describes what a proxy can handle, as reported via gRPC.
// Nil fields mean the proxy never reported this capability.
type Capabilities struct {
@@ -31,6 +34,7 @@ type Proxy struct {
SessionID string `gorm:"type:varchar(36)"`
ClusterAddress string `gorm:"type:varchar(255);not null;index:idx_proxy_cluster_status"`
IPAddress string `gorm:"type:varchar(45)"`
Version string `gorm:"type:varchar(255)"`
AccountID *string `gorm:"type:varchar(255);index:idx_proxy_account_id"`
LastSeen time.Time `gorm:"not null;index:idx_proxy_last_seen"`
ConnectedAt *time.Time
@@ -1,6 +1,7 @@
package proxytoken
import (
"context"
"encoding/json"
"net/http"
"time"
@@ -18,13 +19,29 @@ import (
"github.com/netbirdio/netbird/shared/management/status"
)
// RevocationGuard vetoes the tenant-facing revocation of a proxy access
// token. Implementations are supplied by integrations; none is installed by
// default, so every token the caller's account owns may be revoked. It is
// consulted after the ownership check and before the token is revoked. A
// returned status error is written with util.WriteError: its type selects the
// HTTP status and its message is shown to the caller, so it must not carry
// internal detail. Any other error is reported as a generic internal error.
type RevocationGuard interface {
CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error
}
type handler struct {
store store.Store
permissionsManager permissions.Manager
// revocationGuard vetoes revocations. Optional — when nil every owned
// token may be revoked.
revocationGuard RevocationGuard
}
func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, router *mux.Router) {
h := &handler{store: s, permissionsManager: permissionsManager}
// RegisterEndpoints registers the proxy token endpoints. revocationGuard is
// optional; pass nil for no revocation policy.
func RegisterEndpoints(s store.Store, permissionsManager permissions.Manager, revocationGuard RevocationGuard, router *mux.Router) {
h := &handler{store: s, permissionsManager: permissionsManager, revocationGuard: revocationGuard}
router.HandleFunc("/reverse-proxies/proxy-tokens", h.listTokens).Methods("GET", "OPTIONS")
router.HandleFunc("/reverse-proxies/proxy-tokens", h.createToken).Methods("POST", "OPTIONS")
router.HandleFunc("/reverse-proxies/proxy-tokens/{tokenId}", h.revokeToken).Methods("DELETE", "OPTIONS")
@@ -154,6 +171,13 @@ func (h *handler) revokeToken(w http.ResponseWriter, r *http.Request) {
return
}
if h.revocationGuard != nil {
if err := h.revocationGuard.CheckProxyAccessTokenRevocation(ctx, token); err != nil {
util.WriteError(ctx, err, w)
return
}
}
if err := h.store.RevokeProxyAccessToken(ctx, tokenID); err != nil {
util.WriteErrorResponse("failed to revoke token", http.StatusInternalServerError, w)
return
@@ -4,6 +4,7 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"testing"
@@ -22,6 +23,7 @@ import (
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/auth"
"github.com/netbirdio/netbird/shared/management/http/api"
"github.com/netbirdio/netbird/shared/management/status"
)
func authContext(accountID, userID string) context.Context {
@@ -273,3 +275,152 @@ func TestRevokeToken_ManagementWideToken(t *testing.T) {
h.revokeToken(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
}
type revocationGuardFunc func(ctx context.Context, token *types.ProxyAccessToken) error
func (f revocationGuardFunc) CheckProxyAccessTokenRevocation(ctx context.Context, token *types.ProxyAccessToken) error {
return f(ctx, token)
}
func TestRevokeToken_GuardRefuses(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
// No RevokeProxyAccessToken expectation: a refused revocation must not
// reach the store.
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
var checked *types.ProxyAccessToken
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(_ context.Context, token *types.ProxyAccessToken) error {
checked = token
return status.Errorf(status.PreconditionFailed, "token is in use")
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusPreconditionFailed, w.Code)
assert.Contains(t, w.Body.String(), "token is in use")
require.NotNil(t, checked)
assert.Equal(t, "tok-1", checked.ID)
}
func TestRevokeToken_GuardAllows(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
mockStore.EXPECT().RevokeProxyAccessToken(gomock.Any(), "tok-1").Return(nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
return nil
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusOK, w.Code)
}
func TestRevokeToken_GuardFailure(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
accountID := "acc-123"
// No RevokeProxyAccessToken expectation: a guard that cannot decide must
// not let the revocation through.
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &accountID,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), accountID, "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
return errors.New("connection refused")
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext(accountID, "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusInternalServerError, w.Code)
assert.Contains(t, w.Body.String(), "internal server error")
assert.NotContains(t, w.Body.String(), "connection refused")
}
func TestRevokeToken_GuardNotConsultedForForeignToken(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
otherAccount := "acc-other"
mockStore := store.NewMockStore(ctrl)
mockStore.EXPECT().GetProxyAccessTokenByID(gomock.Any(), store.LockingStrengthNone, "tok-1").Return(&types.ProxyAccessToken{
ID: "tok-1",
AccountID: &otherAccount,
}, nil)
permsMgr := permissions.NewMockManager(ctrl)
permsMgr.EXPECT().ValidateUserPermissions(gomock.Any(), "acc-123", "user-1", modules.Services, operations.Delete).Return(true, context.Background(), nil)
// A foreign token must read as not found, not reveal through the guard's
// answer that it belongs to some account's managed proxy.
h := &handler{
store: mockStore,
permissionsManager: permsMgr,
revocationGuard: revocationGuardFunc(func(context.Context, *types.ProxyAccessToken) error {
t.Fatal("guard consulted for a token the caller does not own")
return nil
}),
}
req := httptest.NewRequest("DELETE", "/reverse-proxies/proxy-tokens/tok-1", nil)
req = req.WithContext(authContext("acc-123", "user-1"))
req = mux.SetURLVars(req, map[string]string{"tokenId": "tok-1"})
w := httptest.NewRecorder()
h.revokeToken(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
}
@@ -30,7 +30,7 @@ func withRealDomainManager(t *testing.T, mgr *Manager, testStore store.Store) {
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", nil, nil)
_, err = proxyMgr.Connect(ctx, "proxy-1", "session-1", validationTestCluster, "127.0.0.1", "", nil, nil)
require.NoError(t, err)
accountMgr := &mock_server.MockAccountManager{
@@ -125,3 +125,53 @@ func TestUpdateService_RefusesMoveToUnvalidatedDomain(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, "app.proven.example.com", stored.Domain, "the service must keep its original domain")
}
func TestCreateService_DomainDeletedBeforeWrite(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
d, err := testStore.CreateCustomDomain(ctx, testAccountID, "proven.example.com", validationTestCluster, true)
require.NoError(t, err)
svc := newTestService("app.proven.example.com")
require.NoError(t, mgr.initializeServiceForCreate(ctx, testAccountID, svc))
// Delete after the initial authorization check, before the service transaction starts.
require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID))
err = mgr.persistNewService(ctx, testAccountID, svc)
require.Error(t, err, "an earlier validation result must not authorize a deleted registration")
sErr, ok := status.FromError(err)
require.True(t, ok, "the caller must receive a typed precondition error")
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the service must require current domain authorization")
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
require.NoError(t, err)
assert.Empty(t, services, "the failed write must not leave a service")
}
func TestUpdateService_DomainDeletedBeforeWrite(t *testing.T) {
ctx := context.Background()
mgr, testStore := setupIntegrationTest(t)
withRealDomainManager(t, mgr, testStore)
_, err := testStore.CreateCustomDomain(ctx, testAccountID, "original.example.com", validationTestCluster, true)
require.NoError(t, err)
d, err := testStore.CreateCustomDomain(ctx, testAccountID, "destination.example.com", validationTestCluster, true)
require.NoError(t, err)
svc, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.original.example.com"))
require.NoError(t, err)
moved := svc.Copy()
moved.Domain = "app.destination.example.com"
cluster, err := mgr.resolveEffectiveCluster(ctx, testAccountID, moved)
require.NoError(t, err)
require.NoError(t, testStore.DeleteCustomDomain(ctx, testAccountID, d.ID))
err = testStore.ExecuteInTransaction(ctx, func(tx store.Store) error {
return mgr.executeServiceUpdate(ctx, tx, testAccountID, moved, &serviceUpdateInfo{}, nil, cluster)
})
require.Error(t, err, "a domain deleted after cluster resolution must reject the update")
sErr, ok := status.FromError(err)
require.True(t, ok, "the caller must receive a typed precondition error")
assert.Equal(t, status.PreconditionFailed, sErr.Type(), "the move must require current domain authorization")
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, svc.ID)
require.NoError(t, err)
assert.Equal(t, svc.Domain, stored.Domain, "the service must retain its authorized domain")
}
@@ -74,6 +74,7 @@ const unknownHostPlaceholder = "unknown"
// ClusterDeriver derives the proxy cluster from a domain.
type ClusterDeriver interface {
DeriveClusterFromDomain(ctx context.Context, accountID, domain string) (string, error)
ValidateServiceDomain(ctx context.Context, tx store.Store, accountID, domain, cluster string) error
GetClusterDomains() []string
}
@@ -332,6 +333,9 @@ func (m *Manager) persistNewService(ctx context.Context, accountID string, svc *
}
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
return err
}
if svc.Domain != "" {
if err := m.checkDomainAvailable(ctx, transaction, svc.Domain, ""); err != nil {
return err
@@ -461,6 +465,9 @@ func (m *Manager) persistNewEphemeralService(ctx context.Context, accountID, pee
}
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
return err
}
if err := m.validateEphemeralPreconditions(ctx, transaction, accountID, peerID, svc); err != nil {
return err
}
@@ -622,6 +629,9 @@ func (m *Manager) resolveEffectiveCluster(ctx context.Context, accountID string,
}
func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.Store, accountID string, service *service.Service, updateInfo *serviceUpdateInfo, customPorts *bool, effectiveCluster string) error {
if err := m.validateServiceDomain(ctx, transaction, accountID, service, effectiveCluster); err != nil {
return err
}
existingService, err := transaction.GetServiceByID(ctx, store.LockingStrengthUpdate, accountID, service.ID)
if err != nil {
return err
@@ -677,6 +687,13 @@ func (m *Manager) executeServiceUpdate(ctx context.Context, transaction store.St
return nil
}
func (m *Manager) validateServiceDomain(ctx context.Context, tx store.Store, accountID string, svc *service.Service, cluster string) error {
if m.clusterDeriver == nil {
return nil
}
return m.clusterDeriver.ValidateServiceDomain(ctx, tx, accountID, svc.Domain, cluster)
}
// validateL4PortDiffOnClusterDiff checks if custom L4 ports are configured and validates port changes across clusters.
// It ensures no port changes if custom ports are unsupported for a given cluster and protocol mode.
// Returns an error if validation fails, otherwise returns nil.
@@ -433,8 +433,8 @@ func TestDeletePeerService_SourcePeerValidation(t *testing.T) {
newProxyServer := func(t *testing.T) *nbgrpc.ProxyServiceServer {
t.Helper()
tokenStore := nbgrpc.NewOneTimeTokenStore(context.Background(), testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(context.Background(), testCacheStore(t))
srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(context.Background(), testCacheStore(t))
srv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
return srv
}
@@ -655,6 +655,10 @@ func (d *testClusterDeriver) GetClusterDomains() []string {
return d.domains
}
func (d *testClusterDeriver) ValidateServiceDomain(context.Context, store.Store, string, string, string) error {
return nil
}
const (
testAccountID = "test-account"
testPeerID = "test-peer-1"
@@ -722,8 +726,8 @@ func setupIntegrationTest(t *testing.T) (*Manager, store.Store) {
}
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
@@ -1146,8 +1150,8 @@ func TestDeleteService_DeletesTargets(t *testing.T) {
mockAcct := account.NewMockManager(ctrl)
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, pkceStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, testCacheStore(t))
proxySrv := nbgrpc.NewProxyServiceServer(nil, tokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
proxyController, err := proxymanager.NewGRPCController(proxySrv, noop.NewMeterProvider().Meter(""))
require.NoError(t, err)
+24 -12
View File
@@ -32,17 +32,18 @@ import (
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
networkmapdbfactory "github.com/netbirdio/netbird/management/internals/network_map_db/factory"
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
"github.com/netbirdio/netbird/management/internals/shared/db"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/activity"
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
nbcache "github.com/netbirdio/netbird/management/server/cache"
nbContext "github.com/netbirdio/netbird/management/server/context"
nbhttp "github.com/netbirdio/netbird/management/server/http"
"github.com/netbirdio/netbird/management/server/http/middleware"
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/telemetry"
mgmtProto "github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/ratelimit"
"github.com/netbirdio/netbird/util/crypt"
)
@@ -84,9 +85,20 @@ func (s *BaseServer) CacheStore() nbcache.Store {
})
}
// DBConn opens the database connection shared by the store and the domain repositories.
func (s *BaseServer) DBConn() *db.Conn {
return Create(s, func() *db.Conn {
conn, err := store.OpenConn(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir)
if err != nil {
log.Fatalf("failed to open database connection: %v", err)
}
return conn
})
}
func (s *BaseServer) Store() store.Store {
return Create(s, func() store.Store {
store, err := store.NewStore(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir, s.Metrics(), false)
store, err := store.NewSqlStore(context.Background(), s.DBConn(), s.Metrics(), false)
if err != nil {
log.Fatalf("failed to create store: %v", err)
}
@@ -147,7 +159,7 @@ func (s *BaseServer) EventStore() activity.Store {
func (s *BaseServer) APIHandler() http.Handler {
return Create(s, func() http.Handler {
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager())
httpAPIHandler, err := nbhttp.NewAPIHandler(context.Background(), s.Router(), s.AccountManager(), s.NetworksManager(), s.ResourcesManager(), s.RoutesManager(), s.GroupsManager(), s.GeoLocationManager(), s.AuthManager(), s.Metrics(), s.PermissionsManager(), s.SettingsManager(), s.ZonesManager(), s.RecordsManager(), s.NetworkMapController(), s.IdpManager(), s.ServiceManager(), s.ReverseProxyDomainManager(), s.AccessLogsManager(), s.ReverseProxyGRPCServer(), s.Config.ReverseProxy.TrustedHTTPProxies, s.RateLimiter(), s.IsValidChildAccount, s.AgentNetworkManager(), nil)
if err != nil {
log.Fatalf("failed to create API handler: %v", err)
}
@@ -171,10 +183,10 @@ func (s *BaseServer) Router() *mux.Router {
})
}
func (s *BaseServer) RateLimiter() *middleware.APIRateLimiter {
return Create(s, func() *middleware.APIRateLimiter {
cfg, enabled := middleware.RateLimiterConfigFromEnv()
limiter := middleware.NewAPIRateLimiter(cfg)
func (s *BaseServer) RateLimiter() *ratelimit.APIRateLimiter {
return Create(s, func() *ratelimit.APIRateLimiter {
cfg, enabled := ratelimit.RateLimiterConfigFromEnv()
limiter := ratelimit.NewAPIRateLimiter(cfg)
limiter.SetEnabled(enabled)
return limiter
})
@@ -236,7 +248,7 @@ func (s *BaseServer) GRPCServer() *grpc.Server {
func (s *BaseServer) ReverseProxyGRPCServer() *nbgrpc.ProxyServiceServer {
return Create(s, func() *nbgrpc.ProxyServiceServer {
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.PKCEVerifierStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.SingleUseStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
s.AfterInit(func(s *BaseServer) {
proxyService.SetServiceManager(s.ServiceManager())
proxyService.SetActivityManager(s.ProxyActivityManager())
@@ -293,9 +305,9 @@ func (s *BaseServer) ProxyTokenStore() *nbgrpc.OneTimeTokenStore {
})
}
func (s *BaseServer) PKCEVerifierStore() *nbgrpc.PKCEVerifierStore {
return Create(s, func() *nbgrpc.PKCEVerifierStore {
return nbgrpc.NewPKCEVerifierStore(context.Background(), s.CacheStore())
func (s *BaseServer) SingleUseStore() *nbgrpc.SingleUseStore {
return Create(s, func() *nbgrpc.SingleUseStore {
return nbgrpc.NewSingleUseStore(context.Background(), s.CacheStore())
})
}
@@ -308,7 +320,7 @@ func (s *BaseServer) ProxyActivityManager() proxyactivity.Manager {
func (s *BaseServer) AccessLogsManager() accesslogs.Manager {
return Create(s, func() accesslogs.Manager {
accessLogManager := accesslogsmanager.NewManager(s.Store(), s.PermissionsManager(), s.GeoLocationManager())
accessLogManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(s.DBConn()), s.Store(), s.PermissionsManager(), s.GeoLocationManager())
accessLogManager.StartPeriodicCleanup(
context.Background(),
s.Config.ReverseProxy.AccessLogRetentionDays,
+18
View File
@@ -23,6 +23,8 @@ import (
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/metrics"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/lifecycle"
"github.com/netbirdio/netbird/shared/profiling"
"github.com/netbirdio/netbird/util/wsproxy"
wsproxyserver "github.com/netbirdio/netbird/util/wsproxy/server"
"github.com/netbirdio/netbird/version"
@@ -36,6 +38,8 @@ const (
DefaultSelfHostedDomain = "netbird.selfhosted"
ContainerKeyBaseServer = "baseServer"
applicationName = "management"
)
type Server interface {
@@ -82,6 +86,8 @@ type BaseServer struct {
errCh chan error
wg sync.WaitGroup
cancel context.CancelFunc
lifecycle.StopHandlers
}
// Config holds the configuration parameters for creating a new server
@@ -117,6 +123,9 @@ func NewServer(cfg *Config) *BaseServer {
}
s.container[ContainerKeyBaseServer] = s
stopProfiling := profiling.Start(applicationName)
s.OnStop(stopProfiling)
return s
}
@@ -126,6 +135,14 @@ func (s *BaseServer) AfterInit(fn func(s *BaseServer)) {
// Start begins listening for HTTP requests on the configured address
func (s *BaseServer) Start(ctx context.Context) error {
if err := s.start(ctx); err != nil {
s.RunStopHandlers()
return err
}
return nil
}
func (s *BaseServer) start(ctx context.Context) error {
srvCtx, cancel := context.WithCancel(ctx)
s.cancel = cancel
s.errCh = make(chan error, 4)
@@ -278,6 +295,7 @@ func (s *BaseServer) setupTLS(ctx context.Context) (bool, error) {
func (s *BaseServer) Stop() error {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
defer s.RunStopHandlers()
if s.domainCleanupStop != nil {
s.domainCleanupStop()
}
+121
View File
@@ -0,0 +1,121 @@
package db
import (
"context"
"fmt"
"os"
"runtime"
"strconv"
"time"
"github.com/jackc/pgx/v5/pgxpool"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
)
const (
defaultTransactionTimeout = 5 * time.Minute
connMaxLifetime = time.Hour
connMaxIdleTime = 3 * time.Minute
)
// TxMetrics receives the duration of every committed top-level transaction.
type TxMetrics interface {
CountTransactionDuration(duration time.Duration)
}
// Conn is the database connection shared by all repositories: one gorm handle,
// the pgx pool of a Postgres deployment and the engine they talk to.
type Conn struct {
db *gorm.DB
pool *pgxpool.Pool
engine Engine
txTimeout time.Duration
metrics TxMetrics
}
// NewConn takes ownership of an open gorm handle and pool once it returns
// without error, applying the connection limits and transaction timeout
// configured through the environment.
func NewConn(ctx context.Context, gormDB *gorm.DB, engine Engine, pool *pgxpool.Pool) (*Conn, error) {
sqlDB, err := gormDB.DB()
if err != nil {
return nil, err
}
txTimeout := defaultTransactionTimeout
if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" {
if parsed, err := time.ParseDuration(v); err == nil {
txTimeout = parsed
}
}
log.WithContext(ctx).Infof("Setting transaction timeout to %v", txTimeout)
conns := runtime.NumCPU()
configuredConns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS"))
connsConfigured := err == nil
if connsConfigured {
conns = configuredConns
}
if engine == SqliteStoreEngine {
if connsConfigured {
log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1")
}
conns = 1
}
sqlDB.SetMaxOpenConns(conns)
sqlDB.SetMaxIdleConns(conns)
sqlDB.SetConnMaxLifetime(connMaxLifetime)
sqlDB.SetConnMaxIdleTime(connMaxIdleTime)
log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v",
conns, conns, connMaxLifetime, connMaxIdleTime)
return &Conn{db: gormDB, pool: pool, engine: engine, txTimeout: txTimeout}, nil
}
// DB returns the handle a query must run on: the transaction when tx is set,
// otherwise the shared connection.
func (c *Conn) DB(tx *Tx) *gorm.DB {
if tx != nil {
return tx.db
}
return c.db
}
// Pool returns the pgx pool for read paths that bypass gorm. It is nil on
// engines other than Postgres and inside a transaction, where the pool would
// not see the uncommitted writes.
func (c *Conn) Pool(tx *Tx) *pgxpool.Pool {
if tx != nil {
return nil
}
return c.pool
}
func (c *Conn) Engine() Engine {
return c.engine
}
// SetTxMetrics registers the sink that receives transaction durations.
func (c *Conn) SetTxMetrics(metrics TxMetrics) {
c.metrics = metrics
}
// AutoMigrate creates or updates the tables of the given models.
func (c *Conn) AutoMigrate(models ...any) error {
return c.db.AutoMigrate(models...)
}
// Close releases the gorm connection and the pgx pool.
func (c *Conn) Close() error {
if c.pool != nil {
c.pool.Close()
}
sqlDB, err := c.db.DB()
if err != nil {
return fmt.Errorf("get db: %w", err)
}
return sqlDB.Close()
}
+148
View File
@@ -0,0 +1,148 @@
package db
import (
"context"
"errors"
"path/filepath"
"testing"
"time"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
type testRow struct {
ID uint `gorm:"primaryKey"`
Name string
}
func openTestConn(t *testing.T) *Conn {
t.Helper()
conn, err := OpenSqliteFile(context.Background(), t.TempDir(), SqliteFileName)
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, conn.Close()) })
require.NoError(t, conn.AutoMigrate(&testRow{}))
return conn
}
func countRows(t *testing.T, conn *Conn) int64 {
t.Helper()
var count int64
require.NoError(t, conn.DB(nil).Model(&testRow{}).Count(&count).Error)
return count
}
func TestNewConn_ReadsTransactionTimeoutFromEnv(t *testing.T) {
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "1s")
conn := openTestConn(t)
assert.Equal(t, time.Second, conn.txTimeout)
assert.Equal(t, SqliteStoreEngine, conn.Engine())
}
func TestRunInTx_CommitsOnSuccess(t *testing.T) {
conn := openTestConn(t)
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
})
require.NoError(t, err)
assert.EqualValues(t, 1, countRows(t, conn))
}
func TestRunInTx_RollsBackOnError(t *testing.T) {
conn := openTestConn(t)
failure := errors.New("boom")
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
return failure
})
require.ErrorIs(t, err, failure)
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_RollsBackOnPanic(t *testing.T) {
conn := openTestConn(t)
require.Panics(t, func() {
_ = conn.RunInTx(context.Background(), func(tx *Tx) error {
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
panic("boom")
})
})
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_FailsWhenTimeoutExceeded(t *testing.T) {
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "50ms")
conn := openTestConn(t)
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
time.Sleep(100 * time.Millisecond)
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
})
require.ErrorIs(t, err, context.DeadlineExceeded)
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_ReportsDurationToMetrics(t *testing.T) {
conn := openTestConn(t)
metrics := &recordingMetrics{}
conn.SetTxMetrics(metrics)
require.NoError(t, conn.RunInTx(context.Background(), func(*Tx) error { return nil }))
assert.Equal(t, 1, metrics.calls)
}
func TestConn_DBSelectsTransactionHandle(t *testing.T) {
conn := openTestConn(t)
assert.Same(t, conn.db, conn.DB(nil))
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
assert.Same(t, tx.db, conn.DB(tx))
assert.NotSame(t, conn.db, conn.DB(tx))
return nil
})
require.NoError(t, err)
}
func TestConn_PoolIsUnavailableInsideTransaction(t *testing.T) {
conn := openTestConn(t)
conn.pool = &pgxpool.Pool{}
defer func() { conn.pool = nil }()
assert.Same(t, conn.pool, conn.Pool(nil))
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
assert.Nil(t, conn.Pool(tx))
return nil
})
require.NoError(t, err)
}
type recordingMetrics struct {
calls int
}
func (m *recordingMetrics) CountTransactionDuration(time.Duration) {
m.calls++
}
func TestNewConn_MaxOpenConnsFromEnv(t *testing.T) {
t.Setenv("NB_SQL_MAX_OPEN_CONNS", "7")
gormDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "store.db")), GormConfig())
require.NoError(t, err)
conn, err := NewConn(context.Background(), gormDB, PostgresStoreEngine, nil)
require.NoError(t, err)
t.Cleanup(func() { _ = conn.Close() })
sqlDB, err := conn.DB(nil).DB()
require.NoError(t, err)
assert.Equal(t, 7, sqlDB.Stats().MaxOpenConnections)
sqliteDB, err := openTestConn(t).DB(nil).DB()
require.NoError(t, err)
assert.Equal(t, 1, sqliteDB.Stats().MaxOpenConnections)
}
@@ -0,0 +1,23 @@
package dbtest
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
// NewConn opens a fresh SQLite database in a temporary directory, migrates the
// given models and closes the connection when the test ends. It ignores
// NB_STORE_ENGINE_SQLITE_FILE, so a developer's configured database is never
// touched, and is safe to call from parallel tests.
func NewConn(t testing.TB, models ...any) *db.Conn {
t.Helper()
conn, err := db.OpenSqliteFile(context.Background(), t.TempDir(), db.SqliteFileName)
require.NoError(t, err)
t.Cleanup(func() { _ = conn.Close() })
require.NoError(t, conn.AutoMigrate(models...))
return conn
}
@@ -0,0 +1,31 @@
package dbtest
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
func TestNewConn_IgnoresSqliteFileOverride(t *testing.T) {
override := filepath.Join(t.TempDir(), "configured.db")
t.Setenv("NB_STORE_ENGINE_SQLITE_FILE", override)
conn := NewConn(t)
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
_, err := os.Stat(override)
require.ErrorIs(t, err, os.ErrNotExist)
}
func TestNewConn_Parallel(t *testing.T) {
t.Parallel()
conn := NewConn(t)
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
}
+10
View File
@@ -0,0 +1,10 @@
package db
// Engine identifies the SQL engine behind a Conn.
type Engine string
const (
SqliteStoreEngine Engine = "sqlite"
PostgresStoreEngine Engine = "postgres"
MysqlStoreEngine Engine = "mysql"
)
+12
View File
@@ -0,0 +1,12 @@
package db
// LockingStrength is the row lock a query holds until its transaction ends.
type LockingStrength string
const (
LockingStrengthUpdate LockingStrength = "UPDATE"
LockingStrengthShare LockingStrength = "SHARE"
LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE"
LockingStrengthKeyShare LockingStrength = "KEY SHARE"
LockingStrengthNone LockingStrength = "NONE"
)
+160
View File
@@ -0,0 +1,160 @@
package db
import (
"context"
"fmt"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"time"
"github.com/jackc/pgx/v5/pgxpool"
"gorm.io/driver/mysql"
"gorm.io/driver/postgres"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// SqliteFileName is the default SQLite database file inside the data directory.
const SqliteFileName = "store.db"
// PoolConfig sizes the pgx pool a Postgres deployment uses for the read paths
// that bypass gorm.
type PoolConfig struct {
MaxConns int32
MinConns int32
MaxConnLifetime time.Duration
HealthCheckPeriod time.Duration
}
var DefaultPoolConfig = PoolConfig{
MaxConns: 30,
MinConns: 1,
MaxConnLifetime: 60 * time.Minute,
HealthCheckPeriod: time.Minute,
}
// GormConfig is the configuration every engine is opened with.
func GormConfig() *gorm.Config {
return &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
CreateBatchSize: 400,
}
}
// OpenSqlite opens the SQLite database in dataDir, or the file named by
// NB_STORE_ENGINE_SQLITE_FILE.
func OpenSqlite(ctx context.Context, dataDir string) (*Conn, error) {
storeFile := SqliteFileName
if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" {
storeFile = envFile
}
return OpenSqliteFile(ctx, dataDir, storeFile)
}
// OpenSqliteFile opens the SQLite database storeFile, resolved against dataDir
// when relative. storeFile may carry SQLite URI query parameters.
func OpenSqliteFile(ctx context.Context, dataDir, storeFile string) (*Conn, error) {
// Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc")
filePath, query, hasQuery := strings.Cut(storeFile, "?")
connStr := filePath
if !filepath.IsAbs(filePath) {
connStr = filepath.Join(dataDir, filePath)
}
// Compose query parameters. User-provided ?_busy_timeout (or its mattn alias
// ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at
// most that long on a lock instead of blocking the only Go-side connection.
// mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so
// the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared
// stays the default on non-Windows for the same reason as before.
parsed, _ := url.ParseQuery(query)
var defaults []string
if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" {
defaults = append(defaults, "_busy_timeout=30000")
}
if !hasQuery && runtime.GOOS != "windows" {
// To avoid `The process cannot access the file because it is being used by another process` on Windows
defaults = append(defaults, "cache=shared")
}
parts := defaults
if hasQuery {
parts = append(parts, query)
}
if len(parts) > 0 {
connStr += "?" + strings.Join(parts, "&")
}
gormDB, err := gorm.Open(sqlite.Open(connStr), GormConfig())
if err != nil {
return nil, err
}
return NewConn(ctx, gormDB, SqliteStoreEngine, nil)
}
// OpenPostgres opens a Postgres database through gorm and a pgx pool sized by pool.
func OpenPostgres(ctx context.Context, dsn string, pool PoolConfig) (*Conn, error) {
gormDB, err := gorm.Open(postgres.Open(dsn), GormConfig())
if err != nil {
return nil, err
}
pgxPool, err := newPgxPool(ctx, dsn, pool)
if err != nil {
closeGorm(gormDB)
return nil, err
}
return NewConn(ctx, gormDB, PostgresStoreEngine, pgxPool)
}
// MysqlDSN adds the connection parameters every MySQL handle needs, keeping
// the options already present in dsn.
func MysqlDSN(dsn string) string {
separator := "?"
if strings.Contains(dsn, "?") {
separator = "&"
}
return dsn + separator + "charset=utf8&parseTime=True&loc=Local"
}
// OpenMysql opens a MySQL database through gorm.
func OpenMysql(ctx context.Context, dsn string) (*Conn, error) {
gormDB, err := gorm.Open(mysql.Open(MysqlDSN(dsn)), GormConfig())
if err != nil {
return nil, err
}
return NewConn(ctx, gormDB, MysqlStoreEngine, nil)
}
func newPgxPool(ctx context.Context, dsn string, cfg PoolConfig) (*pgxpool.Pool, error) {
config, err := pgxpool.ParseConfig(dsn)
if err != nil {
return nil, fmt.Errorf("unable to parse database config: %w", err)
}
config.MaxConns = cfg.MaxConns
config.MinConns = cfg.MinConns
config.MaxConnLifetime = cfg.MaxConnLifetime
config.HealthCheckPeriod = cfg.HealthCheckPeriod
pool, err := pgxpool.NewWithConfig(ctx, config)
if err != nil {
return nil, fmt.Errorf("unable to create connection pool: %w", err)
}
if err := pool.Ping(ctx); err != nil {
pool.Close()
return nil, fmt.Errorf("unable to ping database: %w", err)
}
return pool, nil
}
func closeGorm(gormDB *gorm.DB) {
if sqlDB, err := gormDB.DB(); err == nil {
_ = sqlDB.Close()
}
}
@@ -0,0 +1,12 @@
package db
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestMysqlDSN(t *testing.T) {
assert.Equal(t, "user:pw@tcp(host:3306)/db?charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db"))
assert.Equal(t, "user:pw@tcp(host:3306)/db?tls=true&charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db?tls=true"))
}
@@ -0,0 +1,105 @@
package db
import (
"context"
"errors"
"fmt"
"runtime/debug"
"time"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
)
// Tx is an open transaction handed to repository calls; nil means autocommit.
type Tx struct {
db *gorm.DB
}
// RunInTx runs fn in one transaction that commits when fn returns nil and rolls
// back otherwise, bounded by the configured transaction timeout.
func (c *Conn) RunInTx(ctx context.Context, fn func(tx *Tx) error) error {
timeoutCtx, cancel := context.WithTimeout(ctx, c.txTimeout)
defer cancel()
startTime := time.Now()
tx := c.db.WithContext(timeoutCtx).Begin()
if tx.Error != nil {
return tx.Error
}
defer func() {
if r := recover(); r != nil {
tx.Rollback()
panic(r)
}
}()
if err := c.applyStatementTimeouts(tx); err != nil {
tx.Rollback()
return err
}
err := c.withForeignKeyChecksDisabled(tx, func() error {
return fn(&Tx{db: tx})
})
if err != nil {
tx.Rollback()
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction", startTime)
return err
}
if err := tx.Commit().Error; err != nil {
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction commit", startTime)
return err
}
log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime))
if c.metrics != nil {
c.metrics.CountTransactionDuration(time.Since(startTime))
}
return nil
}
func (c *Conn) applyStatementTimeouts(tx *gorm.DB) error {
if c.engine != PostgresStoreEngine {
return nil
}
if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil {
return fmt.Errorf("failed to set statement timeout: %w", err)
}
if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil {
return fmt.Errorf("failed to set lock timeout: %w", err)
}
return nil
}
// withForeignKeyChecksDisabled runs fn with MySQL's FK checks off, which avoids
// deadlocks on MySQL and Aurora without needing SUPER privilege. The setting is
// session-scoped and survives a rollback, so it is turned back on whenever fn
// returns or panics; otherwise the pooled connection would keep it disabled.
func (c *Conn) withForeignKeyChecksDisabled(tx *gorm.DB, fn func() error) (err error) {
if c.engine != MysqlStoreEngine {
return fn()
}
if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil {
return fmt.Errorf("failed to disable FK checks: %w", err)
}
defer func() {
restoreErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error
if restoreErr == nil {
return
}
if err == nil {
err = fmt.Errorf("failed to re-enable FK checks: %w", restoreErr)
return
}
log.WithContext(tx.Statement.Context).Warnf("failed to re-enable FK checks after failed transaction: %v", restoreErr)
}()
return fn()
}
func (c *Conn) logIfTimedOut(ctx, timeoutCtx context.Context, err error, phase string, startTime time.Time) {
if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) {
log.WithContext(ctx).Warnf("%s exceeded %s timeout after %v, stack: %s", phase, c.txTimeout, time.Since(startTime), debug.Stack())
}
}
@@ -14,6 +14,7 @@ const (
baseBlockDuration = 10 * time.Minute // Duration for which a peer is banned after exceeding the reconnection limit
reconnLimitForBan = 30 // Number of reconnections within the reconnTreshold that triggers a ban
metaChangeLimit = 5 // Number of reconnections with different metadata that triggers a ban of one peer
maxBanLevel = 6 // Highest ban level; the ban duration doubles per level up to this one
)
type lfConfig struct {
@@ -21,6 +22,7 @@ type lfConfig struct {
baseBlockDuration time.Duration
reconnLimitForBan int
metaChangeLimit int
maxBanLevel int
}
func initCfg() *lfConfig {
@@ -29,6 +31,7 @@ func initCfg() *lfConfig {
baseBlockDuration: baseBlockDuration,
reconnLimitForBan: reconnLimitForBan,
metaChangeLimit: metaChangeLimit,
maxBanLevel: maxBanLevel,
}
}
@@ -102,11 +105,18 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
return
}
if state.isBanned && now.After(state.banExpiresAt) {
if state.isBanned {
if now.Before(state.banExpiresAt) {
return
}
state.isBanned = false
}
if state.banLevel > 0 && now.Sub(state.lastSeen) > (2*l.cfg.baseBlockDuration) {
quietSince := state.lastSeen
if state.banExpiresAt.After(quietSince) {
quietSince = state.banExpiresAt
}
if state.banLevel > 0 && now.Sub(quietSince) > (2*l.cfg.baseBlockDuration) {
state.banLevel = 0
}
@@ -124,10 +134,17 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
return
}
if now.Sub(state.sessionStart) >= l.cfg.reconnThreshold {
state.sessionStart = now
state.sessionCounter = 0
}
state.sessionCounter++
if state.sessionCounter > l.cfg.reconnLimitForBan && now.Sub(state.sessionStart) < l.cfg.reconnThreshold {
if state.sessionCounter > l.cfg.reconnLimitForBan {
state.isBanned = true
state.banLevel++
if state.banLevel < l.cfg.maxBanLevel {
state.banLevel++
}
backoffFactor := math.Pow(2, float64(state.banLevel-1))
duration := time.Duration(float64(l.cfg.baseBlockDuration) * backoffFactor)
@@ -20,6 +20,7 @@ func testAdvancedCfg() *lfConfig {
baseBlockDuration: 100 * time.Millisecond,
reconnLimitForBan: 3,
metaChangeLimit: 2,
maxBanLevel: 3,
}
}
@@ -157,6 +158,187 @@ func (s *LoginFilterTestSuite) TestMetaChangeIsAllowedAfterWindowResets() {
s.Equal(1, s.filter.logged[pubKey].metaChangeCounter, "meta change counter should reset")
}
func (s *LoginFilterTestSuite) TestReconnectStormAfterQuietPeriodTriggersBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
s.filter.addLogin(pubKey, meta)
s.Require().Contains(s.filter.logged, pubKey)
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
s.filter.addLogin(pubKey, meta)
s.Equal(1, s.filter.logged[pubKey].sessionCounter, "expired window should restart the count")
for i := 1; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
s.filter.addLogin(pubKey, meta)
s.False(s.filter.allowLogin(pubKey, meta))
s.True(s.filter.logged[pubKey].isBanned)
}
func (s *LoginFilterTestSuite) TestReconnectStormAfterBanExpiresTriggersBanAgain() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.Require().True(s.filter.logged[pubKey].isBanned)
expired := time.Now().Add(-(s.filter.cfg.baseBlockDuration + time.Second))
s.filter.logged[pubKey].banExpiresAt = expired
s.filter.logged[pubKey].sessionStart = expired
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(2, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestSlowReconnectsAcrossWindowsDoNotBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
for i := 0; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
}
func (s *LoginFilterTestSuite) TestBanLevelEscalatesWhenStormResumesRightAfterBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
banTime := time.Now().Add(-3 * s.filter.cfg.baseBlockDuration)
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
isBanned: true,
banLevel: 1,
banExpiresAt: time.Now().Add(-time.Millisecond),
sessionStart: banTime,
lastSeen: banTime,
}
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(2, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestBanLevelResetsAfterQuietPeriodFollowingBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
quiet := 2*s.filter.cfg.baseBlockDuration + time.Second
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
banLevel: 2,
banExpiresAt: time.Now().Add(-s.filter.cfg.baseBlockDuration),
lastSeen: time.Now().Add(-2 * quiet),
}
s.filter.addLogin(pubKey, meta)
s.Equal(2, s.filter.logged[pubKey].banLevel, "ban ended more recently than the quiet period")
s.filter.logged[pubKey].banExpiresAt = time.Now().Add(-quiet)
s.filter.logged[pubKey].lastSeen = time.Now().Add(-2 * quiet)
s.filter.addLogin(pubKey, meta)
s.Equal(0, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestBanDurationIsCappedAtMaxLevel() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
maxLevel := s.filter.cfg.maxBanLevel
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
banLevel: maxLevel,
sessionStart: time.Now(),
lastSeen: time.Now(),
}
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(maxLevel, s.filter.logged[pubKey].banLevel)
expected := s.filter.cfg.baseBlockDuration << (maxLevel - 1)
s.InDelta(expected, s.filter.logged[pubKey].banExpiresAt.Sub(s.filter.logged[pubKey].lastSeen), float64(time.Millisecond))
}
func (s *LoginFilterTestSuite) TestEstablishedPeerReconnectingOnceIsAllowed() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
longAgo := time.Now().Add(-time.Hour)
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
sessionCounter: 1,
sessionStart: longAgo,
lastSeen: longAgo,
metaChangeWindowStart: longAgo,
metaChangeCounter: 1,
}
s.True(s.filter.allowLogin(pubKey, meta))
s.filter.addLogin(pubKey, meta)
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
s.Equal(1, s.filter.logged[pubKey].sessionCounter)
}
func (s *LoginFilterTestSuite) TestLoginsDuringActiveBanDoNotExtendIt() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.Require().True(s.filter.logged[pubKey].isBanned)
expiresAt := time.Now().Add(time.Hour)
s.filter.logged[pubKey].banExpiresAt = expiresAt
lastSeen := s.filter.logged[pubKey].lastSeen
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(1, s.filter.logged[pubKey].banLevel)
s.Equal(expiresAt, s.filter.logged[pubKey].banExpiresAt)
s.Equal(lastSeen, s.filter.logged[pubKey].lastSeen)
s.Equal(0, s.filter.logged[pubKey].sessionCounter)
}
func BenchmarkHashingMethods(b *testing.B) {
meta := nbpeer.PeerSystemMeta{
WtVersion: "1.25.1",
@@ -1,55 +0,0 @@
package grpc
import (
"context"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/store"
log "github.com/sirupsen/logrus"
nbcache "github.com/netbirdio/netbird/management/server/cache"
)
// PKCEVerifierStore manages PKCE verifiers for OAuth flows.
// Supports both in-memory and Redis storage via NB_IDP_CACHE_REDIS_ADDRESS env var.
type PKCEVerifierStore struct {
cache nbcache.Store
ctx context.Context
}
// NewPKCEVerifierStore creates a PKCE verifier store using the provided shared cache store.
func NewPKCEVerifierStore(ctx context.Context, cacheStore nbcache.Store) *PKCEVerifierStore {
return &PKCEVerifierStore{
cache: cacheStore,
ctx: ctx,
}
}
// Store saves a PKCE verifier associated with an OAuth state parameter.
// The verifier is stored with the specified TTL and will be automatically deleted after expiration.
func (s *PKCEVerifierStore) Store(state, verifier string, ttl time.Duration) error {
if err := s.cache.Set(s.ctx, state, verifier, store.WithExpiration(ttl)); err != nil {
return fmt.Errorf("failed to store PKCE verifier: %w", err)
}
log.Debugf("Stored PKCE verifier for state (expires in %s)", ttl)
return nil
}
// LoadAndDelete retrieves and removes a PKCE verifier for the given state.
// Returns the verifier and true if found, or empty string and false if not found.
// This enforces single-use semantics for PKCE verifiers.
func (s *PKCEVerifierStore) LoadAndDelete(state string) (string, bool) {
verifier, found, err := s.cache.GetDel(s.ctx, state)
if err != nil {
log.Warnf("Failed to consume PKCE verifier: %v", err)
return "", false
}
if !found {
log.Debug("PKCE verifier not found for state")
return "", false
}
return verifier, true
}
+61 -25
View File
@@ -27,8 +27,6 @@ import (
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/peers"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
@@ -42,6 +40,7 @@ import (
"github.com/netbirdio/netbird/management/server/users"
proxyauth "github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/shared/hash/argon2id"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/proto"
nbstatus "github.com/netbirdio/netbird/shared/management/status"
)
@@ -142,8 +141,8 @@ type ProxyServiceServer struct {
// OIDC configuration for proxy authentication
oidcConfig ProxyOIDCConfig
// Store for PKCE verifiers
pkceVerifierStore *PKCEVerifierStore
// singleUseStore backs both PKCE verifiers and OIDC session exchange codes.
singleUseStore *SingleUseStore
// tokenTTL is the lifetime of one-time tokens generated for proxy
// authentication. Defaults to defaultProxyTokenTTL when zero.
@@ -158,6 +157,13 @@ type ProxyServiceServer struct {
const pkceVerifierTTL = 10 * time.Minute
const sessionCodeTTL = 60 * time.Second
const sessionCodeCacheNamespace = "proxy:session"
// The signed nonce binds the handoff mode without changing the state format.
const sessionCodeNoncePrefix = "code."
const defaultProxyTokenTTL = 5 * time.Minute
const defaultSnapshotBatchSize = 500
@@ -208,13 +214,13 @@ func enforceAccountScope(ctx context.Context, requestAccountID string) error {
}
// NewProxyServiceServer creates a new proxy service server.
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, pkceStore *PKCEVerifierStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, singleUseStore *SingleUseStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
ctx, cancel := context.WithCancel(context.Background())
s := &ProxyServiceServer{
accessLogManager: accessLogMgr,
oidcConfig: oidcConfig,
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
peersManager: peersManager,
usersManager: usersManager,
idpManager: idpManager,
@@ -306,6 +312,16 @@ func (s *ProxyServiceServer) proxyConnectAuthorizer() ProxyConnectAuthorizer {
return s.connectAuthorizer
}
// GenerateSessionCode creates a single-use code for the given session token.
func (s *ProxyServiceServer) GenerateSessionCode(sessionToken string) (code string, ok bool) {
code, err := s.singleUseStore.Generate(sessionCodeCacheNamespace, sessionToken, sessionCodeTTL)
if err != nil {
log.WithError(err).Error("failed to generate proxy session code")
return "", false
}
return code, true
}
// CheckLLMPolicyLimits is the pre-flight policy gate the proxy calls before
// forwarding an LLM request upstream. Delegates to the agent-network selector,
// which scores applicable policies by remaining headroom and returns the
@@ -414,6 +430,7 @@ func (s *ProxyServiceServer) SetProxyController(proxyController proxy.Controller
type proxyConnectParams struct {
proxyID string
address string
version string
capabilities *proto.ProxyCapabilities
}
@@ -424,6 +441,7 @@ func (s *ProxyServiceServer) GetMappingUpdate(req *proto.GetMappingUpdateRequest
return err
}
params.capabilities = req.GetCapabilities()
params.version = req.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
stream: stream,
@@ -457,6 +475,7 @@ func (s *ProxyServiceServer) SyncMappings(stream proto.ProxyService_SyncMappings
return err
}
params.capabilities = init.GetCapabilities()
params.version = init.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
syncStream: stream,
@@ -568,7 +587,7 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params
}
}
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, accountID, caps)
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, params.version, accountID, caps)
if err != nil {
cancel()
if accountID != nil {
@@ -1533,18 +1552,20 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
log.WithContext(ctx).Errorf("failed to get account services: %v", err)
return nil, status.Errorf(codes.FailedPrecondition, "get account services: %v", err)
}
var found bool
var matchedService *rpservice.Service
for _, service := range services {
if service.Domain == redirectURL.Hostname() {
found = true
matchedService = service
break
}
}
if !found {
if matchedService == nil {
log.WithContext(ctx).Debugf("OIDC redirect URL %q does not match any service domain", redirectURL.Hostname())
return nil, status.Errorf(codes.FailedPrecondition, "service not found in store")
}
useSessionCode := s.proxyManager.ClusterSupportsSessionCode(ctx, matchedService.ProxyCluster)
provider, err := oidc.NewProvider(ctx, s.oidcConfig.Issuer)
if err != nil {
log.WithContext(ctx).Errorf("failed to create OIDC provider: %v", err)
@@ -1564,15 +1585,18 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
return nil, status.Errorf(codes.Internal, "generate nonce: %v", err)
}
nonceB64 := base64.URLEncoding.EncodeToString(nonce)
if useSessionCode {
nonceB64 = sessionCodeNoncePrefix + nonceB64
}
// Using an HMAC here to avoid redirection state being modified.
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
payload := redirectURL.String() + "|" + nonceB64
hmacSum := s.generateHMAC(payload)
state := fmt.Sprintf("%s|%s|%s", base64.URLEncoding.EncodeToString([]byte(redirectURL.String())), nonceB64, hmacSum)
codeVerifier := oauth2.GenerateVerifier()
if err := s.pkceVerifierStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil {
if err := s.singleUseStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil {
log.WithContext(ctx).Errorf("failed to store PKCE verifier: %v", err)
return nil, status.Errorf(codes.Internal, "store PKCE verifier: %v", err)
}
@@ -1609,15 +1633,12 @@ func (s *ProxyServiceServer) generateHMAC(input string) string {
return hex.EncodeToString(mac.Sum(nil))
}
// ValidateState validates the state parameter from an OAuth callback.
// Returns the original redirect URL if valid, or an error if invalid.
// The HMAC is verified before consuming the PKCE verifier to prevent
// an attacker from invalidating a legitimate user's auth flow.
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, err error) {
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
// ValidateState validates and consumes an OIDC state.
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, useSessionCode bool, err error) {
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
parts := strings.Split(state, "|")
if len(parts) != 3 {
return "", "", errors.New("invalid state format")
return "", "", false, errors.New("invalid state format")
}
encodedURL := parts[0]
@@ -1626,7 +1647,7 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
redirectURLBytes, err := base64.URLEncoding.DecodeString(encodedURL)
if err != nil {
return "", "", fmt.Errorf("invalid state encoding: %w", err)
return "", "", false, fmt.Errorf("invalid state encoding: %w", err)
}
redirectURL = string(redirectURLBytes)
@@ -1634,16 +1655,17 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
expectedHMAC := s.generateHMAC(payload)
if !hmac.Equal([]byte(providedHMAC), []byte(expectedHMAC)) {
return "", "", errors.New("invalid state signature")
return "", "", false, errors.New("invalid state signature")
}
useSessionCode = strings.HasPrefix(nonce, sessionCodeNoncePrefix)
// Consume the PKCE verifier only after HMAC validation passes.
verifier, ok := s.pkceVerifierStore.LoadAndDelete(state)
verifier, ok := s.singleUseStore.LoadAndDelete(state)
if !ok {
return "", "", errors.New("no verifier for state")
return "", "", false, errors.New("no verifier for state")
}
return verifier, redirectURL, nil
return verifier, redirectURL, useSessionCode, nil
}
// Denied reasons reported to the proxy when access is refused because of the
@@ -1848,7 +1870,20 @@ func (s *ProxyServiceServer) getAccountServiceByDomain(ctx context.Context, acco
// ValidateSession validates a session token and checks if the user has access to the domain.
func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.ValidateSessionRequest) (*proto.ValidateSessionResponse, error) {
domain := req.GetDomain()
sessionToken := req.GetSessionToken()
sessionToken := req.GetSessionToken() //nolint:staticcheck
// A one-time code from the OIDC callback is redeemed here for the durable
// token, so the token never travels in a redirect URL. The redeemed token
// is returned to the proxy (mintedToken) to install as the session cookie.
mintedToken := ""
if code := req.GetSessionCode(); code != "" {
redeemed, found := s.singleUseStore.LoadAndDelete(singleUseCacheKey(sessionCodeCacheNamespace, code))
if !found {
return deniedSessionResponse("invalid or expired session code"), nil
}
sessionToken = redeemed
mintedToken = redeemed
}
if domain == "" || sessionToken == "" {
return deniedSessionResponse("missing domain or session_token"), nil
@@ -1921,6 +1956,7 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
UserEmail: user.Email,
PeerGroupIds: groupIDs,
PeerGroupNames: groupNames,
SessionToken: mintedToken,
}, nil
}
@@ -0,0 +1,93 @@
package grpc
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/shared/management/proto"
)
const (
versionTestProxyID = "proxy-a"
versionTestCluster = "cluster.example.com"
versionTestVersion = "0.60.0"
)
// hangupStream cancels its context on the first Send, emulating a proxy that
// disconnects right after receiving the initial snapshot. The legacy stream
// carries no proxy-to-management messages, so this is the only way for
// GetMappingUpdate to return.
type hangupStream struct {
recordingStream
ctx context.Context
cancel context.CancelFunc
}
func (s *hangupStream) Send(m *proto.GetMappingUpdateResponse) error {
s.cancel()
return s.recordingStream.Send(m)
}
func (s *hangupStream) Context() context.Context { return s.ctx }
// newVersionTestServer wires a server whose proxy manager only accepts a
// Connect carrying versionTestVersion, so a dropped or mangled version fails
// the test as an unexpected call.
func newVersionTestServer(t *testing.T) *ProxyServiceServer {
t.Helper()
ctrl := gomock.NewController(t)
svcMgr := rpservice.NewMockManager(ctrl)
svcMgr.EXPECT().GetGlobalServices(gomock.Any()).Return(nil, nil)
proxyMgr := proxy.NewMockManager(ctrl)
proxyMgr.EXPECT().
Connect(gomock.Any(), versionTestProxyID, gomock.Any(), versionTestCluster, gomock.Any(), versionTestVersion, gomock.Any(), gomock.Any()).
Return(&proxy.Proxy{ID: versionTestProxyID, Version: versionTestVersion}, nil)
proxyMgr.EXPECT().Disconnect(gomock.Any(), versionTestProxyID, gomock.Any()).Return(nil)
s := newSnapshotTestServer(t, 10)
s.serviceManager = svcMgr
s.proxyManager = proxyMgr
return s
}
func TestSyncMappings_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
// The init carries the version, the ack acknowledges the empty snapshot,
// and the exhausted fake stream then ends the RPC.
stream := &syncRecordingStream{
recvMsgs: []*proto.SyncMappingsRequest{
{Msg: &proto.SyncMappingsRequest_Init{Init: &proto.SyncMappingsInit{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}}},
ackMsg(),
},
}
err := s.SyncMappings(stream)
require.ErrorContains(t, err, "no more recv messages")
}
func TestGetMappingUpdate_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
stream := &hangupStream{ctx: ctx, cancel: cancel}
err := s.GetMappingUpdate(&proto.GetMappingUpdateRequest{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}, stream)
require.ErrorIs(t, err, context.Canceled)
}
+43 -19
View File
@@ -129,11 +129,11 @@ func drainEmpty(ch chan *proto.GetMappingUpdateResponse) bool {
func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -186,11 +186,11 @@ func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) {
func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -220,11 +220,11 @@ func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
func TestSendServiceUpdate_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -272,13 +272,13 @@ func generateState(s *ProxyServiceServer, redirectURL string) string {
func TestOAuthState_NeverTheSame(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
redirectURL := "https://app.example.com/callback"
@@ -300,20 +300,20 @@ func TestOAuthState_NeverTheSame(t *testing.T) {
func TestValidateState_RejectsOldTwoPartFormat(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
// Old format had only 2 parts: base64(url)|hmac
err := s.pkceVerifierStore.Store("base64url|hmac", "test", 10*time.Minute)
err := s.singleUseStore.Store("base64url|hmac", "test", 10*time.Minute)
require.NoError(t, err)
_, _, err = s.ValidateState("base64url|hmac")
_, _, _, err = s.ValidateState("base64url|hmac")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid state format")
}
@@ -372,24 +372,48 @@ func TestEnforceAccountScope_AllowsNoTokenInContext(t *testing.T) {
func TestValidateState_RejectsInvalidHMAC(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
// Store with tampered HMAC
err := s.pkceVerifierStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
err := s.singleUseStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
require.NoError(t, err)
_, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
_, _, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid state signature")
}
func TestSessionCodeCannotConsumeOIDCState(t *testing.T) {
const verifier = "pkce-verifier"
store := NewSingleUseStore(context.Background(), testCacheStore(t))
server := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
singleUseStore: store,
}
state := generateState(server, "https://service.example.com/callback")
require.NoError(t, store.Store(state, verifier, time.Minute))
response, err := server.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
SessionCode: state,
})
require.NoError(t, err)
assert.False(t, response.GetValid())
gotVerifier, _, _, err := server.ValidateState(state)
require.NoError(t, err)
assert.Equal(t, verifier, gotVerifier)
}
func TestSendServiceUpdateToCluster_FiltersOnCapability(t *testing.T) {
tokenStore := NewOneTimeTokenStore(context.Background(), testCacheStore(t))
@@ -0,0 +1,67 @@
package grpc
import (
"context"
"crypto/rand"
"encoding/base64"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/store"
log "github.com/sirupsen/logrus"
nbcache "github.com/netbirdio/netbird/management/server/cache"
)
// SingleUseStore stores short-lived values that can be retrieved only once.
type SingleUseStore struct {
cache nbcache.Store
ctx context.Context
}
// NewSingleUseStore creates a single-use value store over the shared cache.
func NewSingleUseStore(ctx context.Context, cacheStore nbcache.Store) *SingleUseStore {
return &SingleUseStore{
cache: cacheStore,
ctx: ctx,
}
}
// Store saves value under key with the given TTL, after which it is evicted.
func (s *SingleUseStore) Store(key, value string, ttl time.Duration) error {
if err := s.cache.Set(s.ctx, key, value, store.WithExpiration(ttl)); err != nil {
return fmt.Errorf("store single-use value: %w", err)
}
return nil
}
// Generate stores a value under a namespaced random key and returns the random key.
func (s *SingleUseStore) Generate(namespace, value string, ttl time.Duration) (string, error) {
buf := make([]byte, 32)
if _, err := rand.Read(buf); err != nil {
return "", fmt.Errorf("generate single-use key: %w", err)
}
key := base64.RawURLEncoding.EncodeToString(buf)
if err := s.Store(singleUseCacheKey(namespace, key), value, ttl); err != nil {
return "", err
}
return key, nil
}
func singleUseCacheKey(namespace, key string) string {
return namespace + ":" + key
}
// LoadAndDelete retrieves and removes the value for a key.
func (s *SingleUseStore) LoadAndDelete(key string) (string, bool) {
value, found, err := s.cache.GetDel(s.ctx, key)
if err != nil {
log.Warnf("failed to consume single-use value: %v", err)
return "", false
}
if !found {
return "", false
}
return value, true
}
@@ -6,7 +6,7 @@ import (
"time"
)
func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
func TestSingleUseStoreLoadAndDelete(t *testing.T) {
const (
state = "state"
verifier = "verifier"
@@ -14,7 +14,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
)
t.Run("exactly one concurrent caller consumes the verifier", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, time.Minute); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -50,7 +50,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("replayed state is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, time.Minute); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -64,7 +64,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("unknown state is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if got, found := store.LoadAndDelete("never-stored"); found {
t.Fatalf("unknown state should not resolve, got %q", got)
@@ -72,7 +72,7 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
})
t.Run("expired verifier is rejected", func(t *testing.T) {
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
store := NewSingleUseStore(context.Background(), testCacheStore(t))
if err := store.Store(state, verifier, 50*time.Millisecond); err != nil {
t.Fatalf("couldn't store PKCE verifier: %s", err)
}
@@ -83,3 +83,40 @@ func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
}
})
}
func TestSingleUseStore_GenerateAndConsumeOnce(t *testing.T) {
const namespace = "test"
s := NewSingleUseStore(context.Background(), testCacheStore(t))
key, err := s.Generate(namespace, "the-value", time.Minute)
if err != nil {
t.Fatalf("generate: %v", err)
}
if key == "" || key == "the-value" {
t.Fatalf("unexpected key %q", key)
}
value, found := s.LoadAndDelete(singleUseCacheKey(namespace, key))
if !found || value != "the-value" {
t.Fatalf("expected to load the stored value, got %q found=%v", value, found)
}
if _, found := s.LoadAndDelete(singleUseCacheKey(namespace, key)); found {
t.Fatal("value must be consumed on first LoadAndDelete")
}
}
func TestSingleUseStore_GenerateUniqueKeys(t *testing.T) {
s := NewSingleUseStore(context.Background(), testCacheStore(t))
a, err := s.Generate("test", "v", time.Minute)
if err != nil {
t.Fatalf("generate a: %v", err)
}
b, err := s.Generate("test", "v", time.Minute)
if err != nil {
t.Fatalf("generate b: %v", err)
}
if a == b {
t.Fatal("generated keys must be distinct")
}
}
@@ -40,9 +40,9 @@ func setupValidateSessionTest(t *testing.T) *validateSessionTestSetup {
proxyManager := &testValidateSessionProxyManager{}
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
proxyService := NewProxyServiceServer(nil, tokenStore, pkceStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil)
proxyService := NewProxyServiceServer(nil, tokenStore, singleUseStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil)
proxyService.SetServiceManager(serviceManager)
createTestProxies(t, ctx, testStore)
@@ -570,7 +570,7 @@ func (m *testValidateSessionServiceManager) DeleteAccountCluster(_ context.Conte
type testValidateSessionProxyManager struct{}
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
return nil, nil
}
@@ -634,6 +634,10 @@ func (m *testValidateSessionProxyManager) ClusterSupportsPrivate(_ context.Conte
return nil
}
func (m *testValidateSessionProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
return false
}
type testValidateSessionUsersManager struct {
store store.Store
}
@@ -662,3 +666,47 @@ func (m *testValidateSessionUsersManager) GetUserWithGroups(ctx context.Context,
}
return user, groups, nil
}
func TestValidateSession_RedeemsSessionCode(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "testProxyId")
require.NoError(t, err)
token := createSessionToken(t, proxy.SessionPrivateKey, "allowedUserId", "test-proxy.example.com")
code, ok := setup.proxyService.GenerateSessionCode(token)
require.True(t, ok)
require.NotEqual(t, token, code, "code must not be the token itself")
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: code,
})
require.NoError(t, err)
assert.True(t, resp.Valid, "redeemed code should authorize the user")
assert.Equal(t, "allowedUserId", resp.UserId)
assert.Equal(t, token, resp.GetSessionToken(), "response must carry the durable token for the cookie")
// Single-use: the same code must not redeem again.
resp2, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: code,
})
require.NoError(t, err)
assert.False(t, resp2.Valid, "a consumed code must be rejected")
assert.Empty(t, resp2.GetSessionToken())
}
func TestValidateSession_InvalidSessionCode(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: "does-not-exist",
})
require.NoError(t, err)
assert.False(t, resp.Valid)
assert.Empty(t, resp.GetSessionToken())
}
+19 -1
View File
@@ -3,9 +3,11 @@ package auth
import (
"context"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"strings"
"time"
)
@@ -52,6 +54,22 @@ func (s *SessionStore) RegisterToken(ctx context.Context, token string, expiresA
}
func hashToken(token string) string {
sum := sha256.Sum256([]byte(token))
sum := sha256.Sum256([]byte(canonicalizeToken(token)))
return hex.EncodeToString(sum[:])
}
// canonicalizeToken re-encodes the JWT signature segment so noncanonical
// spellings of the same signature map to one stable cache key.
func canonicalizeToken(token string) string {
parts := strings.Split(token, ".")
if len(parts) != 3 {
return token
}
sig, err := base64.RawURLEncoding.DecodeString(parts[2])
if err != nil {
return token
}
return parts[0] + "." + parts[1] + "." + base64.RawURLEncoding.EncodeToString(sig)
}
+51
View File
@@ -2,10 +2,15 @@ package auth
import (
"context"
"crypto/rand"
"crypto/rsa"
"encoding/base64"
"errors"
"strings"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -131,3 +136,49 @@ func TestHashToken_StableAndDoesNotLeak(t *testing.T) {
assert.Len(t, a, 64, "sha256 hex must be 64 chars")
assert.NotContains(t, a, "tokenA", "raw token must not appear in hash")
}
func TestSessionStore_NoncanonicalSpellingIsRejectedAsReplay(t *testing.T) {
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)
token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims{
"sub": "user",
"exp": time.Now().Add(time.Hour).Unix(),
})
canonical, err := token.SignedString(privateKey)
require.NoError(t, err)
parts := strings.Split(canonical, ".")
require.Len(t, parts, 3)
// A 256-byte RSA signature (256 mod 3 == 1) leaves unused bits in the final
// base64url character; flip one without changing the decoded signature.
const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_"
last := strings.IndexByte(alphabet, parts[2][len(parts[2])-1])
require.GreaterOrEqual(t, last, 0)
require.Equal(t, 0, last&3, "unexpected canonical RSA signature encoding")
equivalentSig := parts[2][:len(parts[2])-1] + string(alphabet[last|1])
equivalent := parts[0] + "." + parts[1] + "." + equivalentSig
require.NotEqual(t, canonical, equivalent, "spellings must differ as strings")
// Same decoded signature bytes, so they verify as the same JWT.
canonicalSig, err := base64.RawURLEncoding.DecodeString(parts[2])
require.NoError(t, err)
altSig, err := base64.RawURLEncoding.DecodeString(equivalentSig)
require.NoError(t, err)
require.Equal(t, canonicalSig, altSig, "spellings must decode to identical signature bytes")
// The replay-cache key must be identical for both spellings.
assert.Equal(t, hashToken(canonical), hashToken(equivalent),
"noncanonical spelling must map to the same replay-cache key")
s := newTestSessionStore(t)
ctx := context.Background()
exp := time.Now().Add(time.Hour)
require.NoError(t, s.RegisterToken(ctx, canonical, exp), "first claim should succeed")
err = s.RegisterToken(ctx, equivalent, exp)
require.Error(t, err, "alternate spelling must be treated as a replay")
assert.ErrorIs(t, err, ErrTokenAlreadyUsed)
}
+4 -3
View File
@@ -58,10 +58,11 @@ import (
"github.com/netbirdio/netbird/management/server/networks/resources"
"github.com/netbirdio/netbird/management/server/networks/routers"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/shared/ratelimit"
)
// NewAPIHandler creates the Management service HTTP API handler registering all the available endpoints.
func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *middleware.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager) (http.Handler, error) {
func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager account.Manager, networksManager nbnetworks.Manager, resourceManager resources.Manager, routerManager routers.Manager, groupsManager nbgroups.Manager, LocationManager geolocation.Geolocation, authManager auth.Manager, appMetrics telemetry.AppMetrics, permissionsManager permissions.Manager, settingsManager settings.Manager, zManager zones.Manager, rManager records.Manager, networkMapController network_map.Controller, idpManager idpmanager.Manager, serviceManager service.Manager, reverseProxyDomainManager *manager.Manager, reverseProxyAccessLogsManager accesslogs.Manager, proxyGRPCServer *nbgrpc.ProxyServiceServer, trustedHTTPProxies []netip.Prefix, rateLimiter *ratelimit.APIRateLimiter, isValidChildAccount middleware.IsValidChildAccountFunc, agentNetworkManager agentnetwork.Manager, proxyTokenRevocationGuard proxytoken.RevocationGuard) (http.Handler, error) {
// Register bypass paths for unauthenticated endpoints
if err := bypass.AddBypassPath("/api/instance"); err != nil {
@@ -84,7 +85,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou
if rateLimiter == nil {
log.Warn("NewAPIHandler: nil rate limiter, rate limiting disabled")
rateLimiter = middleware.NewAPIRateLimiter(nil)
rateLimiter = ratelimit.NewAPIRateLimiter(nil)
rateLimiter.SetEnabled(false)
}
@@ -135,7 +136,7 @@ func NewAPIHandler(ctx context.Context, router *mux.Router, accountManager accou
reverseproxymanager.RegisterEndpoints(serviceManager, *reverseProxyDomainManager, reverseProxyAccessLogsManager, permissionsManager, router)
}
proxytoken.RegisterEndpoints(accountManager.GetStore(), permissionsManager, router)
proxytoken.RegisterEndpoints(accountManager.GetStore(), permissionsManager, proxyTokenRevocationGuard, router)
// Register OAuth callback handler for proxy authentication
if proxyGRPCServer != nil {
+16 -7
View File
@@ -16,21 +16,21 @@ import (
"golang.org/x/oauth2"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/http/middleware"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/shared/ratelimit"
)
// AuthCallbackHandler handles OAuth callbacks for proxy authentication.
type AuthCallbackHandler struct {
proxyService *nbgrpc.ProxyServiceServer
rateLimiter *middleware.APIRateLimiter
rateLimiter *ratelimit.APIRateLimiter
trustedProxies []netip.Prefix
}
// NewAuthCallbackHandler creates a new OAuth callback handler.
func NewAuthCallbackHandler(proxyService *nbgrpc.ProxyServiceServer, trustedProxies []netip.Prefix) *AuthCallbackHandler {
rateLimiterConfig := &middleware.RateLimiterConfig{
rateLimiterConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 10,
Burst: 15,
CleanupInterval: 5 * time.Minute,
@@ -39,7 +39,7 @@ func NewAuthCallbackHandler(proxyService *nbgrpc.ProxyServiceServer, trustedProx
return &AuthCallbackHandler{
proxyService: proxyService,
rateLimiter: middleware.NewAPIRateLimiter(rateLimiterConfig),
rateLimiter: ratelimit.NewAPIRateLimiter(rateLimiterConfig),
trustedProxies: trustedProxies,
}
}
@@ -59,7 +59,7 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
state := r.URL.Query().Get("state")
codeVerifier, originalURL, err := h.proxyService.ValidateState(state)
codeVerifier, originalURL, useSessionCode, err := h.proxyService.ValidateState(state)
if err != nil {
log.WithError(err).Error("OAuth callback state validation failed")
http.Error(w, "Invalid state parameter", http.StatusBadRequest)
@@ -119,10 +119,19 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
redirectURL.Scheme = "https"
query := redirectURL.Query()
query.Set("session_token", sessionToken)
if useSessionCode {
code, ok := h.proxyService.GenerateSessionCode(sessionToken)
if !ok {
http.Error(w, "Failed to create session", http.StatusInternalServerError)
return
}
query.Set("session_code", code)
} else {
query.Set("session_token", sessionToken)
}
redirectURL.RawQuery = query.Encode()
log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user with session token")
log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user to proxy")
http.Redirect(w, r, redirectURL.String(), http.StatusFound)
}
@@ -181,6 +181,10 @@ func (m *testAccessLogManager) GetAllAccessLogs(_ context.Context, _, _ string,
}
func setupAuthCallbackTest(t *testing.T) *testSetup {
return setupAuthCallbackTestWithProxyManager(t, testSessionCodeManager{})
}
func setupAuthCallbackTestWithProxyManager(t *testing.T, proxyManager nbproxy.Manager) *testSetup {
t.Helper()
ctx := context.Background()
@@ -197,7 +201,7 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
require.NoError(t, err)
tokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore)
pkceStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore)
usersManager := users.NewManager(testStore)
@@ -212,12 +216,12 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
proxyService := nbgrpc.NewProxyServiceServer(
&testAccessLogManager{},
tokenStore,
pkceStore,
singleUseStore,
oidcConfig,
nil,
usersManager,
nil,
nil,
proxyManager,
nil,
)
@@ -242,6 +246,15 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
}
}
type testSessionCodeManager struct {
nbproxy.Manager
supported bool
}
func (m testSessionCodeManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
return m.supported
}
func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store.Store) {
t.Helper()
@@ -252,10 +265,11 @@ func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store
privKey := base64.StdEncoding.EncodeToString(priv)
testProxy := &service.Service{
ID: "testProxyId",
AccountID: "testAccountId",
Name: "Test Proxy",
Domain: "test-proxy.example.com",
ID: "testProxyId",
AccountID: "testAccountId",
Name: "Test Proxy",
Domain: "test-proxy.example.com",
ProxyCluster: "cluster.example.com",
Targets: []*service.Target{{
Path: strPtr("/"),
Host: "localhost",
@@ -512,29 +526,56 @@ func createTestState(t *testing.T, ps *nbgrpc.ProxyServiceServer, redirectURL st
}
func TestAuthCallback_UserAllowedToLogin(t *testing.T) {
setup := setupAuthCallbackTest(t)
defer setup.cleanup()
tests := []struct {
name string
manager nbproxy.Manager
wantParam string
absentParam string
}{
{name: "legacy proxy", manager: testSessionCodeManager{}, wantParam: "session_token", absentParam: "session_code"},
{name: "compatible proxy", manager: testSessionCodeManager{supported: true}, wantParam: "session_code", absentParam: "session_token"},
}
setup.oidcServer.tokenSubject = "allowedUserId"
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
setup := setupAuthCallbackTestWithProxyManager(t, tt.manager)
defer setup.cleanup()
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
setup.oidcServer.tokenSubject = "allowedUserId"
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
rec := httptest.NewRecorder()
setup.router.ServeHTTP(rec, req)
require.Equal(t, http.StatusFound, rec.Code)
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
rec := httptest.NewRecorder()
location, err := url.Parse(rec.Header().Get("Location"))
require.NoError(t, err)
require.Equal(t, "test-proxy.example.com", location.Host)
require.NotEmpty(t, location.Query().Get(tt.wantParam))
require.Empty(t, location.Query().Get(tt.absentParam))
require.Empty(t, location.Query().Get("error"))
setup.router.ServeHTTP(rec, req)
if tt.wantParam == "session_code" {
code := location.Query().Get("session_code")
response, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: location.Hostname(),
SessionCode: code,
})
require.NoError(t, err)
require.True(t, response.GetValid())
require.NotEmpty(t, response.GetSessionToken())
require.NotEqual(t, code, response.GetSessionToken())
require.Equal(t, http.StatusFound, rec.Code)
location := rec.Header().Get("Location")
require.NotEmpty(t, location)
parsedLocation, err := url.Parse(location)
require.NoError(t, err)
require.Equal(t, "test-proxy.example.com", parsedLocation.Host)
require.NotEmpty(t, parsedLocation.Query().Get("session_token"), "Should include session token")
require.Empty(t, parsedLocation.Query().Get("error"), "Should not have error parameter")
replayed, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: location.Hostname(),
SessionCode: code,
})
require.NoError(t, err)
require.False(t, replayed.GetValid())
require.Empty(t, replayed.GetSessionToken())
}
})
}
}
// TestAuthCallback_UserDeniedByAccountStatus asserts that a user whose account
@@ -11,15 +11,15 @@ import (
"github.com/netbirdio/netbird/management/server/account"
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/management/server/http/middleware"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/http/api"
"github.com/netbirdio/netbird/shared/management/http/util"
"github.com/netbirdio/netbird/shared/management/status"
"github.com/netbirdio/netbird/shared/ratelimit"
)
// publicInviteRateLimiter limits public invite requests by IP address to prevent brute-force attacks
var publicInviteRateLimiter = middleware.NewAPIRateLimiter(&middleware.RateLimiterConfig{
var publicInviteRateLimiter = ratelimit.NewAPIRateLimiter(&ratelimit.RateLimiterConfig{
RequestsPerMinute: 10, // 10 attempts per minute per IP
Burst: 5, // Allow burst of 5 requests
CleanupInterval: 10 * time.Minute,
@@ -18,6 +18,7 @@ import (
"github.com/netbirdio/netbird/shared/auth"
"github.com/netbirdio/netbird/shared/management/http/util"
"github.com/netbirdio/netbird/shared/management/status"
"github.com/netbirdio/netbird/shared/ratelimit"
)
type EnsureAccountFunc func(ctx context.Context, userAuth auth.UserAuth) (string, string, error)
@@ -33,7 +34,7 @@ type AuthMiddleware struct {
ensureAccount EnsureAccountFunc
getUserFromUserAuth GetUserFromUserAuthFunc
syncUserJWTGroups SyncUserJWTGroupsFunc
rateLimiter *APIRateLimiter
rateLimiter *ratelimit.APIRateLimiter
patUsageTracker *PATUsageTracker
isValidChildAccount IsValidChildAccountFunc
}
@@ -44,7 +45,7 @@ func NewAuthMiddleware(
ensureAccount EnsureAccountFunc,
syncUserJWTGroups SyncUserJWTGroupsFunc,
getUserFromUserAuth GetUserFromUserAuthFunc,
rateLimiter *APIRateLimiter,
rateLimiter *ratelimit.APIRateLimiter,
meter metric.Meter,
isValidChildAccount IsValidChildAccountFunc,
) *AuthMiddleware {
@@ -18,6 +18,7 @@ import (
"github.com/netbirdio/netbird/management/server/util"
nbauth "github.com/netbirdio/netbird/shared/auth"
nbjwt "github.com/netbirdio/netbird/shared/auth/jwt"
"github.com/netbirdio/netbird/shared/ratelimit"
)
const (
@@ -196,7 +197,7 @@ func TestAuthMiddleware_Handler(t *testing.T) {
GetPATInfoFunc: mockGetAccountInfoFromPAT,
}
disabledLimiter := NewAPIRateLimiter(nil)
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
disabledLimiter.SetEnabled(false)
authMiddleware := NewAuthMiddleware(
mockAuth,
@@ -260,7 +261,7 @@ func TestAuthMiddleware_SyncUserJWTGroupsDetachedFromRequestCancellation(t *test
GetPATInfoFunc: mockGetAccountInfoFromPAT,
}
disabledLimiter := NewAPIRateLimiter(nil)
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
disabledLimiter.SetEnabled(false)
authMiddleware := NewAuthMiddleware(
@@ -311,7 +312,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("PAT Token Rate Limiting - Burst Works", func(t *testing.T) {
// Configure rate limiter: 10 requests per minute with burst of 5
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 10,
Burst: 5,
CleanupInterval: 5 * time.Minute,
@@ -329,7 +330,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -364,7 +365,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("PAT Token Rate Limiting - Rate Limit Enforced", func(t *testing.T) {
// Configure very low rate limit: 1 request per minute
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -382,7 +383,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -408,7 +409,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("Bearer Token Not Rate Limited", func(t *testing.T) {
// Configure strict rate limit
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -426,7 +427,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -453,7 +454,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("PAT Token Rate Limiting Per Token", func(t *testing.T) {
// Configure rate limiter
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -471,7 +472,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -518,7 +519,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
t.Run("Rate Limiter Cleanup", func(t *testing.T) {
// Configure rate limiter with short cleanup interval and TTL for testing
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 60,
Burst: 1,
CleanupInterval: 100 * time.Millisecond,
@@ -536,7 +537,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -578,7 +579,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
})
t.Run("Terraform User Agent Not Rate Limited", func(t *testing.T) {
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -596,7 +597,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -634,7 +635,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
})
t.Run("Non-Terraform User Agent With PAT Is Rate Limited", func(t *testing.T) {
rateLimitConfig := &RateLimiterConfig{
rateLimitConfig := &ratelimit.RateLimiterConfig{
RequestsPerMinute: 1,
Burst: 1,
CleanupInterval: 5 * time.Minute,
@@ -652,7 +653,7 @@ func TestAuthMiddleware_RateLimiting(t *testing.T) {
func(ctx context.Context, userAuth nbauth.UserAuth) (*types.User, error) {
return &types.User{}, nil
},
NewAPIRateLimiter(rateLimitConfig),
ratelimit.NewAPIRateLimiter(rateLimitConfig),
nil,
func(_ context.Context, _, _, _ string) bool { return false },
)
@@ -740,7 +741,7 @@ func TestAuthMiddleware_Handler_Child(t *testing.T) {
GetPATInfoFunc: mockGetAccountInfoFromPAT,
}
disabledLimiter := NewAPIRateLimiter(nil)
disabledLimiter := ratelimit.NewAPIRateLimiter(nil)
disabledLimiter.SetEnabled(false)
authMiddleware := NewAuthMiddleware(
mockAuth,
@@ -46,14 +46,14 @@ import (
"github.com/netbirdio/netbird/management/server/networks/routers"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/management/server/store"
nbstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/management/server/users"
"github.com/netbirdio/netbird/shared/auth"
)
func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPeerUpdate *network_map.UpdateMessage, validateUpdate bool) (http.Handler, account.Manager, chan struct{}) {
store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir())
store, cleanup, err := nbstore.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir())
if err != nil {
t.Fatalf("Failed to create test store: %v", err)
}
@@ -108,15 +108,15 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee
t.Fatalf("Failed to create manager: %v", err)
}
accessLogsManager := accesslogsmanager.NewManager(store, permissionsManager, nil)
accessLogsManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(store.(*nbstore.SqlStore).Conn()), store, permissionsManager, nil)
proxyTokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore)
pkceverifierStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore)
noopMeter := noop.NewMeterProvider().Meter("")
proxyMgr, err := proxymanager.NewManager(store, noopMeter)
if err != nil {
t.Fatalf("Failed to create proxy manager: %v", err)
}
proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, pkceverifierStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil)
proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil)
// NewProxyServiceServer starts cleanupStaleProxies on a context it derives
// from context.Background(), independent of the cancellable ctx above;
// Close() cancels it so the goroutine does not outlive the test.
@@ -147,7 +147,7 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter()
apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil)
apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil, nil)
if err != nil {
t.Fatalf("Failed to create API handler: %v", err)
}
@@ -204,7 +204,7 @@ func PeerShouldNotReceiveAnyUpdate(t testing_tools.TB, updateMessage <-chan *net
// BuildApiBlackBoxWithDBStateAndPeerChannel creates the API handler and returns
// the peer update channel directly so tests can verify updates inline.
func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile string) (http.Handler, account.Manager, <-chan *network_map.UpdateMessage) {
store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir())
store, cleanup, err := nbstore.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir())
if err != nil {
t.Fatalf("Failed to create test store: %v", err)
}
@@ -248,15 +248,15 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin
t.Fatalf("Failed to create manager: %v", err)
}
accessLogsManager := accesslogsmanager.NewManager(store, permissionsManager, nil)
accessLogsManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(store.(*nbstore.SqlStore).Conn()), store, permissionsManager, nil)
proxyTokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore)
pkceverifierStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore)
singleUseStore := nbgrpc.NewSingleUseStore(ctx, cacheStore)
noopMeter := noop.NewMeterProvider().Meter("")
proxyMgr, err := proxymanager.NewManager(store, noopMeter)
if err != nil {
t.Fatalf("Failed to create proxy manager: %v", err)
}
proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, pkceverifierStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil)
proxyServiceServer := nbgrpc.NewProxyServiceServer(accessLogsManager, proxyTokenStore, singleUseStore, nbgrpc.ProxyOIDCConfig{}, peersManager, userManager, nil, proxyMgr, nil)
// NewProxyServiceServer starts cleanupStaleProxies on a context it derives
// from context.Background(), independent of the cancellable ctx above;
// Close() cancels it so the goroutine does not outlive the test.
@@ -287,7 +287,7 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin
zoneRecordsManager := recordsManager.NewManager(store, am, permissionsManager)
apiRouter := mux.NewRouter().PathPrefix("/api").Subrouter()
apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil)
apiHandler, err := http2.NewAPIHandler(ctx, apiRouter, am, networksManager, resourcesManager, routersManager, groupsManager, geoMock, authManagerMock, metrics, permissionsManager, settingsManager, customZonesManager, zoneRecordsManager, networkMapController, nil, serviceManager, nil, nil, nil, nil, nil, nil, nil, nil)
if err != nil {
t.Fatalf("Failed to create API handler: %v", err)
}

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