mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-30 02:29:08 +02:00
Merge branch 'main' into loopback-wg-proxy
This commit is contained in:
@@ -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"
|
||||
|
||||
@@ -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') }}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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 })
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
//
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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 "$@"
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
+51
-1
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+42
-5
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,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
Reference in New Issue
Block a user