mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-06 05:29:07 +02:00
Merge remote-tracking branch 'origin/main' into feat-post_quantum_ml_kem
This commit is contained in:
@@ -33,7 +33,7 @@ jobs:
|
||||
with:
|
||||
usesh: true
|
||||
copyback: false
|
||||
release: "15.0"
|
||||
release: "15.1"
|
||||
envs: "GO_VERSION"
|
||||
prepare: |
|
||||
pkg install -y curl pkgconf xorg
|
||||
|
||||
@@ -0,0 +1,199 @@
|
||||
name: Red Hat Certification
|
||||
|
||||
# Certify published UBI images in the Red Hat Ecosystem Catalog. Called by
|
||||
# release.yml on stable tags, or run by hand to (re)certify any released
|
||||
# version. preflight submits every architecture of an image's manifest list
|
||||
# to Pyxis; auto-publish on the component makes it public once certified.
|
||||
#
|
||||
# Each component's Partner Connect ID comes from the REDHAT_CERT_ID_<NAME>
|
||||
# repository variable, e.g. REDHAT_CERT_ID_CLIENT_ROOTLESS. The run fails
|
||||
# before certifying anything if a selected component's variable is not set.
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
component:
|
||||
type: string
|
||||
required: true
|
||||
version:
|
||||
type: string
|
||||
required: true
|
||||
secrets:
|
||||
PYXIS_API_TOKEN:
|
||||
required: true
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
component:
|
||||
description: "Component to certify"
|
||||
type: choice
|
||||
required: true
|
||||
default: all
|
||||
options:
|
||||
- all
|
||||
- client-rootless
|
||||
- reverse-proxy
|
||||
version:
|
||||
description: "Released version, e.g. v0.80.0"
|
||||
type: string
|
||||
required: true
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
resolve:
|
||||
name: Resolve components
|
||||
runs-on: ubuntu-24.04
|
||||
outputs:
|
||||
version: ${{ steps.resolve.outputs.version }}
|
||||
matrix: ${{ steps.resolve.outputs.matrix }}
|
||||
steps:
|
||||
- name: Resolve components and images
|
||||
id: resolve
|
||||
env:
|
||||
COMPONENT: ${{ inputs.component }}
|
||||
INPUT_VERSION: ${{ inputs.version }}
|
||||
REPO_VARS: ${{ toJSON(vars) }}
|
||||
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
|
||||
# name, image repository, tag suffix (must match .goreleaser.yaml).
|
||||
# Keep the names in sync with the workflow_dispatch options above.
|
||||
components=(
|
||||
"client-rootless ghcr.io/netbirdio/netbird -rootless-ubi"
|
||||
"reverse-proxy ghcr.io/netbirdio/reverse-proxy -ubi"
|
||||
)
|
||||
matrix="[]"
|
||||
missing=()
|
||||
for c in "${components[@]}"; do
|
||||
read -r name repo suffix <<< "$c"
|
||||
[[ "$COMPONENT" == all || "$COMPONENT" == "$name" ]] || continue
|
||||
var="REDHAT_CERT_ID_${name^^}"; var="${var//-/_}"
|
||||
id="$(jq -r --arg v "$var" '.[$v] // empty' <<< "$REPO_VARS")"
|
||||
if [[ -z "$id" ]]; then
|
||||
missing+=("$var")
|
||||
continue
|
||||
fi
|
||||
matrix="$(jq -c --arg n "$name" --arg t "${version}${suffix}" --arg r "${repo}:${version}${suffix}" --arg i "$id" \
|
||||
'. + [{component: $n, tag: $t, ref: $r, component_id: $i}]' <<< "$matrix")"
|
||||
done
|
||||
if (( ${#missing[@]} )); then
|
||||
echo "::error::Set these repository variables to the Partner Connect component IDs: ${missing[*]}"
|
||||
exit 1
|
||||
fi
|
||||
if [[ "$matrix" == "[]" ]]; then
|
||||
echo "::error::No component to certify for '${COMPONENT}'"
|
||||
exit 1
|
||||
fi
|
||||
echo "Components to certify: ${matrix}"
|
||||
echo "version=${version}" >> "$GITHUB_OUTPUT"
|
||||
echo "matrix=${matrix}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
certify:
|
||||
name: "Certify ${{ matrix.component }} UBI image"
|
||||
needs: resolve
|
||||
runs-on: ubuntu-24.04
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include: ${{ fromJSON(needs.resolve.outputs.matrix) }}
|
||||
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"
|
||||
steps:
|
||||
- name: Verify the multi-arch image is on ghcr.io
|
||||
env:
|
||||
IMAGE_REF: ${{ matrix.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: ${{ matrix.ref }}
|
||||
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
|
||||
PFLT_CERTIFICATION_COMPONENT_ID: ${{ matrix.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 does not create the PFLT_LOGFILE directory, and --submit
|
||||
# fails if the log file is missing.
|
||||
mkdir -p artifacts
|
||||
./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-${{ matrix.component }}-${{ needs.resolve.outputs.version }}
|
||||
path: artifacts/
|
||||
retention-days: 30
|
||||
|
||||
- name: Wait for Pyxis to mark both architectures certified
|
||||
env:
|
||||
TAG: ${{ matrix.tag }}
|
||||
COMPONENT_ID: ${{ matrix.component_id }}
|
||||
PFLT_PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
# Filter on the tag server-side so older versions are found past the first page.
|
||||
url="https://catalog.redhat.com/api/containers/v1/projects/certification/id/${COMPONENT_ID}/images?filter=repositories.tags.name==${TAG}&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 "::error::Pyxis has not marked both architectures certified after 10 minutes. Check https://connect.redhat.com/component/view/${COMPONENT_ID}/images"
|
||||
exit 1
|
||||
+10
-118
@@ -69,7 +69,7 @@ jobs:
|
||||
with:
|
||||
usesh: true
|
||||
copyback: false
|
||||
release: "15.0"
|
||||
release: "15.1"
|
||||
envs: "GO_VERSION"
|
||||
prepare: |
|
||||
# Install required packages
|
||||
@@ -380,131 +380,23 @@ 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.
|
||||
# Certify the UBI images in the Red Hat Ecosystem Catalog on stable tags.
|
||||
# See redhat-certify.yml, which can also be run by hand for any released version.
|
||||
redhat_certification:
|
||||
name: "Red Hat / Certify rootless UBI image"
|
||||
name: "Red Hat"
|
||||
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"
|
||||
uses: ./.github/workflows/redhat-certify.yml
|
||||
with:
|
||||
component: all
|
||||
version: ${{ github.ref_name }}
|
||||
secrets:
|
||||
PYXIS_API_TOKEN: ${{ secrets.PYXIS_API_TOKEN }}
|
||||
|
||||
release_ui:
|
||||
runs-on: ubuntu-latest
|
||||
|
||||
+52
-33
@@ -3,56 +3,75 @@ package base62
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
||||
base = uint32(len(alphabet))
|
||||
alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
||||
base = uint32(len(alphabet))
|
||||
maxBase62Digits = 6 // max number of digits required to encode MaxUint32
|
||||
|
||||
)
|
||||
|
||||
var (
|
||||
ErrEmptyString = fmt.Errorf("empty string")
|
||||
ErrInvalidChar = fmt.Errorf("invalid character")
|
||||
ErrOverflow = fmt.Errorf("integer overflow")
|
||||
)
|
||||
|
||||
// Fixed-size arrays have better performance and lower memory overhead compared to maps for small static sets of data
|
||||
var charToIndex [123]int8 // Assuming ASCII from '\0' - 'z'
|
||||
|
||||
func init() {
|
||||
for i := range charToIndex {
|
||||
charToIndex[i] = -1
|
||||
}
|
||||
for i, c := range alphabet {
|
||||
charToIndex[c] = int8(i)
|
||||
}
|
||||
}
|
||||
|
||||
// Encode encodes a uint32 value to a base62 string.
|
||||
func Encode(num uint32) string {
|
||||
if num == 0 {
|
||||
return string(alphabet[0])
|
||||
// The returned string will be between 1-6 characters long.
|
||||
func Encode(n uint32) string {
|
||||
if n < base {
|
||||
return string(alphabet[n])
|
||||
}
|
||||
// avoid dynamic memory usage for small, fixed size data
|
||||
buf := [maxBase62Digits]byte{}
|
||||
idx := len(buf)
|
||||
|
||||
for n > 0 {
|
||||
idx--
|
||||
buf[idx] = alphabet[n%base]
|
||||
n /= base
|
||||
}
|
||||
|
||||
var encoded strings.Builder
|
||||
|
||||
for num > 0 {
|
||||
remainder := num % base
|
||||
encoded.WriteByte(alphabet[remainder])
|
||||
num /= base
|
||||
}
|
||||
|
||||
// Reverse the encoded string
|
||||
encodedString := encoded.String()
|
||||
reversed := reverse(encodedString)
|
||||
return reversed
|
||||
return string(buf[idx:])
|
||||
}
|
||||
|
||||
// Decode decodes a base62 string to a uint32 value.
|
||||
// Returns an error if the input string is empty, contains invalid characters,
|
||||
// or would result in integer overflow.
|
||||
func Decode(encoded string) (uint32, error) {
|
||||
if len(encoded) == 0 {
|
||||
return 0, ErrEmptyString
|
||||
}
|
||||
var decoded uint32
|
||||
strLen := len(encoded)
|
||||
|
||||
for i, char := range encoded {
|
||||
index := strings.IndexRune(alphabet, char)
|
||||
for _, char := range encoded {
|
||||
index := int8(-1)
|
||||
if int(char) < len(charToIndex) {
|
||||
index = charToIndex[char]
|
||||
}
|
||||
if index < 0 {
|
||||
return 0, fmt.Errorf("invalid character: %c", char)
|
||||
return 0, fmt.Errorf("%w: %c", ErrInvalidChar, char)
|
||||
}
|
||||
// Add overflow check when calculating the decoded value to prevent silent overflow of uint32
|
||||
if decoded > (math.MaxUint32-uint32(index))/base {
|
||||
return 0, fmt.Errorf("%w: %s", ErrOverflow, encoded)
|
||||
}
|
||||
|
||||
decoded += uint32(index) * uint32(math.Pow(float64(base), float64(strLen-i-1)))
|
||||
decoded = decoded*base + uint32(index)
|
||||
}
|
||||
|
||||
return decoded, nil
|
||||
}
|
||||
|
||||
// Reverse a string.
|
||||
func reverse(s string) string {
|
||||
runes := []rune(s)
|
||||
for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 {
|
||||
runes[i], runes[j] = runes[j], runes[i]
|
||||
}
|
||||
return string(runes)
|
||||
}
|
||||
|
||||
+50
-14
@@ -1,31 +1,67 @@
|
||||
package base62
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEncodeDecode(t *testing.T) {
|
||||
tests := []struct {
|
||||
num uint32
|
||||
testCases := []struct {
|
||||
input uint32
|
||||
expected string
|
||||
}{
|
||||
{0},
|
||||
{1},
|
||||
{42},
|
||||
{12345},
|
||||
{99999},
|
||||
{123456789},
|
||||
{0, "0"},
|
||||
{1, "1"},
|
||||
{5, "5"},
|
||||
{9, "9"},
|
||||
{10, "A"},
|
||||
{42, "g"},
|
||||
{61, "z"},
|
||||
{62, "10"},
|
||||
{'0', "m"},
|
||||
{'9', "v"},
|
||||
{'A', "13"},
|
||||
{'Z', "1S"},
|
||||
{'a', "1Z"},
|
||||
{'z', "1y"},
|
||||
{99999, "Q0t"},
|
||||
{12345, "3D7"},
|
||||
{123456789, "8M0kX"},
|
||||
{math.MaxUint32, "4gfFC3"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
encoded := Encode(tt.num)
|
||||
for _, tc := range testCases {
|
||||
encoded := Encode(tc.input)
|
||||
if encoded != tc.expected {
|
||||
t.Errorf("Encode(%d) = %s; want %s", tc.input, encoded, tc.expected)
|
||||
}
|
||||
decoded, err := Decode(encoded)
|
||||
|
||||
if err != nil {
|
||||
t.Errorf("Decode error: %v", err)
|
||||
t.Errorf("Expected error nil, got %v", err)
|
||||
}
|
||||
|
||||
if decoded != tt.num {
|
||||
t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tt.num)
|
||||
if decoded != tc.input {
|
||||
t.Errorf("Decode(%v) = %v, want %v", encoded, decoded, tc.input)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Decode handles empty string input with appropriate error
|
||||
func TestDecodeEmptyString(t *testing.T) {
|
||||
if _, err := Decode(""); !errors.Is(err, ErrEmptyString) {
|
||||
t.Errorf("Expected error %v, got %v", ErrEmptyString, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeOverflow(t *testing.T) {
|
||||
if _, err := Decode("4gfFC4"); !errors.Is(err, ErrOverflow) {
|
||||
t.Errorf("Expected error %v, got %v", ErrOverflow, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeInvalid(t *testing.T) {
|
||||
if _, err := Decode("/"); !errors.Is(err, ErrInvalidChar) {
|
||||
t.Errorf("Expected error %v, got %v", ErrInvalidChar, err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -213,6 +213,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
|
||||
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
|
||||
internal.WithNetEvents(c.netMgr))
|
||||
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
||||
connectClient.SetSyncResponsePersistence(true)
|
||||
// This path runs the interactive SSO flow, so reaching here means the peer
|
||||
// is authenticated again — release the latch Status() reports from. Clear
|
||||
// only once the fresh connect client is installed: until then Status()
|
||||
@@ -256,6 +257,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
|
||||
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
|
||||
internal.WithNetEvents(c.netMgr))
|
||||
c.setState(cfg, cacheDir, cfgFile, connectClient)
|
||||
connectClient.SetSyncResponsePersistence(true)
|
||||
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
|
||||
}
|
||||
|
||||
@@ -327,6 +329,19 @@ func (c *Client) NotifyNetworkChange() {
|
||||
// or "strict"; strict also anonymizes internal IP ranges, peer names, and
|
||||
// WireGuard public keys, and implies anonymize.
|
||||
func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
|
||||
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, true)
|
||||
}
|
||||
|
||||
// DebugBundleFile generates a debug bundle and returns the path of the zip in
|
||||
// the cache directory instead of uploading it, so the app can hand the file to
|
||||
// the user for inspection. The caller owns the file and removes it once done;
|
||||
// the stale-bundle cleanup of later runs removes it only after a day.
|
||||
// anonymize and anonymizeLevel behave as in DebugBundle.
|
||||
func (c *Client) DebugBundleFile(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string) (string, error) {
|
||||
return c.debugBundle(platformFiles, anonymize, anonymizeLevel, false)
|
||||
}
|
||||
|
||||
func (c *Client) debugBundle(platformFiles PlatformFiles, anonymize bool, anonymizeLevel string, upload bool) (string, error) {
|
||||
cfg, cacheDir, cc := c.stateSnapshot()
|
||||
|
||||
// If the engine hasn't been started, load config from disk
|
||||
@@ -342,6 +357,11 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
|
||||
cacheDir = platformFiles.CacheDir()
|
||||
}
|
||||
|
||||
// Clear what an interrupted earlier run may have left in the cache before
|
||||
// adding to it. Remote debug jobs write to the same directory, so anything
|
||||
// younger than an hour is treated as possibly still in use.
|
||||
debug.RemoveStaleBundles(cacheDir, time.Hour)
|
||||
|
||||
deps := debug.GeneratorDependencies{
|
||||
InternalConfig: cfg,
|
||||
StatusRecorder: c.recorder,
|
||||
@@ -379,6 +399,9 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool, anonym
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("generate debug bundle: %w", err)
|
||||
}
|
||||
if !upload {
|
||||
return debug.ExportBundle(path)
|
||||
}
|
||||
defer func() {
|
||||
if err := os.Remove(path); err != nil {
|
||||
log.Errorf("failed to remove debug bundle file: %v", err)
|
||||
@@ -475,6 +498,7 @@ func (c *Client) Networks() *NetworkArray {
|
||||
routesMap := routeManager.GetClientRoutesWithNetID()
|
||||
v6Merged := route.V6ExitMergeSet(routesMap)
|
||||
resolvedDomains := c.recorder.GetResolvedDomainsStates()
|
||||
activeRoutePeers := c.recorder.GetActiveRoutePeers()
|
||||
|
||||
networkArray := &NetworkArray{
|
||||
items: make([]Network, 0),
|
||||
@@ -488,7 +512,7 @@ func (c *Client) Networks() *NetworkArray {
|
||||
continue
|
||||
}
|
||||
|
||||
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged)
|
||||
network := c.buildNetwork(id, routes, routeSelector.IsSelected(id), resolvedDomains, v6Merged, activeRoutePeers)
|
||||
if network == nil {
|
||||
continue
|
||||
}
|
||||
@@ -497,14 +521,14 @@ func (c *Client) Networks() *NetworkArray {
|
||||
return networkArray
|
||||
}
|
||||
|
||||
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}) *Network {
|
||||
func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bool, resolvedDomains map[domain.Domain]peer.ResolvedDomainInfo, v6Merged map[route.NetID]struct{}, activeRoutePeers map[route.HAUniqueID]string) *Network {
|
||||
r := routes[0]
|
||||
netStr := r.Network.String()
|
||||
if r.IsDynamic() {
|
||||
netStr = r.Domains.SafeString()
|
||||
}
|
||||
|
||||
routePeer, err := c.findBestRoutePeer(routes)
|
||||
routePeer, err := c.findBestRoutePeer(routes, activeRoutePeers)
|
||||
if err != nil {
|
||||
log.Errorf("could not get peer info for route %s: %v", id, err)
|
||||
return nil
|
||||
@@ -528,12 +552,9 @@ func (c *Client) buildNetwork(id route.NetID, routes []*route.Route, selected bo
|
||||
|
||||
// findBestRoutePeer returns the peer actively routing traffic for the given
|
||||
// HA route group. Falls back to the first connected peer, then the first peer.
|
||||
func (c *Client) findBestRoutePeer(routes []*route.Route) (peer.State, error) {
|
||||
netStr := routes[0].Network.String()
|
||||
|
||||
fullStatus := c.recorder.GetFullStatus()
|
||||
for _, p := range fullStatus.Peers {
|
||||
if _, ok := p.GetRoutes()[netStr]; ok {
|
||||
func (c *Client) findBestRoutePeer(routes []*route.Route, activeRoutePeers map[route.HAUniqueID]string) (peer.State, error) {
|
||||
if peerKey, ok := activeRoutePeers[routes[0].GetHAUniqueID()]; ok {
|
||||
if p, err := c.recorder.GetPeer(peerKey); err == nil {
|
||||
return p, nil
|
||||
}
|
||||
}
|
||||
|
||||
+16
-36
@@ -40,14 +40,18 @@ func init() {
|
||||
peerPubKey = peerPrivateKey.PublicKey().String()
|
||||
}
|
||||
|
||||
// testIFaceBlackList mirrors the prefixes profilemanager.DefaultInterfaceBlacklist
|
||||
// carries for the overlay interface. These tests create their own utun device, and
|
||||
// stdnet's filter probes with wgctrl every interface it is not told to skip, which
|
||||
// on a userspace WireGuard platform reaches the UAPI socket of this same process.
|
||||
// Declared here rather than imported because profilemanager imports this package.
|
||||
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
|
||||
|
||||
func TestWGIface_UpdateAddr(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||
addr := "100.64.0.1/8"
|
||||
wgPort := 33100
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -127,10 +131,7 @@ func getIfaceAddrs(ifaceName string) ([]net.Addr, error) {
|
||||
func Test_CreateInterface(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+1)
|
||||
wgIP := "10.99.99.1/32"
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
Address: wgaddr.MustParseWGAddress(wgIP),
|
||||
@@ -170,10 +171,7 @@ func Test_Close(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
|
||||
wgIP := "10.99.99.2/32"
|
||||
wgPort := 33100
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -215,10 +213,7 @@ func TestRecreation(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+2)
|
||||
wgIP := "10.99.99.2/32"
|
||||
wgPort := 33100
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -288,10 +283,7 @@ func Test_ConfigureInterface(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+3)
|
||||
wgIP := "10.99.99.5/30"
|
||||
wgPort := 33100
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
Address: wgaddr.MustParseWGAddress(wgIP),
|
||||
@@ -343,10 +335,7 @@ func Test_ConfigureInterface(t *testing.T) {
|
||||
func Test_UpdatePeer(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||
wgIP := "10.99.99.9/30"
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -413,10 +402,7 @@ func Test_UpdatePeer(t *testing.T) {
|
||||
func Test_RemovePeer(t *testing.T) {
|
||||
ifaceName := fmt.Sprintf("utun%d", WgIntNumber+4)
|
||||
wgIP := "10.99.99.13/30"
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := WGIFaceOpts{
|
||||
IFaceName: ifaceName,
|
||||
@@ -477,10 +463,7 @@ func Test_ConnectPeers(t *testing.T) {
|
||||
peer2wgPort := 33200
|
||||
|
||||
keepAlive := 1 * time.Second
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
guid := fmt.Sprintf("{%s}", uuid.New().String())
|
||||
device.CustomWindowsGUIDString = strings.ToLower(guid)
|
||||
@@ -516,10 +499,7 @@ func Test_ConnectPeers(t *testing.T) {
|
||||
guid = fmt.Sprintf("{%s}", uuid.New().String())
|
||||
device.CustomWindowsGUIDString = strings.ToLower(guid)
|
||||
|
||||
newNet, err = stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet = stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
optsPeer2 := WGIFaceOpts{
|
||||
IFaceName: peer2ifaceName,
|
||||
|
||||
@@ -200,10 +200,7 @@ func (m *SingleSocketUDPMux) updateLocalAddresses() {
|
||||
}
|
||||
if len(networks) > 0 {
|
||||
if m.params.Net == nil {
|
||||
var err error
|
||||
if m.params.Net, err = stdnet.NewNet(context.Background(), nil); err != nil {
|
||||
m.params.Logger.Errorf("failed to get create network: %v", err)
|
||||
}
|
||||
m.params.Net = stdnet.NewNet(context.Background(), nil)
|
||||
}
|
||||
|
||||
ips, err := localInterfaces(m.params.Net, m.params.InterfaceFilter, nil, networks, true)
|
||||
|
||||
@@ -380,9 +380,38 @@ func NewBundleGenerator(deps GeneratorDependencies, cfg BundleConfig) *BundleGen
|
||||
}
|
||||
}
|
||||
|
||||
// bundleFilePattern names the bundle zips Generate creates in tempDir; the
|
||||
// asterisk is filled in by os.CreateTemp.
|
||||
const bundleFilePattern = "netbird.debug.*.zip"
|
||||
|
||||
const exportedBundlePrefix = "netbird.debug-file."
|
||||
|
||||
const exportedBundleMaxAge = 24 * time.Hour
|
||||
|
||||
// RemoveStaleBundles deletes bundle zips that an interrupted generation or
|
||||
// upload left behind in dir. Only files older than maxAge go, so a bundle that
|
||||
// another caller is still writing or uploading in the same directory survives.
|
||||
// Exported bundles are kept for exportedBundleMaxAge instead.
|
||||
func RemoveStaleBundles(dir string, maxAge time.Duration) {
|
||||
removeStaleFiles(dir, bundleFilePattern, maxAge)
|
||||
removeStaleFiles(dir, exportedBundlePrefix+"*.zip", exportedBundleMaxAge)
|
||||
}
|
||||
|
||||
// ExportBundle renames a generated bundle out of the RemoveStaleBundles pattern
|
||||
// and returns the new path. The caller owns the file from then on; an export
|
||||
// abandoned for longer than exportedBundleMaxAge is removed by RemoveStaleBundles.
|
||||
func ExportBundle(path string) (string, error) {
|
||||
base := strings.TrimPrefix(filepath.Base(path), strings.SplitN(bundleFilePattern, "*", 2)[0])
|
||||
exported := filepath.Join(filepath.Dir(path), exportedBundlePrefix+base)
|
||||
if err := os.Rename(path, exported); err != nil {
|
||||
return "", fmt.Errorf("export debug bundle: %w", err)
|
||||
}
|
||||
return exported, nil
|
||||
}
|
||||
|
||||
// Generate creates a debug bundle and returns the location.
|
||||
func (g *BundleGenerator) Generate() (resp string, err error) {
|
||||
bundlePath, err := os.CreateTemp(g.tempDir, "netbird.debug.*.zip")
|
||||
bundlePath, err := os.CreateTemp(g.tempDir, bundleFilePattern)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("create zip file: %w", err)
|
||||
}
|
||||
@@ -1729,3 +1758,26 @@ func anonymizeSlice(v []any, anonymizer *anonymize.Anonymizer) []any {
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func removeStaleFiles(dir, pattern string, maxAge time.Duration) {
|
||||
matches, err := filepath.Glob(filepath.Join(dir, pattern))
|
||||
if err != nil {
|
||||
log.Debugf("glob stale debug bundles in %s: %v", dir, err)
|
||||
return
|
||||
}
|
||||
|
||||
cutoff := time.Now().Add(-maxAge)
|
||||
for _, path := range matches {
|
||||
info, err := os.Stat(path)
|
||||
if err != nil || info.ModTime().After(cutoff) {
|
||||
continue
|
||||
}
|
||||
if err := os.Remove(path); err != nil {
|
||||
if !errors.Is(err, fs.ErrNotExist) {
|
||||
log.Warnf("remove stale debug bundle %s: %v", path, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
log.Infof("removed stale debug bundle %s", path)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
@@ -969,3 +970,52 @@ func renderAddConfigSpecific(g *BundleGenerator) string {
|
||||
func newAnonymizerForTest() *anonymize.Anonymizer {
|
||||
return anonymize.NewAnonymizer(anonymize.DefaultAddresses())
|
||||
}
|
||||
|
||||
func TestRemoveStaleBundles(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
stale := filepath.Join(dir, "netbird.debug.111.zip")
|
||||
fresh := filepath.Join(dir, "netbird.debug.222.zip")
|
||||
other := filepath.Join(dir, "netbird.debug.333.txt")
|
||||
owned := filepath.Join(dir, "netbird.debug.444.zip")
|
||||
abandoned := filepath.Join(dir, "netbird.debug.555.zip")
|
||||
for _, p := range []string{stale, fresh, other, owned, abandoned} {
|
||||
require.NoError(t, os.WriteFile(p, []byte("x"), 0o600))
|
||||
}
|
||||
exported, err := ExportBundle(owned)
|
||||
require.NoError(t, err)
|
||||
exportedAbandoned, err := ExportBundle(abandoned)
|
||||
require.NoError(t, err)
|
||||
old := time.Now().Add(-2 * time.Hour)
|
||||
for _, p := range []string{stale, other, exported} {
|
||||
require.NoError(t, os.Chtimes(p, old, old))
|
||||
}
|
||||
ancient := time.Now().Add(-exportedBundleMaxAge - time.Hour)
|
||||
require.NoError(t, os.Chtimes(exportedAbandoned, ancient, ancient))
|
||||
|
||||
RemoveStaleBundles(dir, time.Hour)
|
||||
|
||||
assert.NoFileExists(t, stale, "bundle older than maxAge should be removed")
|
||||
assert.FileExists(t, fresh, "bundle younger than maxAge must survive, it may still be uploading")
|
||||
assert.FileExists(t, other, "files outside the bundle pattern must not be touched")
|
||||
assert.NoFileExists(t, owned)
|
||||
assert.FileExists(t, exported, "exported bundle is caller-owned and must survive maxAge")
|
||||
assert.NoFileExists(t, exportedAbandoned, "exported bundle older than exportedBundleMaxAge is abandoned")
|
||||
}
|
||||
|
||||
func TestBundleIncludesNetworkMap(t *testing.T) {
|
||||
for _, anonymize := range []bool{false, true} {
|
||||
t.Run(fmt.Sprintf("anonymize=%t", anonymize), func(t *testing.T) {
|
||||
g := NewBundleGenerator(GeneratorDependencies{
|
||||
SyncResponse: &mgmProto.SyncResponse{NetworkMap: &mgmProto.NetworkMap{Serial: 1}},
|
||||
}, BundleConfig{Anonymize: anonymize})
|
||||
|
||||
require.Contains(t, bundleEntries(t, g), "network_map.json")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBundleOmitsNetworkMapWithoutSyncResponse(t *testing.T) {
|
||||
g := NewBundleGenerator(GeneratorDependencies{}, BundleConfig{})
|
||||
|
||||
require.NotContains(t, bundleEntries(t, g), "network_map.json")
|
||||
}
|
||||
|
||||
@@ -9,9 +9,9 @@ import (
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/miekg/dns"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface"
|
||||
@@ -24,6 +24,10 @@ import (
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
)
|
||||
|
||||
// testIFaceBlackList mirrors the overlay prefixes profilemanager.DefaultInterfaceBlacklist
|
||||
// carries. Declared here rather than imported because profilemanager imports this package.
|
||||
var testIFaceBlackList = []string{"wt", "utun", "tun0"}
|
||||
|
||||
func TestUpdateDNSServer(t *testing.T) {
|
||||
|
||||
nameServers := []nbdns.NameServer{
|
||||
@@ -243,10 +247,7 @@ func TestUpdateDNSServer(t *testing.T) {
|
||||
for n, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
privKey, _ := wgtypes.GenerateKey()
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), testIFaceBlackList)
|
||||
|
||||
opts := iface.WGIFaceOpts{
|
||||
IFaceName: fmt.Sprintf("utun230%d", n),
|
||||
@@ -348,11 +349,7 @@ func TestDNSFakeResolverHandleUpdates(t *testing.T) {
|
||||
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
|
||||
|
||||
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
|
||||
newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"})
|
||||
if err != nil {
|
||||
t.Errorf("create stdnet: %v", err)
|
||||
return
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), []string{"utun2301"})
|
||||
|
||||
privKey, _ := wgtypes.GeneratePrivateKey()
|
||||
opts := iface.WGIFaceOpts{
|
||||
|
||||
@@ -394,11 +394,7 @@ func createWgInterfaceWithBind(t *testing.T) (*iface.WGIface, error) {
|
||||
defer t.Setenv("NB_WG_KERNEL_DISABLED", ov)
|
||||
|
||||
t.Setenv("NB_WG_KERNEL_DISABLED", "true")
|
||||
newNet, err := stdnet.NewNet(context.Background(), []string{"utun2301"})
|
||||
if err != nil {
|
||||
t.Fatalf("create stdnet: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), []string{"utun2301"})
|
||||
|
||||
privKey, _ := wgtypes.GeneratePrivateKey()
|
||||
|
||||
|
||||
@@ -2249,10 +2249,7 @@ func (e *Engine) close() {
|
||||
}
|
||||
|
||||
func (e *Engine) newWgIface() (*iface.WGIface, error) {
|
||||
transportNet, err := e.newStdNet()
|
||||
if err != nil {
|
||||
log.Errorf("failed to create pion's stdnet: %s", err)
|
||||
}
|
||||
transportNet := e.newStdNet()
|
||||
|
||||
opts := iface.WGIFaceOpts{
|
||||
IFaceName: e.config.WgIfaceName,
|
||||
|
||||
@@ -12,12 +12,12 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/google/uuid"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/keepalive"
|
||||
@@ -27,6 +27,7 @@ import (
|
||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||
"github.com/netbirdio/netbird/client/internal/dns"
|
||||
"github.com/netbirdio/netbird/client/internal/peer"
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
nbssh "github.com/netbirdio/netbird/client/ssh"
|
||||
"github.com/netbirdio/netbird/client/system"
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
@@ -81,6 +82,7 @@ func TestEngine_SSH(t *testing.T) {
|
||||
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
|
||||
WgPrivateKey: key,
|
||||
WgPort: 33100,
|
||||
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
|
||||
ServerSSHAllowed: true,
|
||||
MTU: iface.DefaultMTU,
|
||||
SSHKey: sshKey,
|
||||
@@ -204,11 +206,12 @@ func TestEngine_Sync(t *testing.T) {
|
||||
}
|
||||
relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU)
|
||||
engine := NewEngine(ctx, cancel, &EngineConfig{
|
||||
WgIfaceName: "utun103",
|
||||
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
|
||||
WgPrivateKey: key,
|
||||
WgPort: 33100,
|
||||
MTU: iface.DefaultMTU,
|
||||
WgIfaceName: "utun103",
|
||||
WgAddr: wgaddr.MustParseWGAddress("100.64.0.1/24"),
|
||||
WgPrivateKey: key,
|
||||
WgPort: 33100,
|
||||
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
|
||||
MTU: iface.DefaultMTU,
|
||||
}, EngineServices{
|
||||
SignalClient: &signal.MockClient{},
|
||||
MgmClient: &mgmt.MockClient{SyncFunc: syncFunc},
|
||||
@@ -412,11 +415,12 @@ func createEngine(ctx context.Context, cancel context.CancelFunc, setupKey strin
|
||||
|
||||
wgPort := 33100 + i
|
||||
conf := &EngineConfig{
|
||||
WgIfaceName: ifaceName,
|
||||
WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address),
|
||||
WgPrivateKey: key,
|
||||
WgPort: wgPort,
|
||||
MTU: iface.DefaultMTU,
|
||||
WgIfaceName: ifaceName,
|
||||
WgAddr: wgaddr.MustParseWGAddress(resp.PeerConfig.Address),
|
||||
WgPrivateKey: key,
|
||||
WgPort: wgPort,
|
||||
IFaceBlackList: profilemanager.DefaultInterfaceBlacklist,
|
||||
MTU: iface.DefaultMTU,
|
||||
}
|
||||
|
||||
relayMgr := relayClient.NewManager(ctx, nil, key.PublicKey().String(), iface.DefaultMTU)
|
||||
|
||||
@@ -6,6 +6,6 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
)
|
||||
|
||||
func (e *Engine) newStdNet() (*stdnet.Net, error) {
|
||||
func (e *Engine) newStdNet() *stdnet.Net {
|
||||
return stdnet.NewNet(e.clientCtx, e.config.IFaceBlackList)
|
||||
}
|
||||
|
||||
@@ -2,6 +2,6 @@ package internal
|
||||
|
||||
import "github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
|
||||
func (e *Engine) newStdNet() (*stdnet.Net, error) {
|
||||
func (e *Engine) newStdNet() *stdnet.Net {
|
||||
return stdnet.NewNetWithDiscover(e.clientCtx, e.mobileDep.IFaceDiscover, e.config.IFaceBlackList)
|
||||
}
|
||||
|
||||
@@ -161,7 +161,6 @@ func (m *MockWGIface) GetProxy() wgproxy.Proxy {
|
||||
return m.GetProxyFunc()
|
||||
}
|
||||
|
||||
|
||||
func (m *MockWGIface) GetNet() *netstack.Net {
|
||||
return m.GetNetFunc()
|
||||
}
|
||||
@@ -689,10 +688,7 @@ func TestEngine_UpdateNetworkMapWithRoutes(t *testing.T) {
|
||||
StatusRecorder: peer.NewRecorder("https://mgm"),
|
||||
}, MobileDependency{})
|
||||
engine.ctx = ctx
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
|
||||
|
||||
opts := iface.WGIFaceOpts{
|
||||
IFaceName: wgIfaceName,
|
||||
@@ -897,10 +893,7 @@ func TestEngine_UpdateNetworkMapWithDNSUpdate(t *testing.T) {
|
||||
}, MobileDependency{})
|
||||
engine.ctx = ctx
|
||||
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
|
||||
opts := iface.WGIFaceOpts{
|
||||
IFaceName: wgIfaceName,
|
||||
Address: wgaddr.MustParseWGAddress(wgAddr),
|
||||
@@ -1500,3 +1493,21 @@ func TestOverlayAddrsFromAllowedIPs(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEngine_SyncResponsePersistence(t *testing.T) {
|
||||
e := &Engine{}
|
||||
|
||||
_, err := e.GetLatestSyncResponse()
|
||||
require.Error(t, err, "persistence is disabled by default")
|
||||
|
||||
e.SetSyncResponsePersistence(true)
|
||||
e.persistSyncResponse(&mgmtProto.SyncResponse{NetworkMap: &mgmtProto.NetworkMap{Serial: 7}})
|
||||
|
||||
got, err := e.GetLatestSyncResponse()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(7), got.GetNetworkMap().GetSerial())
|
||||
|
||||
e.SetSyncResponsePersistence(false)
|
||||
_, err = e.GetLatestSyncResponse()
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
@@ -39,10 +39,7 @@ func NewAgent(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, c
|
||||
iceFailedTimeout := iceFailedTimeout()
|
||||
iceRelayAcceptanceMinWait := iceRelayAcceptanceMinWait()
|
||||
|
||||
transportNet, err := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList)
|
||||
if err != nil {
|
||||
log.Errorf("failed to create pion's stdnet: %s", err)
|
||||
}
|
||||
transportNet := newStdNet(ctx, iFaceDiscover, config.InterfaceBlackList)
|
||||
|
||||
fac := logging.NewDefaultLoggerFactory()
|
||||
|
||||
|
||||
@@ -8,6 +8,6 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
)
|
||||
|
||||
func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) {
|
||||
func newStdNet(ctx context.Context, _ stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net {
|
||||
return stdnet.NewNet(ctx, ifaceBlacklist)
|
||||
}
|
||||
|
||||
@@ -6,6 +6,6 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
)
|
||||
|
||||
func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) (*stdnet.Net, error) {
|
||||
func newStdNet(ctx context.Context, iFaceDiscover stdnet.ExternalIFaceDiscover, ifaceBlacklist []string) *stdnet.Net {
|
||||
return stdnet.NewNetWithDiscover(ctx, iFaceDiscover, ifaceBlacklist)
|
||||
}
|
||||
|
||||
@@ -205,6 +205,7 @@ type Status struct {
|
||||
muxRelays sync.RWMutex
|
||||
peers map[string]State
|
||||
ipToKey map[string]string
|
||||
activeRoutePeers map[route.HAUniqueID]string
|
||||
changeNotify map[string]map[string]*StatusChangeSubscription // map[peerID]map[subscriptionID]*StatusChangeSubscription
|
||||
signalState bool
|
||||
signalError error
|
||||
@@ -268,6 +269,7 @@ func NewRecorder(mgmAddress string) *Status {
|
||||
return &Status{
|
||||
peers: make(map[string]State),
|
||||
ipToKey: make(map[string]string),
|
||||
activeRoutePeers: make(map[route.HAUniqueID]string),
|
||||
changeNotify: make(map[string]map[string]*StatusChangeSubscription),
|
||||
eventStreams: make(map[string]chan *proto.SystemEvent),
|
||||
eventQueue: NewEventQueue(eventQueueSize),
|
||||
@@ -492,6 +494,24 @@ func (d *Status) RemovePeerStateRoute(peer string, route string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *Status) AddActiveRoutePeer(haID route.HAUniqueID, peer string) {
|
||||
d.mux.Lock()
|
||||
defer d.mux.Unlock()
|
||||
d.activeRoutePeers[haID] = peer
|
||||
}
|
||||
|
||||
func (d *Status) RemoveActiveRoutePeer(haID route.HAUniqueID) {
|
||||
d.mux.Lock()
|
||||
defer d.mux.Unlock()
|
||||
delete(d.activeRoutePeers, haID)
|
||||
}
|
||||
|
||||
func (d *Status) GetActiveRoutePeers() map[route.HAUniqueID]string {
|
||||
d.mux.RLock()
|
||||
defer d.mux.RUnlock()
|
||||
return maps.Clone(d.activeRoutePeers)
|
||||
}
|
||||
|
||||
// CheckRoutes checks if the source and destination addresses are within the same route
|
||||
// and returns the resource ID of the route that contains the addresses
|
||||
func (d *Status) CheckRoutes(ip netip.Addr) ([]byte, bool) {
|
||||
|
||||
@@ -9,6 +9,8 @@ import (
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/route"
|
||||
)
|
||||
|
||||
func TestAddPeer(t *testing.T) {
|
||||
@@ -372,3 +374,24 @@ func TestMarkServerStateDoesNotNotifyWhenUnchanged(t *testing.T) {
|
||||
status.MarkManagementDisconnected(err)
|
||||
assert.False(t, notified(ch), "redundant disconnect should not notify")
|
||||
}
|
||||
|
||||
func TestActiveRoutePeers(t *testing.T) {
|
||||
status := NewRecorder("https://mgm")
|
||||
netA := route.HAUniqueID("net-a-10.0.0.0/24")
|
||||
netB := route.HAUniqueID("net-b-10.0.0.0/24")
|
||||
|
||||
status.AddActiveRoutePeer(netA, "peerA")
|
||||
status.AddActiveRoutePeer(netB, "peerB")
|
||||
|
||||
active := status.GetActiveRoutePeers()
|
||||
assert.Equal(t, "peerA", active[netA])
|
||||
assert.Equal(t, "peerB", active[netB])
|
||||
|
||||
status.RemoveActiveRoutePeer(netA)
|
||||
delete(active, netB)
|
||||
|
||||
active = status.GetActiveRoutePeers()
|
||||
_, ok := active[netA]
|
||||
assert.False(t, ok)
|
||||
assert.Equal(t, "peerB", active[netB])
|
||||
}
|
||||
|
||||
@@ -201,11 +201,7 @@ func (p *StunTurnProbe) probeSTUN(ctx context.Context, uri *stun.URI) (addr stri
|
||||
}
|
||||
}()
|
||||
|
||||
net, err := stdnet.NewNet(ctx, nil)
|
||||
if err != nil {
|
||||
probeErr = fmt.Errorf("new net: %w", err)
|
||||
return
|
||||
}
|
||||
net := stdnet.NewNet(ctx, nil)
|
||||
|
||||
client, err := stun.DialURI(uri, &stun.DialConfig{
|
||||
Net: net,
|
||||
@@ -290,11 +286,7 @@ func (p *StunTurnProbe) probeTURN(ctx context.Context, uri *stun.URI) (addr stri
|
||||
}
|
||||
}()
|
||||
|
||||
net, err := stdnet.NewNet(ctx, nil)
|
||||
if err != nil {
|
||||
probeErr = fmt.Errorf("new net: %w", err)
|
||||
return
|
||||
}
|
||||
net := stdnet.NewNet(ctx, nil)
|
||||
cfg := &turn.ClientConfig{
|
||||
STUNServerAddr: turnServerAddr,
|
||||
TURNServerAddr: turnServerAddr,
|
||||
|
||||
@@ -294,6 +294,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error {
|
||||
return fmt.Errorf("add allowed IPs for peer %s: %w", route.Peer, err)
|
||||
}
|
||||
|
||||
w.statusRecorder.AddActiveRoutePeer(route.GetHAUniqueID(), route.Peer)
|
||||
if err := w.statusRecorder.AddPeerStateRoute(route.Peer, w.handler.String(), route.GetResourceID()); err != nil {
|
||||
log.Warnf("Failed to update peer state: %v", err)
|
||||
}
|
||||
@@ -303,6 +304,7 @@ func (w *Watcher) addAllowedIPs(route *route.Route) error {
|
||||
}
|
||||
|
||||
func (w *Watcher) removeAllowedIPs(route *route.Route, rsn reason) error {
|
||||
w.statusRecorder.RemoveActiveRoutePeer(route.GetHAUniqueID())
|
||||
if err := w.statusRecorder.RemovePeerStateRoute(route.Peer, w.handler.String()); err != nil {
|
||||
log.Warnf("Failed to update peer state: %v", err)
|
||||
}
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
|
||||
@@ -406,10 +407,7 @@ func TestManagerUpdateRoutes(t *testing.T) {
|
||||
for n, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
peerPrivateKey, _ := wgtypes.GeneratePrivateKey()
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
|
||||
opts := iface.WGIFaceOpts{
|
||||
IFaceName: fmt.Sprintf("utun43%d", n),
|
||||
Address: wgaddr.MustParseWGAddress("100.65.65.2/24"),
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"syscall"
|
||||
"testing"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -436,8 +437,7 @@ func createWGInterface(t *testing.T, interfaceName, ipAddressCIDR string, listen
|
||||
peerPrivateKey, err := wgtypes.GeneratePrivateKey()
|
||||
require.NoError(t, err)
|
||||
|
||||
newNet, err := stdnet.NewNet(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
newNet := stdnet.NewNet(context.Background(), profilemanager.DefaultInterfaceBlacklist)
|
||||
|
||||
opts := iface.WGIFaceOpts{
|
||||
IFaceName: interfaceName,
|
||||
|
||||
@@ -45,7 +45,7 @@ type Net struct {
|
||||
}
|
||||
|
||||
// NewNetWithDiscover creates a new StdNet instance.
|
||||
func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) (*Net, error) {
|
||||
func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover, disallowList []string) *Net {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
@@ -60,20 +60,19 @@ func NewNetWithDiscover(ctx context.Context, iFaceDiscover ExternalIFaceDiscover
|
||||
} else {
|
||||
n.iFaceDiscover = newMobileIFaceDiscover(iFaceDiscover)
|
||||
}
|
||||
return n, n.UpdateInterfaces()
|
||||
return n
|
||||
}
|
||||
|
||||
// NewNet creates a new StdNet instance.
|
||||
func NewNet(ctx context.Context, disallowList []string) (*Net, error) {
|
||||
func NewNet(ctx context.Context, disallowList []string) *Net {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
n := &Net{
|
||||
return &Net{
|
||||
iFaceDiscover: pionDiscover{},
|
||||
interfaceFilter: InterfaceFilter(disallowList),
|
||||
ctx: ctx,
|
||||
}
|
||||
return n, n.UpdateInterfaces()
|
||||
}
|
||||
|
||||
// resolveAddr performs DNS resolution with context support and timeout.
|
||||
@@ -122,45 +121,18 @@ func (n *Net) resolveAddr(network, address string) (netip.AddrPort, error) {
|
||||
return netip.AddrPortFrom(addrs[0], uint16(port)), nil
|
||||
}
|
||||
|
||||
// UpdateInterfaces updates the internal list of network interfaces
|
||||
// and associated addresses filtering them by name.
|
||||
// The interfaces are discovered by an external iFaceDiscover function or by a default discoverer if the external one
|
||||
// wasn't specified.
|
||||
func (n *Net) UpdateInterfaces() (err error) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
|
||||
return n.updateInterfaces()
|
||||
}
|
||||
|
||||
func (n *Net) updateInterfaces() (err error) {
|
||||
allIfaces, err := n.iFaceDiscover.iFaces()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
n.interfaces = n.filterInterfaces(allIfaces)
|
||||
|
||||
n.lastUpdate = time.Now()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Interfaces returns a slice of interfaces which are available on the
|
||||
// system
|
||||
func (n *Net) Interfaces() ([]*transport.Interface, error) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
|
||||
if time.Since(n.lastUpdate) < updateInterval {
|
||||
return slices.Clone(n.interfaces), nil
|
||||
iFaces, err := n.freshInterfacesLocked()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := n.updateInterfaces(); err != nil {
|
||||
return nil, fmt.Errorf("update interfaces: %w", err)
|
||||
}
|
||||
|
||||
return slices.Clone(n.interfaces), nil
|
||||
return slices.Clone(iFaces), nil
|
||||
}
|
||||
|
||||
// InterfaceByIndex returns the interface specified by index.
|
||||
@@ -171,7 +143,13 @@ func (n *Net) Interfaces() ([]*transport.Interface, error) {
|
||||
func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
for _, ifc := range n.interfaces {
|
||||
|
||||
iFaces, err := n.freshInterfacesLocked()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, ifc := range iFaces {
|
||||
if ifc.Index == index {
|
||||
return ifc, nil
|
||||
}
|
||||
@@ -184,7 +162,13 @@ func (n *Net) InterfaceByIndex(index int) (*transport.Interface, error) {
|
||||
func (n *Net) InterfaceByName(name string) (*transport.Interface, error) {
|
||||
n.mu.Lock()
|
||||
defer n.mu.Unlock()
|
||||
for _, ifc := range n.interfaces {
|
||||
|
||||
iFaces, err := n.freshInterfacesLocked()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, ifc := range iFaces {
|
||||
if ifc.Name == name {
|
||||
return ifc, nil
|
||||
}
|
||||
@@ -193,6 +177,31 @@ func (n *Net) InterfaceByName(name string) (*transport.Interface, error) {
|
||||
return nil, fmt.Errorf("%w: %s", transport.ErrInterfaceNotFound, name)
|
||||
}
|
||||
|
||||
func (n *Net) freshInterfacesLocked() ([]*transport.Interface, error) {
|
||||
if time.Since(n.lastUpdate) < updateInterval {
|
||||
return n.interfaces, nil
|
||||
}
|
||||
|
||||
if err := n.updateInterfacesLocked(); err != nil {
|
||||
return nil, fmt.Errorf("update interfaces: %w", err)
|
||||
}
|
||||
|
||||
return n.interfaces, nil
|
||||
}
|
||||
|
||||
func (n *Net) updateInterfacesLocked() error {
|
||||
allIFaces, err := n.iFaceDiscover.iFaces()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
n.interfaces = n.filterInterfaces(allIFaces)
|
||||
|
||||
n.lastUpdate = time.Now()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (n *Net) filterInterfaces(interfaces []*transport.Interface) []*transport.Interface {
|
||||
if n.interfaceFilter == nil {
|
||||
return interfaces
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
package stdnet
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/pion/transport/v3"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type countingDiscover struct {
|
||||
calls int
|
||||
list []*transport.Interface
|
||||
err error
|
||||
}
|
||||
|
||||
func (d *countingDiscover) iFaces() ([]*transport.Interface, error) {
|
||||
d.calls++
|
||||
if d.err != nil {
|
||||
return nil, d.err
|
||||
}
|
||||
return d.list, nil
|
||||
}
|
||||
|
||||
func newTestNet(t *testing.T, d iFaceDiscover) *Net {
|
||||
t.Helper()
|
||||
return &Net{
|
||||
iFaceDiscover: d,
|
||||
ctx: context.Background(),
|
||||
}
|
||||
}
|
||||
|
||||
func testIFace(index int, name string) *transport.Interface {
|
||||
return transport.NewInterface(net.Interface{Index: index, Name: name})
|
||||
}
|
||||
|
||||
func TestNet_InterfacesDiscoversLazilyAndCaches(t *testing.T) {
|
||||
d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}}
|
||||
n := newTestNet(t, d)
|
||||
|
||||
require.Zero(t, d.calls, "construction must not discover interfaces")
|
||||
|
||||
iFaces, err := n.Interfaces()
|
||||
require.NoError(t, err)
|
||||
require.Len(t, iFaces, 1)
|
||||
assert.Equal(t, 1, d.calls)
|
||||
|
||||
_, err = n.Interfaces()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 1, d.calls)
|
||||
}
|
||||
|
||||
func TestNewNet_DoesNotDiscoverAtConstruction(t *testing.T) {
|
||||
n := NewNet(context.Background(), nil)
|
||||
require.NotNil(t, n)
|
||||
assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold")
|
||||
}
|
||||
|
||||
func TestNewNetWithDiscover_DoesNotDiscoverAtConstruction(t *testing.T) {
|
||||
n := NewNetWithDiscover(context.Background(), nil, nil)
|
||||
require.NotNil(t, n)
|
||||
assert.True(t, n.lastUpdate.IsZero(), "constructor must leave the cache cold")
|
||||
}
|
||||
|
||||
func TestNet_InterfacesRetryAfterDiscoveryFailure(t *testing.T) {
|
||||
discoverErr := errors.New("discover failed")
|
||||
d := &countingDiscover{err: discoverErr}
|
||||
n := newTestNet(t, d)
|
||||
|
||||
_, err := n.Interfaces()
|
||||
require.ErrorIs(t, err, discoverErr)
|
||||
|
||||
d.err = nil
|
||||
d.list = []*transport.Interface{testIFace(1, "eth0")}
|
||||
|
||||
iFaces, err := n.Interfaces()
|
||||
require.NoError(t, err)
|
||||
require.Len(t, iFaces, 1)
|
||||
assert.Equal(t, 2, d.calls)
|
||||
}
|
||||
|
||||
func TestNet_InterfaceByNameRefreshes(t *testing.T) {
|
||||
d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}}
|
||||
n := newTestNet(t, d)
|
||||
|
||||
ifc, err := n.InterfaceByName("eth0")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "eth0", ifc.Name)
|
||||
assert.Equal(t, 1, d.calls)
|
||||
|
||||
_, err = n.InterfaceByName("nope")
|
||||
require.ErrorIs(t, err, transport.ErrInterfaceNotFound)
|
||||
}
|
||||
|
||||
func TestNet_InterfaceByIndexRefreshes(t *testing.T) {
|
||||
d := &countingDiscover{list: []*transport.Interface{testIFace(3, "eth0")}}
|
||||
n := newTestNet(t, d)
|
||||
|
||||
ifc, err := n.InterfaceByIndex(3)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "eth0", ifc.Name)
|
||||
assert.Equal(t, 1, d.calls)
|
||||
|
||||
_, err = n.InterfaceByIndex(99)
|
||||
require.ErrorIs(t, err, transport.ErrInterfaceNotFound)
|
||||
}
|
||||
|
||||
func TestNet_InterfaceLookupPropagatesDiscoveryError(t *testing.T) {
|
||||
discoverErr := errors.New("discover failed")
|
||||
n := newTestNet(t, &countingDiscover{err: discoverErr})
|
||||
|
||||
_, err := n.InterfaceByName("eth0")
|
||||
require.ErrorIs(t, err, discoverErr)
|
||||
|
||||
_, err = n.InterfaceByIndex(1)
|
||||
require.ErrorIs(t, err, discoverErr)
|
||||
}
|
||||
|
||||
func TestNet_InterfacesReturnsCopy(t *testing.T) {
|
||||
d := &countingDiscover{list: []*transport.Interface{testIFace(1, "eth0")}}
|
||||
n := newTestNet(t, d)
|
||||
|
||||
iFaces, err := n.Interfaces()
|
||||
require.NoError(t, err)
|
||||
require.Len(t, iFaces, 1)
|
||||
|
||||
iFaces[0] = testIFace(2, "tampered")
|
||||
|
||||
iFaces, err = n.Interfaces()
|
||||
require.NoError(t, err)
|
||||
require.Len(t, iFaces, 1)
|
||||
assert.Equal(t, "eth0", iFaces[0].Name)
|
||||
}
|
||||
@@ -1,4 +1,5 @@
|
||||
import { useEffect, useRef } from "react";
|
||||
import { useSearchParams } from "react-router-dom";
|
||||
import { Events } from "@wailsio/runtime";
|
||||
import { useStatus } from "@/contexts/StatusContext.tsx";
|
||||
|
||||
@@ -6,13 +7,15 @@ const EVENT_WINDOW_PAINTED = "netbird:window-painted";
|
||||
|
||||
export const ReadySignal = () => {
|
||||
const { isReady } = useStatus();
|
||||
const sent = useRef(false);
|
||||
const [params] = useSearchParams();
|
||||
const generation = params.get("gen") ?? "";
|
||||
const sent = useRef<string | null>(null);
|
||||
|
||||
useEffect(() => {
|
||||
if (!isReady || sent.current) return;
|
||||
sent.current = true;
|
||||
void Events.Emit(EVENT_WINDOW_PAINTED);
|
||||
}, [isReady]);
|
||||
if (!isReady || sent.current === generation) return;
|
||||
sent.current = generation;
|
||||
void Events.Emit(EVENT_WINDOW_PAINTED, generation);
|
||||
}, [isReady, generation]);
|
||||
|
||||
return null;
|
||||
};
|
||||
|
||||
@@ -1,23 +1,27 @@
|
||||
import { useLayoutEffect, useRef } from "react";
|
||||
import { Window } from "@wailsio/runtime";
|
||||
import { useSearchParams } from "react-router-dom";
|
||||
import { Events, Window } from "@wailsio/runtime";
|
||||
import i18next from "@/lib/i18n";
|
||||
import { isLinux } from "@/lib/platform";
|
||||
|
||||
const EVENT_WINDOW_PAINTED = "netbird:window-painted";
|
||||
|
||||
// Sizes the current Wails window to the measured content height (keeping `width`),
|
||||
// then shows it. Re-applies on content resize and language change.
|
||||
// then reports it as painted so Go shows it. Re-applies on content resize and language change.
|
||||
export function useAutoSizeWindow<T extends HTMLElement>(width: number, ready: boolean = true) {
|
||||
const ref = useRef<T | null>(null);
|
||||
const [params] = useSearchParams();
|
||||
const generation = params.get("gen") ?? "";
|
||||
useLayoutEffect(() => {
|
||||
const el = ref.current;
|
||||
if (!el) return;
|
||||
let shown = false;
|
||||
let painted = false;
|
||||
let raf1 = 0;
|
||||
let raf2 = 0;
|
||||
const showOnce = () => {
|
||||
if (shown) return;
|
||||
shown = true;
|
||||
Window.Show().catch(() => {});
|
||||
Window.Focus().catch(() => {});
|
||||
const paintedOnce = () => {
|
||||
if (painted) return;
|
||||
painted = true;
|
||||
Events.Emit(EVENT_WINDOW_PAINTED, generation).catch(() => {});
|
||||
};
|
||||
const apply = async () => {
|
||||
if (!ready) return;
|
||||
@@ -33,7 +37,7 @@ export function useAutoSizeWindow<T extends HTMLElement>(width: number, ready: b
|
||||
await Window.SetMaxSize(width, targetH);
|
||||
}
|
||||
await Window.SetSize(width, targetH);
|
||||
showOnce();
|
||||
paintedOnce();
|
||||
} catch {
|
||||
// window gone / not ready — ignore
|
||||
}
|
||||
@@ -55,6 +59,6 @@ export function useAutoSizeWindow<T extends HTMLElement>(width: number, ready: b
|
||||
cancelAnimationFrame(raf2);
|
||||
i18next.off("languageChanged", scheduleApply);
|
||||
};
|
||||
}, [width, ready]);
|
||||
}, [width, ready, generation]);
|
||||
return ref;
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { useCallback, useEffect, useRef } from "react";
|
||||
import { useCallback } from "react";
|
||||
import { useTranslation } from "react-i18next";
|
||||
import { useSearchParams } from "react-router-dom";
|
||||
import { Events } from "@wailsio/runtime";
|
||||
@@ -21,7 +21,6 @@ export default function LoginWaitingForBrowserDialog() {
|
||||
const [params] = useSearchParams();
|
||||
const uri = params.get("uri") ?? "";
|
||||
const contentRef = useAutoSizeWindow<HTMLDivElement>(WINDOW_WIDTH);
|
||||
const openedRef = useRef(false);
|
||||
|
||||
const reportOpenFailure = useCallback(
|
||||
(e: unknown) => {
|
||||
@@ -33,13 +32,6 @@ export default function LoginWaitingForBrowserDialog() {
|
||||
[t],
|
||||
);
|
||||
|
||||
// Open the browser only after mount, or it lands on top of the still-hidden popup.
|
||||
useEffect(() => {
|
||||
if (!uri || openedRef.current) return;
|
||||
openedRef.current = true;
|
||||
Connection.OpenURL(uri).catch(reportOpenFailure);
|
||||
}, [uri, reportOpenFailure]);
|
||||
|
||||
const tryAgain = useCallback(() => {
|
||||
if (!uri) return;
|
||||
Connection.OpenURL(uri).catch(reportOpenFailure);
|
||||
|
||||
@@ -205,19 +205,7 @@ func (s *Connection) Down(ctx context.Context) error {
|
||||
// window.open, so the SSO verification page can't pop inline. Honors $BROWSER
|
||||
// before the platform default.
|
||||
func (s *Connection) OpenURL(url string) error {
|
||||
if browser := os.Getenv("BROWSER"); browser != "" {
|
||||
return exec.Command(browser, url).Start()
|
||||
}
|
||||
switch runtime.GOOS {
|
||||
case "windows":
|
||||
return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start()
|
||||
case "darwin":
|
||||
return exec.Command("open", url).Start()
|
||||
case "linux":
|
||||
return exec.Command("xdg-open", url).Start()
|
||||
default:
|
||||
return fmt.Errorf("unsupported platform")
|
||||
}
|
||||
return openURL(url)
|
||||
}
|
||||
|
||||
func (s *Connection) Logout(ctx context.Context, p LogoutParams) error {
|
||||
@@ -288,3 +276,19 @@ func (s *Connection) waitSSOLogin(ctx context.Context, p WaitSSOParams) (string,
|
||||
func (s *Connection) classifyDaemonError(err error) *ClientError {
|
||||
return s.classifier.classify(err)
|
||||
}
|
||||
|
||||
func openURL(url string) error {
|
||||
if browser := os.Getenv("BROWSER"); browser != "" {
|
||||
return exec.Command(browser, url).Start()
|
||||
}
|
||||
switch runtime.GOOS {
|
||||
case "windows":
|
||||
return exec.Command("rundll32", "url.dll,FileProtocolHandler", url).Start()
|
||||
case "darwin":
|
||||
return exec.Command("open", url).Start()
|
||||
case "linux":
|
||||
return exec.Command("xdg-open", url).Start()
|
||||
default:
|
||||
return fmt.Errorf("unsupported platform")
|
||||
}
|
||||
}
|
||||
|
||||
+364
-120
@@ -5,6 +5,7 @@ package services
|
||||
import (
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
@@ -26,6 +27,16 @@ type windowOp func(w *application.WebviewWindow, created bool)
|
||||
|
||||
type windowCloser func(w *application.WebviewWindow)
|
||||
|
||||
// hideableWindow is the slice of application.Window the hide/restore bookkeeping needs.
|
||||
// Narrow enough to fake in tests, which application.Window itself is not: it carries
|
||||
// unexported methods.
|
||||
type hideableWindow interface {
|
||||
Show() application.Window
|
||||
Hide() application.Window
|
||||
IsVisible() bool
|
||||
Name() string
|
||||
}
|
||||
|
||||
// EventTriggerLogin asks the frontend's startLogin() to begin an SSO flow.
|
||||
const EventTriggerLogin = "trigger-login"
|
||||
|
||||
@@ -37,7 +48,10 @@ const EventSettingsOpen = "netbird:settings:open"
|
||||
|
||||
const EventWindowPainted = "netbird:window-painted"
|
||||
|
||||
const paintedFallback = 2 * time.Second
|
||||
// generationParam carries the painted-report token in each dialog's start URL.
|
||||
const generationParam = "gen"
|
||||
|
||||
const paintedFallback = 3 * time.Second
|
||||
|
||||
const headlessTeardownDelay = 2 * time.Second
|
||||
|
||||
@@ -201,6 +215,12 @@ func DialogWindowOptions(name, title, url string, linuxIcon []byte) application.
|
||||
}
|
||||
}
|
||||
|
||||
// hiddenWindow records a window hidden by owner, the name of the popup that hid it.
|
||||
type hiddenWindow struct {
|
||||
win hideableWindow
|
||||
owner string
|
||||
}
|
||||
|
||||
type WindowManager struct {
|
||||
app *application.App
|
||||
mainWindow *application.WebviewWindow
|
||||
@@ -213,19 +233,35 @@ type WindowManager struct {
|
||||
installProgress *application.WebviewWindow
|
||||
welcome *application.WebviewWindow
|
||||
errorDialog *application.WebviewWindow
|
||||
// hiddenForLogin holds windows hidden while the BrowserLogin popup is open, restored on close.
|
||||
hiddenForLogin []application.Window
|
||||
mu sync.Mutex
|
||||
newMain func(startURL string) *application.WebviewWindow
|
||||
creating map[string]bool
|
||||
pendingOps map[string][]windowOp
|
||||
pendingClose map[string]windowCloser
|
||||
restoreGen uint64
|
||||
ready map[uint]bool
|
||||
// hiddenWindows holds windows hidden while a popup owns the screen, each tagged with
|
||||
// the popup that hid it so closing one popup cannot restore what another still hides.
|
||||
hiddenWindows []hiddenWindow
|
||||
hiding map[string]bool
|
||||
// allWindows and raiseMain are the seams the hide/restore tests replace; both are nil
|
||||
// in production, where the Wails app and the platform helper are used directly.
|
||||
allWindows func() []hideableWindow
|
||||
raiseMain func()
|
||||
mu sync.Mutex
|
||||
newMain func(startURL string) *application.WebviewWindow
|
||||
creating map[string]bool
|
||||
pendingOps map[string][]windowOp
|
||||
pendingClose map[string]windowCloser
|
||||
restoreGen map[string]uint64
|
||||
// painted gates showing a window: set by the frontend's first render, or by the
|
||||
// fallback timer so a webview that never wakes up still becomes visible.
|
||||
painted map[uint]bool
|
||||
// mounted gates emitting to a window: set only by a real frontend report, since an
|
||||
// event emitted to a frontend that has not subscribed yet is dropped, not queued.
|
||||
mounted map[uint]bool
|
||||
showPending map[uint]bool
|
||||
pendingTab map[uint]string
|
||||
pendingEmits map[uint][]string
|
||||
fallbackTimers map[uint]*time.Timer
|
||||
afterShow map[uint]func()
|
||||
// generation maps a window name to the token stamped into its current start URL, so a
|
||||
// painted report from a replaced window can be told apart from the live one's.
|
||||
generation map[string]uint64
|
||||
lastGeneration uint64
|
||||
headlessMain bool
|
||||
headlessTimer *time.Timer
|
||||
// recenterOnShow is set only on the minimal-WM/XEmbed path, where the WM neither centers nor
|
||||
@@ -243,11 +279,16 @@ func NewWindowManager(app *application.App, mainWindow *application.WebviewWindo
|
||||
creating: map[string]bool{},
|
||||
pendingOps: map[string][]windowOp{},
|
||||
pendingClose: map[string]windowCloser{},
|
||||
ready: map[uint]bool{},
|
||||
restoreGen: map[string]uint64{},
|
||||
hiding: map[string]bool{},
|
||||
painted: map[uint]bool{},
|
||||
mounted: map[uint]bool{},
|
||||
showPending: map[uint]bool{},
|
||||
pendingTab: map[uint]string{},
|
||||
pendingEmits: map[uint][]string{},
|
||||
fallbackTimers: map[uint]*time.Timer{},
|
||||
afterShow: map[uint]func(){},
|
||||
generation: map[string]uint64{},
|
||||
}
|
||||
s.watchPainted()
|
||||
s.watchTriggerLogin()
|
||||
@@ -307,13 +348,13 @@ func (s *WindowManager) OpenSettings(tab string) {
|
||||
|
||||
s.withWindow(windowSettings, &s.settings, s.newSettingsWindow, func(w *application.WebviewWindow, _ bool) {
|
||||
s.mu.Lock()
|
||||
ready := s.ready[w.ID()]
|
||||
if !ready {
|
||||
mounted := s.mounted[w.ID()]
|
||||
if !mounted {
|
||||
s.pendingTab[w.ID()] = target
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
if ready {
|
||||
if mounted {
|
||||
s.app.Event.Emit(EventSettingsOpen, target)
|
||||
}
|
||||
s.showWhenReady(w)
|
||||
@@ -327,21 +368,37 @@ func (s *WindowManager) OpenBrowserLogin(uri string) {
|
||||
startURL = "/#/dialog/browser-login?uri=" + url.QueryEscape(uri)
|
||||
}
|
||||
s.withWindow(windowBrowserLogin, &s.browserLogin, func() *application.WebviewWindow {
|
||||
return s.newBrowserLoginWindow(startURL)
|
||||
return s.newBrowserLoginWindow(s.stampGeneration(windowBrowserLogin, startURL))
|
||||
}, func(w *application.WebviewWindow, created bool) {
|
||||
if created {
|
||||
s.centerOnCursorScreen(w)
|
||||
return
|
||||
}
|
||||
if uri != "" {
|
||||
w.SetURL(startURL)
|
||||
if !created && uri != "" {
|
||||
w.SetURL(s.stampGeneration(windowBrowserLogin, startURL))
|
||||
}
|
||||
s.centerOnCursorScreen(w)
|
||||
w.Show()
|
||||
w.Focus()
|
||||
s.showThenOpenBrowser(w, uri)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *WindowManager) showThenOpenBrowser(w *application.WebviewWindow, uri string) {
|
||||
if uri != "" {
|
||||
s.mu.Lock()
|
||||
s.afterShow[w.ID()] = func() { s.openBrowser(uri) }
|
||||
s.mu.Unlock()
|
||||
}
|
||||
s.showWhenReady(w)
|
||||
}
|
||||
|
||||
func (s *WindowManager) openBrowser(uri string) {
|
||||
if uri == "" {
|
||||
return
|
||||
}
|
||||
go func() {
|
||||
if err := openURL(uri); err != nil {
|
||||
log.Errorf("open browser for SSO login: %v", err)
|
||||
s.OpenError(s.title("browserLogin.openFailedTitle"), err.Error(), "")
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func (s *WindowManager) newBrowserLoginWindow(startURL string) *application.WebviewWindow {
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
opts := DialogWindowOptions(windowBrowserLogin, s.title("window.title.signIn"), startURL, s.linuxIcon)
|
||||
@@ -360,12 +417,14 @@ func (s *WindowManager) newBrowserLoginWindow(startURL string) *application.Webv
|
||||
if userClosed {
|
||||
s.browserLogin = nil
|
||||
}
|
||||
s.forgetWindowLocked(w)
|
||||
s.mu.Unlock()
|
||||
if userClosed {
|
||||
s.restoreHiddenWindows()
|
||||
s.restoreHiddenWindows(windowBrowserLogin)
|
||||
s.app.Event.Emit(EventBrowserLoginCancel)
|
||||
}
|
||||
})
|
||||
s.armReady(w)
|
||||
return w
|
||||
}
|
||||
|
||||
@@ -386,13 +445,11 @@ func (s *WindowManager) InstallProgressWindow() *application.WebviewWindow {
|
||||
}
|
||||
|
||||
func (s *WindowManager) CloseBrowserLogin() {
|
||||
// The WindowClosing hook no-ops on a programmatic close, so restore here —
|
||||
// but only if a popup was actually open. The frontend calls this even when no
|
||||
// popup was ever shown (e.g. resetDialog() after an early RequestExtend failure,
|
||||
// or connection.ts's catch path), and hiddenForLogin is shared with
|
||||
// OpenInstallProgress, so an unconditional restore could re-show windows a
|
||||
// still-running install-progress is hiding.
|
||||
s.closeWindow(windowBrowserLogin, &s.browserLogin, s.restoreAndClose)
|
||||
// The WindowClosing hook no-ops on a programmatic close, so the closer restores.
|
||||
// The frontend calls this even when no popup was ever shown (resetDialog() after an
|
||||
// early RequestExtend failure, or connection.ts's catch path); closeWindow skips the
|
||||
// closer then, and an owner-scoped restore cannot touch what install-progress hides.
|
||||
s.closeWindow(windowBrowserLogin, &s.browserLogin, s.restoringCloser(windowBrowserLogin))
|
||||
}
|
||||
|
||||
// OpenSessionExpiration shows the countdown warning on the cursor's display; seconds seeds
|
||||
@@ -404,16 +461,13 @@ func (s *WindowManager) OpenSessionExpiration(seconds int, deadlineUnixMilli int
|
||||
startURL += "&deadline=" + strconv.FormatInt(deadlineUnixMilli, 10)
|
||||
}
|
||||
s.withWindow(windowSessionExpiration, &s.sessionExpiration, func() *application.WebviewWindow {
|
||||
return s.newSessionExpirationWindow(startURL)
|
||||
return s.newSessionExpirationWindow(s.stampGeneration(windowSessionExpiration, startURL))
|
||||
}, func(w *application.WebviewWindow, created bool) {
|
||||
if created {
|
||||
s.centerOnCursorScreen(w)
|
||||
return
|
||||
if !created {
|
||||
w.SetURL(s.stampGeneration(windowSessionExpiration, startURL))
|
||||
}
|
||||
w.SetURL(startURL)
|
||||
s.centerOnCursorScreen(w)
|
||||
w.Show()
|
||||
w.Focus()
|
||||
s.showWhenReady(w)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -427,8 +481,10 @@ func (s *WindowManager) newSessionExpirationWindow(startURL string) *application
|
||||
if s.sessionExpiration == w {
|
||||
s.sessionExpiration = nil
|
||||
}
|
||||
s.forgetWindowLocked(w)
|
||||
s.mu.Unlock()
|
||||
})
|
||||
s.armReady(w)
|
||||
return w
|
||||
}
|
||||
|
||||
@@ -440,20 +496,20 @@ func (s *WindowManager) CloseSessionExpiration() {
|
||||
// closes the browser-login popup and the session-expiration window together.
|
||||
func (s *WindowManager) CloseRenewFlow() {
|
||||
s.mu.Lock()
|
||||
bl := s.takeWindowLocked(windowBrowserLogin, &s.browserLogin, s.restoreAndClose)
|
||||
bl := s.takeWindowLocked(windowBrowserLogin, &s.browserLogin, s.restoringCloser(windowBrowserLogin))
|
||||
se := s.takeWindowLocked(windowSessionExpiration, &s.sessionExpiration, closeOnly)
|
||||
if se != nil {
|
||||
kept := s.hiddenForLogin[:0]
|
||||
for _, w := range s.hiddenForLogin {
|
||||
if w != se {
|
||||
kept = append(kept, w)
|
||||
kept := s.hiddenWindows[:0]
|
||||
for _, hidden := range s.hiddenWindows {
|
||||
if !sameWindow(hidden.win, se) {
|
||||
kept = append(kept, hidden)
|
||||
}
|
||||
}
|
||||
s.hiddenForLogin = kept
|
||||
s.hiddenWindows = kept
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
s.restoreHiddenWindows()
|
||||
s.restoreHiddenWindows(windowBrowserLogin)
|
||||
// Close after unlock so the re-entrant handlers can take s.mu.
|
||||
if bl != nil {
|
||||
bl.Close()
|
||||
@@ -471,14 +527,12 @@ func (s *WindowManager) OpenInstallProgress(version string) {
|
||||
startURL = "/#/dialog/install-progress?version=" + url.QueryEscape(version)
|
||||
}
|
||||
s.withWindow(windowInstallProgress, &s.installProgress, func() *application.WebviewWindow {
|
||||
return s.newInstallProgressWindow(startURL)
|
||||
return s.newInstallProgressWindow(s.stampGeneration(windowInstallProgress, startURL))
|
||||
}, func(w *application.WebviewWindow, created bool) {
|
||||
if !created {
|
||||
w.SetURL(startURL)
|
||||
w.Show()
|
||||
w.Focus()
|
||||
w.SetURL(s.stampGeneration(windowInstallProgress, startURL))
|
||||
}
|
||||
s.centerWhenReady(w)
|
||||
s.showWhenReady(w)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -489,32 +543,33 @@ func (s *WindowManager) newInstallProgressWindow(startURL string) *application.W
|
||||
)
|
||||
w.OnWindowEvent(events.Common.WindowClosing, func(_ *application.WindowEvent) {
|
||||
s.mu.Lock()
|
||||
if s.installProgress == w {
|
||||
userClosed := s.installProgress == w
|
||||
if userClosed {
|
||||
s.installProgress = nil
|
||||
}
|
||||
s.forgetWindowLocked(w)
|
||||
s.mu.Unlock()
|
||||
s.restoreHiddenWindows()
|
||||
if userClosed {
|
||||
s.restoreHiddenWindows(windowInstallProgress)
|
||||
}
|
||||
})
|
||||
s.armReady(w)
|
||||
return w
|
||||
}
|
||||
|
||||
func (s *WindowManager) CloseInstallProgress() {
|
||||
s.closeWindow(windowInstallProgress, &s.installProgress, closeOnly)
|
||||
s.closeWindow(windowInstallProgress, &s.installProgress, s.restoringCloser(windowInstallProgress))
|
||||
}
|
||||
|
||||
// OpenWelcome shows the first-launch onboarding window. Singleton, destroyed on close.
|
||||
func (s *WindowManager) OpenWelcome() {
|
||||
s.withWindow(windowWelcome, &s.welcome, s.newWelcomeWindow, func(w *application.WebviewWindow, created bool) {
|
||||
if !created {
|
||||
w.Show()
|
||||
w.Focus()
|
||||
}
|
||||
s.centerWhenReady(w)
|
||||
s.withWindow(windowWelcome, &s.welcome, s.newWelcomeWindow, func(w *application.WebviewWindow, _ bool) {
|
||||
s.showWhenReady(w)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *WindowManager) newWelcomeWindow() *application.WebviewWindow {
|
||||
opts := DialogWindowOptions(windowWelcome, s.title("window.title.welcome"), "/#/dialog/welcome", s.linuxIcon)
|
||||
opts := DialogWindowOptions(windowWelcome, s.title("window.title.welcome"), s.stampGeneration(windowWelcome, "/#/dialog/welcome"), s.linuxIcon)
|
||||
opts.Width = 420
|
||||
opts.InitialPosition = application.WindowCentered
|
||||
w := s.app.Window.NewWithOptions(opts)
|
||||
@@ -523,8 +578,10 @@ func (s *WindowManager) newWelcomeWindow() *application.WebviewWindow {
|
||||
if s.welcome == w {
|
||||
s.welcome = nil
|
||||
}
|
||||
s.forgetWindowLocked(w)
|
||||
s.mu.Unlock()
|
||||
})
|
||||
s.armReady(w)
|
||||
return w
|
||||
}
|
||||
|
||||
@@ -542,14 +599,12 @@ func (s *WindowManager) OpenError(title, message, command string) {
|
||||
}
|
||||
startURL := errorDialogURL(title, message, command)
|
||||
s.withWindow(windowError, &s.errorDialog, func() *application.WebviewWindow {
|
||||
return s.newErrorWindow(startURL)
|
||||
return s.newErrorWindow(s.stampGeneration(windowError, startURL))
|
||||
}, func(w *application.WebviewWindow, created bool) {
|
||||
if !created {
|
||||
w.SetURL(startURL)
|
||||
w.Show()
|
||||
w.Focus()
|
||||
w.SetURL(s.stampGeneration(windowError, startURL))
|
||||
}
|
||||
s.centerWhenReady(w)
|
||||
s.showWhenReady(w)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -562,8 +617,10 @@ func (s *WindowManager) newErrorWindow(startURL string) *application.WebviewWind
|
||||
if s.errorDialog == w {
|
||||
s.errorDialog = nil
|
||||
}
|
||||
s.forgetWindowLocked(w)
|
||||
s.mu.Unlock()
|
||||
})
|
||||
s.armReady(w)
|
||||
return w
|
||||
}
|
||||
|
||||
@@ -589,14 +646,14 @@ func (s *WindowManager) ShowMainAndEmit(event string) {
|
||||
s.ensureMain("/", func(w *application.WebviewWindow, _ bool) {
|
||||
id := w.ID()
|
||||
s.mu.Lock()
|
||||
ready := s.ready[id]
|
||||
if !ready {
|
||||
mounted := s.mounted[id]
|
||||
if !mounted {
|
||||
s.pendingEmits[id] = append(s.pendingEmits[id], event)
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
s.showWhenReady(w)
|
||||
if ready {
|
||||
if mounted {
|
||||
s.app.Event.Emit(event)
|
||||
}
|
||||
})
|
||||
@@ -741,31 +798,66 @@ func (s *WindowManager) releaseCreationLocked(name string) {
|
||||
delete(s.pendingClose, name)
|
||||
}
|
||||
|
||||
func (s *WindowManager) restoreAndClose(w *application.WebviewWindow) {
|
||||
s.restoreHiddenWindows()
|
||||
w.Close()
|
||||
func (s *WindowManager) restoringCloser(owner string) windowCloser {
|
||||
return func(w *application.WebviewWindow) {
|
||||
s.restoreHiddenWindows(owner)
|
||||
w.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// armReady starts the fallback that shows w even if its frontend never reports a first
|
||||
// render. The timer starts at creation, because a hidden webview can be suspended before
|
||||
// it reaches WindowRuntimeReady — the very case this fallback covers. That makes the first
|
||||
// budget cover webview boot as well, so the runtime-ready hook rearms it to give the
|
||||
// frontend its own full budget to mount and paint.
|
||||
func (s *WindowManager) armReady(w *application.WebviewWindow) {
|
||||
if w == nil {
|
||||
return
|
||||
}
|
||||
s.armPaintedFallback(w)
|
||||
w.RegisterHook(events.Common.WindowRuntimeReady, func(_ *application.WindowEvent) {
|
||||
timer := time.AfterFunc(paintedFallback, func() {
|
||||
log.Warnf("window %q never reported a first render, showing it anyway", w.Name())
|
||||
s.markReady(w)
|
||||
})
|
||||
s.mu.Lock()
|
||||
s.fallbackTimers[w.ID()] = timer
|
||||
s.mu.Unlock()
|
||||
s.armPaintedFallback(w)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *WindowManager) armPaintedFallback(w *application.WebviewWindow) {
|
||||
id := w.ID()
|
||||
timer := time.AfterFunc(paintedFallback, func() {
|
||||
s.mu.Lock()
|
||||
painted := s.painted[id]
|
||||
s.mu.Unlock()
|
||||
if painted {
|
||||
return
|
||||
}
|
||||
log.Warnf("window %q never reported a first render, showing it anyway", w.Name())
|
||||
s.markPainted(w)
|
||||
})
|
||||
|
||||
s.mu.Lock()
|
||||
if prev := s.fallbackTimers[id]; prev != nil {
|
||||
prev.Stop()
|
||||
}
|
||||
if s.painted[id] {
|
||||
timer.Stop()
|
||||
delete(s.fallbackTimers, id)
|
||||
} else {
|
||||
s.fallbackTimers[id] = timer
|
||||
}
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
func (s *WindowManager) watchPainted() {
|
||||
s.app.Event.On(EventWindowPainted, func(e *application.CustomEvent) {
|
||||
if w := s.windowByName(e.Sender); w != nil {
|
||||
s.markReady(w)
|
||||
w := s.windowByName(e.Sender)
|
||||
if w == nil {
|
||||
return
|
||||
}
|
||||
if !s.matchesGeneration(e.Sender, paintedGeneration(e.Data)) {
|
||||
log.Debugf("ignoring stale painted report for window %q", e.Sender)
|
||||
return
|
||||
}
|
||||
s.markPainted(w)
|
||||
s.markMounted(w)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -777,7 +869,7 @@ func (s *WindowManager) watchTriggerLogin() {
|
||||
s.headlessTimer = nil
|
||||
}
|
||||
w := s.mainWindow
|
||||
ready := w != nil && s.ready[w.ID()]
|
||||
ready := w != nil && s.mounted[w.ID()]
|
||||
s.mu.Unlock()
|
||||
if ready {
|
||||
return
|
||||
@@ -788,7 +880,7 @@ func (s *WindowManager) watchTriggerLogin() {
|
||||
if created {
|
||||
s.headlessMain = true
|
||||
}
|
||||
pending := !s.ready[w.ID()]
|
||||
pending := !s.mounted[w.ID()]
|
||||
if pending {
|
||||
s.pendingEmits[w.ID()] = append(s.pendingEmits[w.ID()], EventTriggerLogin)
|
||||
}
|
||||
@@ -850,18 +942,67 @@ func (s *WindowManager) forgetWindowLocked(w *application.WebviewWindow) {
|
||||
timer.Stop()
|
||||
}
|
||||
delete(s.fallbackTimers, id)
|
||||
delete(s.ready, id)
|
||||
delete(s.painted, id)
|
||||
delete(s.mounted, id)
|
||||
delete(s.showPending, id)
|
||||
delete(s.pendingTab, id)
|
||||
delete(s.pendingEmits, id)
|
||||
delete(s.afterShow, id)
|
||||
|
||||
kept := s.hiddenForLogin[:0]
|
||||
for _, hidden := range s.hiddenForLogin {
|
||||
if hidden != application.Window(w) {
|
||||
kept := s.hiddenWindows[:0]
|
||||
for _, hidden := range s.hiddenWindows {
|
||||
if !sameWindow(hidden.win, w) {
|
||||
kept = append(kept, hidden)
|
||||
}
|
||||
}
|
||||
s.hiddenForLogin = kept
|
||||
s.hiddenWindows = kept
|
||||
}
|
||||
|
||||
func (s *WindowManager) stampGeneration(name, startURL string) string {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.lastGeneration++
|
||||
s.generation[name] = s.lastGeneration
|
||||
return appendGeneration(startURL, s.lastGeneration)
|
||||
}
|
||||
|
||||
func (s *WindowManager) matchesGeneration(name string, gen uint64) bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
want, tracked := s.generation[name]
|
||||
if !tracked {
|
||||
return true
|
||||
}
|
||||
return want == gen
|
||||
}
|
||||
|
||||
func (s *WindowManager) hideableWindows() []hideableWindow {
|
||||
if s.allWindows != nil {
|
||||
return s.allWindows()
|
||||
}
|
||||
all := s.app.Window.GetAll()
|
||||
windows := make([]hideableWindow, 0, len(all))
|
||||
for _, w := range all {
|
||||
windows = append(windows, w)
|
||||
}
|
||||
return windows
|
||||
}
|
||||
|
||||
func (s *WindowManager) isMainWindow(w hideableWindow, mainWindow *application.WebviewWindow) bool {
|
||||
if s.allWindows != nil {
|
||||
return w != nil && w.Name() == windowMain
|
||||
}
|
||||
return sameWindow(w, mainWindow)
|
||||
}
|
||||
|
||||
func (s *WindowManager) raiseMainWindow(mainWindow *application.WebviewWindow) {
|
||||
if s.raiseMain != nil {
|
||||
s.raiseMain()
|
||||
return
|
||||
}
|
||||
if mainWindow != nil {
|
||||
raiseToForeground(mainWindow)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *WindowManager) windowByName(name string) *application.WebviewWindow {
|
||||
@@ -872,24 +1013,50 @@ func (s *WindowManager) windowByName(name string) *application.WebviewWindow {
|
||||
return s.mainWindow
|
||||
case windowSettings:
|
||||
return s.settings
|
||||
case windowBrowserLogin:
|
||||
return s.browserLogin
|
||||
case windowSessionExpiration:
|
||||
return s.sessionExpiration
|
||||
case windowInstallProgress:
|
||||
return s.installProgress
|
||||
case windowWelcome:
|
||||
return s.welcome
|
||||
case windowError:
|
||||
return s.errorDialog
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (s *WindowManager) markReady(w *application.WebviewWindow) {
|
||||
func (s *WindowManager) markPainted(w *application.WebviewWindow) {
|
||||
id := w.ID()
|
||||
s.mu.Lock()
|
||||
already := s.ready[id]
|
||||
s.ready[id] = true
|
||||
already := s.painted[id]
|
||||
s.painted[id] = true
|
||||
wanted := s.showPending[id]
|
||||
tab, hasTab := s.pendingTab[id]
|
||||
emits := s.pendingEmits[id]
|
||||
delete(s.showPending, id)
|
||||
if timer := s.fallbackTimers[id]; timer != nil {
|
||||
timer.Stop()
|
||||
delete(s.fallbackTimers, id)
|
||||
}
|
||||
delete(s.showPending, id)
|
||||
s.mu.Unlock()
|
||||
|
||||
if already || !wanted {
|
||||
return
|
||||
}
|
||||
s.showNow(w)
|
||||
}
|
||||
|
||||
// markMounted records that the window's frontend is subscribed, and flushes the events
|
||||
// held back for it. The fallback timer never calls this: showing a blank window is
|
||||
// recoverable, emitting into a frontend that cannot hear it is not.
|
||||
func (s *WindowManager) markMounted(w *application.WebviewWindow) {
|
||||
id := w.ID()
|
||||
s.mu.Lock()
|
||||
already := s.mounted[id]
|
||||
s.mounted[id] = true
|
||||
tab, hasTab := s.pendingTab[id]
|
||||
emits := s.pendingEmits[id]
|
||||
delete(s.pendingTab, id)
|
||||
delete(s.pendingEmits, id)
|
||||
s.mu.Unlock()
|
||||
@@ -902,10 +1069,6 @@ func (s *WindowManager) markReady(w *application.WebviewWindow) {
|
||||
s.app.Event.Emit(EventSettingsOpen, tab)
|
||||
}
|
||||
|
||||
if wanted {
|
||||
s.showNow(w)
|
||||
}
|
||||
|
||||
for _, event := range emits {
|
||||
s.app.Event.Emit(event)
|
||||
}
|
||||
@@ -918,18 +1081,19 @@ func (s *WindowManager) showWhenReady(w *application.WebviewWindow) {
|
||||
|
||||
id := w.ID()
|
||||
s.mu.Lock()
|
||||
ready := s.ready[id]
|
||||
if !ready {
|
||||
painted := s.painted[id]
|
||||
if !painted {
|
||||
s.showPending[id] = true
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
if ready {
|
||||
if painted {
|
||||
s.showNow(w)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *WindowManager) showNow(w *application.WebviewWindow) {
|
||||
id := w.ID()
|
||||
s.mu.Lock()
|
||||
if w == s.mainWindow {
|
||||
s.headlessMain = false
|
||||
@@ -938,10 +1102,15 @@ func (s *WindowManager) showNow(w *application.WebviewWindow) {
|
||||
s.headlessTimer = nil
|
||||
}
|
||||
}
|
||||
after := s.afterShow[id]
|
||||
delete(s.afterShow, id)
|
||||
s.mu.Unlock()
|
||||
w.Show()
|
||||
w.Focus()
|
||||
s.centerWhenReady(w)
|
||||
if after != nil {
|
||||
after()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *WindowManager) ShowMainAt(url string) {
|
||||
@@ -1070,13 +1239,19 @@ func (s *WindowManager) retitleAll() {
|
||||
}
|
||||
}
|
||||
|
||||
// hideOtherWindows hides every visible window except keepName, recording them against
|
||||
// keepName so only its own restore brings them back. A window already hidden by an
|
||||
// earlier popup is skipped, leaving it tagged to the popup that actually hid it. The
|
||||
// per-owner generation catches a restore for keepName that ran between the snapshot and
|
||||
// the record, in which case the windows are re-shown rather than stranded.
|
||||
func (s *WindowManager) hideOtherWindows(keepName string) {
|
||||
s.mu.Lock()
|
||||
gen := s.restoreGen
|
||||
s.hiding[keepName] = true
|
||||
gen := s.restoreGen[keepName]
|
||||
s.mu.Unlock()
|
||||
|
||||
var hidden []application.Window
|
||||
for _, w := range s.app.Window.GetAll() {
|
||||
var hidden []hideableWindow
|
||||
for _, w := range s.hideableWindows() {
|
||||
if w == nil || w.Name() == keepName || !w.IsVisible() {
|
||||
continue
|
||||
}
|
||||
@@ -1088,9 +1263,11 @@ func (s *WindowManager) hideOtherWindows(keepName string) {
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
restored := s.restoreGen != gen
|
||||
restored := s.restoreGen[keepName] != gen
|
||||
if !restored {
|
||||
s.hiddenForLogin = append(s.hiddenForLogin, hidden...)
|
||||
for _, w := range hidden {
|
||||
s.hiddenWindows = append(s.hiddenWindows, hiddenWindow{win: w, owner: keepName})
|
||||
}
|
||||
}
|
||||
s.mu.Unlock()
|
||||
if !restored {
|
||||
@@ -1101,33 +1278,58 @@ func (s *WindowManager) hideOtherWindows(keepName string) {
|
||||
}
|
||||
}
|
||||
|
||||
// restoreHiddenWindows re-shows windows hidden by hideOtherWindows. If the main
|
||||
// window was among them, raiseToForeground lifts it above the SSO browser, which
|
||||
// still owns the foreground — a plain Show/Focus would be demoted to a taskbar
|
||||
// flash and leave it stranded behind.
|
||||
func (s *WindowManager) restoreHiddenWindows() {
|
||||
// restoreHiddenWindows re-shows the windows owner hid, unless another popup still covers
|
||||
// them, in which case they are handed to that popup. If the main window was among them,
|
||||
// raiseToForeground lifts it above the SSO browser, which still owns the foreground — a
|
||||
// plain Show/Focus would be demoted to a taskbar flash and leave it stranded behind.
|
||||
func (s *WindowManager) restoreHiddenWindows(owner string) {
|
||||
s.mu.Lock()
|
||||
hidden := s.hiddenForLogin
|
||||
s.hiddenForLogin = nil
|
||||
s.restoreGen++
|
||||
mainWindow := s.mainWindow
|
||||
delete(s.hiding, owner)
|
||||
var restore []hideableWindow
|
||||
kept := s.hiddenWindows[:0]
|
||||
for _, hidden := range s.hiddenWindows {
|
||||
if hidden.owner != owner {
|
||||
kept = append(kept, hidden)
|
||||
continue
|
||||
}
|
||||
if coverer, covered := s.coveringPopupLocked(hidden.win); covered {
|
||||
hidden.owner = coverer
|
||||
kept = append(kept, hidden)
|
||||
continue
|
||||
}
|
||||
if hidden.win != nil {
|
||||
restore = append(restore, hidden.win)
|
||||
}
|
||||
}
|
||||
s.hiddenWindows = kept
|
||||
s.restoreGen[owner]++
|
||||
s.mu.Unlock()
|
||||
|
||||
mainRestored := false
|
||||
for _, w := range hidden {
|
||||
if w == nil {
|
||||
continue
|
||||
}
|
||||
for _, w := range restore {
|
||||
w.Show()
|
||||
if w == mainWindow {
|
||||
if s.isMainWindow(w, mainWindow) {
|
||||
mainRestored = true
|
||||
}
|
||||
}
|
||||
if mainRestored && mainWindow != nil {
|
||||
raiseToForeground(mainWindow)
|
||||
if mainRestored {
|
||||
s.raiseMainWindow(mainWindow)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *WindowManager) coveringPopupLocked(w hideableWindow) (string, bool) {
|
||||
if w == nil {
|
||||
return "", false
|
||||
}
|
||||
for name := range s.hiding {
|
||||
if name != w.Name() {
|
||||
return name, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// getScreenBasedOnCursorPosition returns the cursor's display, falling back to the
|
||||
// main-window screen, then nil (OS-default placement).
|
||||
func (s *WindowManager) getScreenBasedOnCursorPosition() *application.Screen {
|
||||
@@ -1169,6 +1371,48 @@ func errorDialogURL(title, message, command string) string {
|
||||
return startURL
|
||||
}
|
||||
|
||||
// appendGeneration adds the painted-report token to a dialog start URL, keeping any
|
||||
// existing query params intact across the "/#/path?params" hash-router form.
|
||||
func appendGeneration(startURL string, gen uint64) string {
|
||||
sep := "?"
|
||||
if strings.Contains(startURL, "?") {
|
||||
sep = "&"
|
||||
}
|
||||
return startURL + sep + generationParam + "=" + strconv.FormatUint(gen, 10)
|
||||
}
|
||||
|
||||
// paintedGeneration reads the token a painted report carries back, returning 0 when the
|
||||
// frontend sent none (an older bundle, or the main window, which is never stamped).
|
||||
func paintedGeneration(data any) uint64 {
|
||||
switch v := data.(type) {
|
||||
case string:
|
||||
gen, err := strconv.ParseUint(v, 10, 64)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return gen
|
||||
case float64:
|
||||
return uint64(v)
|
||||
case []any:
|
||||
if len(v) == 0 {
|
||||
return 0
|
||||
}
|
||||
return paintedGeneration(v[0])
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
// sameWindow reports whether a hidden entry refers to w, comparing through the interface
|
||||
// so a nil entry never matches a live window.
|
||||
func sameWindow(hidden hideableWindow, w *application.WebviewWindow) bool {
|
||||
if hidden == nil || w == nil {
|
||||
return false
|
||||
}
|
||||
other, ok := hidden.(*application.WebviewWindow)
|
||||
return ok && other == w
|
||||
}
|
||||
|
||||
// u32ptr returns a pointer to v, for the optional *uint32 Wails theme fields.
|
||||
func u32ptr(v uint32) *uint32 { return &v }
|
||||
|
||||
|
||||
@@ -17,9 +17,65 @@ func newTestWindowManager() *WindowManager {
|
||||
creating: map[string]bool{},
|
||||
pendingOps: map[string][]windowOp{},
|
||||
pendingClose: map[string]windowCloser{},
|
||||
restoreGen: map[string]uint64{},
|
||||
hiding: map[string]bool{},
|
||||
generation: map[string]uint64{},
|
||||
}
|
||||
}
|
||||
|
||||
type fakeWindow struct {
|
||||
name string
|
||||
visible bool
|
||||
shown int
|
||||
hidden int
|
||||
}
|
||||
|
||||
func newFakeWindow(name string) *fakeWindow {
|
||||
return &fakeWindow{name: name, visible: true}
|
||||
}
|
||||
|
||||
func (f *fakeWindow) Show() application.Window {
|
||||
f.visible = true
|
||||
f.shown++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeWindow) Hide() application.Window {
|
||||
f.visible = false
|
||||
f.hidden++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeWindow) IsVisible() bool { return f.visible }
|
||||
|
||||
func (f *fakeWindow) Name() string { return f.name }
|
||||
|
||||
type fakeDesktop struct {
|
||||
windows []*fakeWindow
|
||||
raised int
|
||||
}
|
||||
|
||||
func newFakeDesktop(s *WindowManager, windows ...*fakeWindow) *fakeDesktop {
|
||||
d := &fakeDesktop{windows: windows}
|
||||
s.allWindows = func() []hideableWindow {
|
||||
all := make([]hideableWindow, 0, len(d.windows))
|
||||
for _, w := range d.windows {
|
||||
all = append(all, w)
|
||||
}
|
||||
return all
|
||||
}
|
||||
s.raiseMain = func() { d.raised++ }
|
||||
return d
|
||||
}
|
||||
|
||||
func ownersOf(hidden []hiddenWindow) []string {
|
||||
owners := make([]string, 0, len(hidden))
|
||||
for _, h := range hidden {
|
||||
owners = append(owners, h.owner)
|
||||
}
|
||||
return owners
|
||||
}
|
||||
|
||||
func waitDone(t *testing.T, done <-chan struct{}, msg string) {
|
||||
t.Helper()
|
||||
select {
|
||||
@@ -339,12 +395,245 @@ func TestCloseRenewFlowDuringBrowserLoginCreationRestoresHiddenWindows(t *testin
|
||||
// Seeded after the call so the deferred closer, not CloseRenewFlow's own
|
||||
// immediate restore, is what has to drain it. A nil entry is skipped by
|
||||
// restoreHiddenWindows, so no Wails window is needed.
|
||||
s.hiddenForLogin = []application.Window{nil}
|
||||
s.hiddenWindows = []hiddenWindow{{owner: windowBrowserLogin}}
|
||||
return &application.WebviewWindow{}
|
||||
}, func(*application.WebviewWindow, bool) {})
|
||||
|
||||
require.Nil(t, s.browserLogin)
|
||||
require.Empty(t, s.hiddenForLogin)
|
||||
require.Empty(t, s.hiddenWindows)
|
||||
require.Empty(t, s.creating)
|
||||
require.Empty(t, s.pendingClose)
|
||||
}
|
||||
|
||||
func TestHideOtherWindowsSkipsKeepNameAndInvisible(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
settings := newFakeWindow(windowSettings)
|
||||
settings.visible = false
|
||||
popup := newFakeWindow(windowBrowserLogin)
|
||||
s := newTestWindowManager()
|
||||
newFakeDesktop(s, main, settings, popup)
|
||||
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
|
||||
require.False(t, main.visible)
|
||||
require.Equal(t, 1, main.hidden)
|
||||
require.Equal(t, 0, settings.hidden, "an already hidden window must not be recorded")
|
||||
require.Equal(t, 0, popup.hidden, "the popup itself must stay visible")
|
||||
require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows))
|
||||
}
|
||||
|
||||
func TestInstallDuringLoginKeepsMainHiddenUntilLoginCloses(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
login := newFakeWindow(windowBrowserLogin)
|
||||
install := newFakeWindow(windowInstallProgress)
|
||||
install.visible = false
|
||||
s := newTestWindowManager()
|
||||
d := newFakeDesktop(s, main, login, install)
|
||||
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
require.False(t, main.visible)
|
||||
|
||||
install.visible = true
|
||||
s.hideOtherWindows(windowInstallProgress)
|
||||
require.False(t, login.visible, "the install popup hides the login popup")
|
||||
|
||||
s.restoreHiddenWindows(windowInstallProgress)
|
||||
require.True(t, login.visible, "the install popup restores the login popup it hid")
|
||||
require.False(t, main.visible, "the main window stays hidden for the login popup")
|
||||
require.Equal(t, 0, d.raised)
|
||||
require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows))
|
||||
|
||||
s.restoreHiddenWindows(windowBrowserLogin)
|
||||
require.True(t, main.visible)
|
||||
require.Equal(t, 1, d.raised, "restoring the main window raises it above the SSO browser")
|
||||
require.Empty(t, s.hiddenWindows)
|
||||
}
|
||||
|
||||
func TestLoginClosingUnderInstallHandsMainToInstall(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
login := newFakeWindow(windowBrowserLogin)
|
||||
install := newFakeWindow(windowInstallProgress)
|
||||
install.visible = false
|
||||
s := newTestWindowManager()
|
||||
d := newFakeDesktop(s, main, login, install)
|
||||
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
install.visible = true
|
||||
s.hideOtherWindows(windowInstallProgress)
|
||||
require.False(t, main.visible)
|
||||
require.False(t, login.visible, "the install popup hides the login popup")
|
||||
|
||||
// The login popup closes while the install popup is still up: the main window it
|
||||
// hid must not resurface under the install popup, it is handed over instead.
|
||||
s.restoreHiddenWindows(windowBrowserLogin)
|
||||
require.False(t, main.visible, "the install popup still covers the main window")
|
||||
require.Equal(t, 0, d.raised)
|
||||
require.Equal(t, []string{windowInstallProgress, windowInstallProgress}, ownersOf(s.hiddenWindows))
|
||||
|
||||
s.restoreHiddenWindows(windowInstallProgress)
|
||||
require.True(t, main.visible, "the install popup restores the handed-over main window")
|
||||
require.Equal(t, 1, d.raised)
|
||||
require.Empty(t, s.hiddenWindows)
|
||||
}
|
||||
|
||||
func TestInstallClosingUnderLoginHandsMainToLogin(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
install := newFakeWindow(windowInstallProgress)
|
||||
login := newFakeWindow(windowBrowserLogin)
|
||||
login.visible = false
|
||||
s := newTestWindowManager()
|
||||
d := newFakeDesktop(s, main, install, login)
|
||||
|
||||
s.hideOtherWindows(windowInstallProgress)
|
||||
login.visible = true
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
require.False(t, install.visible, "the login popup hides the install popup")
|
||||
|
||||
s.restoreHiddenWindows(windowInstallProgress)
|
||||
require.False(t, main.visible, "the login popup still covers the main window")
|
||||
require.Equal(t, 0, d.raised)
|
||||
require.Equal(t, []string{windowBrowserLogin, windowBrowserLogin}, ownersOf(s.hiddenWindows))
|
||||
|
||||
s.restoreHiddenWindows(windowBrowserLogin)
|
||||
require.True(t, main.visible)
|
||||
require.Equal(t, 1, d.raised)
|
||||
require.Empty(t, s.hiddenWindows)
|
||||
}
|
||||
|
||||
func TestPopupClosingReshowsTheCoveringPopupItself(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
install := newFakeWindow(windowInstallProgress)
|
||||
login := newFakeWindow(windowBrowserLogin)
|
||||
login.visible = false
|
||||
s := newTestWindowManager()
|
||||
d := newFakeDesktop(s, main, install, login)
|
||||
|
||||
s.hideOtherWindows(windowInstallProgress)
|
||||
login.visible = true
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
|
||||
// The login popup hid the install popup itself; closing the login popup must bring
|
||||
// the install popup back rather than hand it over to its own owner.
|
||||
s.restoreHiddenWindows(windowBrowserLogin)
|
||||
require.True(t, install.visible, "a popup is never handed over to itself")
|
||||
require.False(t, main.visible, "the main window stays with the install popup")
|
||||
require.Equal(t, 0, d.raised)
|
||||
require.Equal(t, []string{windowInstallProgress}, ownersOf(s.hiddenWindows))
|
||||
}
|
||||
|
||||
func TestRestoreHiddenWindowsUnknownOwnerKeepsEverything(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
s := newTestWindowManager()
|
||||
d := newFakeDesktop(s, main)
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
|
||||
s.restoreHiddenWindows(windowWelcome)
|
||||
|
||||
require.False(t, main.visible)
|
||||
require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows))
|
||||
require.Equal(t, 0, d.raised)
|
||||
}
|
||||
|
||||
func TestRestoreHiddenWindowsWithoutMainDoesNotRaise(t *testing.T) {
|
||||
settings := newFakeWindow(windowSettings)
|
||||
s := newTestWindowManager()
|
||||
d := newFakeDesktop(s, settings)
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
|
||||
s.restoreHiddenWindows(windowBrowserLogin)
|
||||
|
||||
require.True(t, settings.visible)
|
||||
require.Equal(t, 0, d.raised)
|
||||
}
|
||||
|
||||
func TestRestoreHiddenWindowsEmptyIsNoop(t *testing.T) {
|
||||
s := newTestWindowManager()
|
||||
require.NotPanics(t, func() { s.restoreHiddenWindows(windowBrowserLogin) })
|
||||
require.Empty(t, s.hiddenWindows)
|
||||
}
|
||||
|
||||
func TestHideOtherWindowsRacingOwnRestoreReshowsWhatItHid(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
s := newTestWindowManager()
|
||||
d := newFakeDesktop(s, main)
|
||||
enumerate := s.allWindows
|
||||
// A restore for the same owner lands between the generation snapshot and the record.
|
||||
s.allWindows = func() []hideableWindow {
|
||||
s.restoreHiddenWindows(windowBrowserLogin)
|
||||
return enumerate()
|
||||
}
|
||||
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
|
||||
require.True(t, main.visible)
|
||||
require.Equal(t, 1, main.hidden)
|
||||
require.Empty(t, s.hiddenWindows)
|
||||
require.Equal(t, 0, d.raised)
|
||||
}
|
||||
|
||||
func TestHideOtherWindowsIgnoresRestoreOfAnotherOwner(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
s := newTestWindowManager()
|
||||
newFakeDesktop(s, main)
|
||||
enumerate := s.allWindows
|
||||
s.allWindows = func() []hideableWindow {
|
||||
s.restoreHiddenWindows(windowInstallProgress)
|
||||
return enumerate()
|
||||
}
|
||||
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
|
||||
require.False(t, main.visible)
|
||||
require.Equal(t, []string{windowBrowserLogin}, ownersOf(s.hiddenWindows))
|
||||
}
|
||||
|
||||
func TestRestoringCloserRestoresOnlyItsOwner(t *testing.T) {
|
||||
main := newFakeWindow(windowMain)
|
||||
s := newTestWindowManager()
|
||||
newFakeDesktop(s, main)
|
||||
s.hideOtherWindows(windowBrowserLogin)
|
||||
s.hiddenWindows = append(s.hiddenWindows, hiddenWindow{owner: windowInstallProgress})
|
||||
|
||||
s.restoringCloser(windowBrowserLogin)(&application.WebviewWindow{})
|
||||
|
||||
require.True(t, main.visible)
|
||||
require.Equal(t, []string{windowInstallProgress}, ownersOf(s.hiddenWindows))
|
||||
}
|
||||
|
||||
func TestStampGenerationTracksLatestPerWindow(t *testing.T) {
|
||||
s := newTestWindowManager()
|
||||
|
||||
first := s.stampGeneration(windowBrowserLogin, "/#/dialog/browser-login")
|
||||
require.Equal(t, "/#/dialog/browser-login?gen=1", first)
|
||||
require.True(t, s.matchesGeneration(windowBrowserLogin, 1))
|
||||
|
||||
second := s.stampGeneration(windowBrowserLogin, "/#/dialog/browser-login?uri=x")
|
||||
require.Equal(t, "/#/dialog/browser-login?uri=x&gen=2", second)
|
||||
require.False(t, s.matchesGeneration(windowBrowserLogin, 1))
|
||||
require.True(t, s.matchesGeneration(windowBrowserLogin, 2))
|
||||
}
|
||||
|
||||
func TestMatchesGenerationUntrackedWindowAccepts(t *testing.T) {
|
||||
s := newTestWindowManager()
|
||||
require.True(t, s.matchesGeneration(windowMain, 0))
|
||||
}
|
||||
|
||||
func TestPaintedGeneration(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
data any
|
||||
want uint64
|
||||
}{
|
||||
{"string", "7", 7},
|
||||
{"float", float64(7), 7},
|
||||
{"slice", []any{"7"}, 7},
|
||||
{"empty slice", []any{}, 0},
|
||||
{"unparsable", "abc", 0},
|
||||
{"nil", nil, 0},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
require.Equal(t, tc.want, paintedGeneration(tc.data))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,14 +1,14 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
|
||||
if ! which curl >/dev/null 2>&1; then
|
||||
if ! command -v curl >/dev/null 2>&1; then
|
||||
echo "This script uses curl fetch OpenID configuration from IDP."
|
||||
echo "Please install curl and re-run the script https://curl.se/"
|
||||
echo ""
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if ! which jq >/dev/null 2>&1; then
|
||||
if ! command -v jq >/dev/null 2>&1; then
|
||||
echo "This script uses jq to load OpenID configuration from IDP."
|
||||
echo "Please install jq and re-run the script https://stedolan.github.io/jq/"
|
||||
echo ""
|
||||
@@ -18,13 +18,13 @@ fi
|
||||
source setup.env
|
||||
source base.setup.env
|
||||
|
||||
if ! which envsubst >/dev/null 2>&1; then
|
||||
if ! command -v envsubst >/dev/null 2>&1; then
|
||||
echo "envsubst is needed to run this script"
|
||||
if [[ $(uname) == "Darwin" ]]; then
|
||||
echo "you can install it with homebrew (https://brew.sh):"
|
||||
echo "brew install gettext"
|
||||
else
|
||||
if which apt-get >/dev/null 2>&1; then
|
||||
if command -v apt-get >/dev/null 2>&1; then
|
||||
echo "you can install it by running"
|
||||
echo "apt-get update && apt-get install gettext-base"
|
||||
else
|
||||
|
||||
@@ -44,7 +44,7 @@ func (r *repository) GetAccountNetwork(ctx context.Context, accountID string) (*
|
||||
}
|
||||
|
||||
func (r *repository) GetAccountPeers(ctx context.Context, accountID string) ([]*peer.Peer, error) {
|
||||
return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
return r.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
}
|
||||
|
||||
func (r *repository) GetAccountByPeerID(ctx context.Context, peerID string) (*types.Account, error) {
|
||||
|
||||
@@ -97,7 +97,7 @@ func (m *managerImpl) GetAllPeers(ctx context.Context, accountID, userID string)
|
||||
return m.store.GetUserPeers(ctx, store.LockingStrengthNone, accountID, userID)
|
||||
}
|
||||
|
||||
return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
return m.store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
}
|
||||
|
||||
func (m *managerImpl) GetPeerAccountID(ctx context.Context, peerID string) (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
|
||||
ClusterAllProxiesPrivate(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)
|
||||
|
||||
@@ -23,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
|
||||
GetClusterAllProxiesPrivate(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)
|
||||
@@ -149,6 +150,11 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string)
|
||||
return m.store.GetClusterSupportsPrivate(ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterAllProxiesPrivate reports whether every active proxy claims the private capability (nil = unreported).
|
||||
func (m Manager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
|
||||
return m.store.GetClusterAllProxiesPrivate(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)
|
||||
|
||||
@@ -105,6 +105,9 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo
|
||||
func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
func (m *mockStore) GetClusterAllProxiesPrivate(_ 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)
|
||||
|
||||
@@ -56,6 +56,20 @@ func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *go
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration)
|
||||
}
|
||||
|
||||
// ClusterAllProxiesPrivate mocks base method.
|
||||
func (m *MockManager) ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ClusterAllProxiesPrivate", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].(*bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ClusterAllProxiesPrivate indicates an expected call of ClusterAllProxiesPrivate.
|
||||
func (mr *MockManagerMockRecorder) ClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterAllProxiesPrivate", reflect.TypeOf((*MockManager)(nil).ClusterAllProxiesPrivate), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterRequireSubdomain mocks base method.
|
||||
func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -84,6 +84,7 @@ type CapabilityProvider interface {
|
||||
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
}
|
||||
|
||||
type Manager struct {
|
||||
@@ -332,6 +333,10 @@ func (m *Manager) persistNewService(ctx context.Context, accountID string, svc *
|
||||
return err
|
||||
}
|
||||
|
||||
if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
|
||||
return err
|
||||
@@ -369,6 +374,43 @@ func (m *Manager) clusterCustomPorts(ctx context.Context, svc *service.Service)
|
||||
return m.capabilities.ClusterSupportsCustomPorts(ctx, svc.ProxyCluster)
|
||||
}
|
||||
|
||||
// validatePrivateClusterTargets rejects cluster and direct upstream targets unless
|
||||
// every active proxy in the service's cluster reports the private capability. The
|
||||
// mapping reaches all proxies in the cluster, so one non-private proxy would serve
|
||||
// these targets too. An unreported capability is treated as unsupported. Must be
|
||||
// called outside a transaction, like clusterCustomPorts.
|
||||
func (m *Manager) validatePrivateClusterTargets(ctx context.Context, targets []*service.Target, cluster string) error {
|
||||
target := firstPrivateClusterTarget(targets)
|
||||
if target == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if private := m.capabilities.ClusterAllProxiesPrivate(ctx, cluster); private != nil && *private {
|
||||
return nil
|
||||
}
|
||||
|
||||
if target.TargetType == service.TargetTypeCluster {
|
||||
return status.Errorf(status.InvalidArgument,
|
||||
"target_type %q requires a proxy cluster with private mode enabled, cluster %s does not support it",
|
||||
service.TargetTypeCluster, cluster)
|
||||
}
|
||||
return status.Errorf(status.InvalidArgument,
|
||||
"direct_upstream requires a proxy cluster with private mode enabled, cluster %s does not support it", cluster)
|
||||
}
|
||||
|
||||
// firstPrivateClusterTarget returns the first target that only a private cluster may serve.
|
||||
func firstPrivateClusterTarget(targets []*service.Target) *service.Target {
|
||||
for _, target := range targets {
|
||||
if target == nil {
|
||||
continue
|
||||
}
|
||||
if target.TargetType == service.TargetTypeCluster || target.Options.DirectUpstream {
|
||||
return target
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ensureL4Port auto-assigns a listen port when needed and validates cluster support.
|
||||
// customPorts must be pre-computed via clusterCustomPorts before entering a transaction.
|
||||
func (m *Manager) ensureL4Port(ctx context.Context, tx store.Store, svc *service.Service, customPorts *bool, serviceUpdate bool) error {
|
||||
@@ -464,6 +506,10 @@ func (m *Manager) persistNewEphemeralService(ctx context.Context, accountID, pee
|
||||
return err
|
||||
}
|
||||
|
||||
if err := m.validatePrivateClusterTargets(ctx, svc.Targets, svc.ProxyCluster); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
if err := m.validateServiceDomain(ctx, transaction, accountID, svc, svc.ProxyCluster); err != nil {
|
||||
return err
|
||||
@@ -584,6 +630,10 @@ func (m *Manager) persistServiceUpdate(ctx context.Context, accountID string, se
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := m.validatePrivateClusterTargets(ctx, service.Targets, effectiveCluster); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Validate subdomain requirement *before* the transaction: the underlying
|
||||
// capability lookup talks to the main DB pool, and SQLite's single-connection
|
||||
// pool would self-deadlock if this ran while the tx already held the only
|
||||
|
||||
@@ -0,0 +1,218 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel/metric/noop"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
// setupPrivateClusterTest wires the real proxy manager as the capability
|
||||
// provider and connects one proxy to testCluster reporting the given private
|
||||
// capability. A nil private connects no proxy, so the capability is unreported.
|
||||
func setupPrivateClusterTest(t *testing.T, private *bool) (*Manager, store.Store) {
|
||||
t.Helper()
|
||||
|
||||
mgr, testStore := setupIntegrationTest(t)
|
||||
|
||||
proxyMgr, err := proxymanager.NewManager(testStore, noop.NewMeterProvider().Meter(""))
|
||||
require.NoError(t, err)
|
||||
mgr.capabilities = proxyMgr
|
||||
|
||||
if private != nil {
|
||||
connectTestProxy(t, proxyMgr, "proxy-1", &proxy.Capabilities{Private: private})
|
||||
}
|
||||
|
||||
return mgr, testStore
|
||||
}
|
||||
|
||||
func connectTestProxy(t *testing.T, proxyMgr *proxymanager.Manager, proxyID string, caps *proxy.Capabilities) {
|
||||
t.Helper()
|
||||
_, err := proxyMgr.Connect(context.Background(), proxyID, "session-"+proxyID, testCluster, "127.0.0.1", "", nil, caps)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func clusterTarget() *rpservice.Target {
|
||||
return &rpservice.Target{
|
||||
TargetId: testCluster,
|
||||
TargetType: rpservice.TargetTypeCluster,
|
||||
Host: "backend.lan",
|
||||
Port: 8080,
|
||||
Protocol: "http",
|
||||
Enabled: true,
|
||||
Options: rpservice.TargetOptions{DirectUpstream: true},
|
||||
}
|
||||
}
|
||||
|
||||
func directUpstreamPeerTarget() *rpservice.Target {
|
||||
return &rpservice.Target{
|
||||
TargetId: testPeerID,
|
||||
TargetType: rpservice.TargetTypePeer,
|
||||
Host: "backend.lan",
|
||||
Port: 8080,
|
||||
Protocol: "http",
|
||||
Enabled: true,
|
||||
Options: rpservice.TargetOptions{DirectUpstream: true},
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateService_PrivateClusterTargets(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
private *bool
|
||||
target *rpservice.Target
|
||||
wantErr string
|
||||
}{
|
||||
{name: "cluster target on private cluster", private: boolPtr(true), target: clusterTarget()},
|
||||
{name: "direct upstream on private cluster", private: boolPtr(true), target: directUpstreamPeerTarget()},
|
||||
{name: "cluster target on non-private cluster", private: boolPtr(false), target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
|
||||
{name: "direct upstream on non-private cluster", private: boolPtr(false), target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
|
||||
{name: "cluster target with unreported capability", private: nil, target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
|
||||
{name: "direct upstream with unreported capability", private: nil, target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupPrivateClusterTest(t, tc.private)
|
||||
|
||||
svc := newTestService("app.test.netbird.io")
|
||||
svc.Targets = []*rpservice.Target{tc.target}
|
||||
|
||||
_, err := mgr.CreateService(ctx, testAccountID, testUserID, svc)
|
||||
|
||||
services, listErr := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
|
||||
require.NoError(t, listErr)
|
||||
|
||||
if tc.wantErr == "" {
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, services, 1, "the service should be persisted")
|
||||
return
|
||||
}
|
||||
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tc.wantErr)
|
||||
sErr, ok := status.FromError(err)
|
||||
require.True(t, ok, "the caller must receive a typed error")
|
||||
assert.Equal(t, status.InvalidArgument, sErr.Type(), "the rejection should be an invalid argument")
|
||||
assert.Empty(t, services, "a rejected service must not be persisted")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A cluster where only some proxies run in private mode must not accept these
|
||||
// targets: the mapping is delivered to every proxy in the cluster, so the
|
||||
// non-private ones would serve the target from their host network as well.
|
||||
func TestCreateService_MixedClusterRejectsPrivateTargets(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
secondCaps *proxy.Capabilities
|
||||
}{
|
||||
{name: "second proxy reports not private", secondCaps: &proxy.Capabilities{Private: boolPtr(false)}},
|
||||
{name: "second proxy predates capability reporting", secondCaps: nil},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
for _, target := range []*rpservice.Target{clusterTarget(), directUpstreamPeerTarget()} {
|
||||
t.Run(tc.name+"/"+string(target.TargetType), func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupPrivateClusterTest(t, boolPtr(true))
|
||||
connectTestProxy(t, mgr.capabilities.(*proxymanager.Manager), "proxy-2", tc.secondCaps)
|
||||
|
||||
svc := newTestService("app.test.netbird.io")
|
||||
svc.Targets = []*rpservice.Target{target}
|
||||
|
||||
_, err := mgr.CreateService(ctx, testAccountID, testUserID, svc)
|
||||
require.Error(t, err, "a cluster with a non-private proxy must not accept the target")
|
||||
assert.Contains(t, err.Error(), "requires a proxy cluster with private mode enabled")
|
||||
|
||||
services, err := testStore.GetAccountServices(ctx, store.LockingStrengthNone, testAccountID)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, services, "a rejected service must not be persisted")
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateService_RegularTargetIgnoresPrivateCapability(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, _ := setupPrivateClusterTest(t, boolPtr(false))
|
||||
|
||||
_, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
|
||||
require.NoError(t, err, "a peer target without direct upstream must not need a private cluster")
|
||||
}
|
||||
|
||||
func TestUpdateService_PrivateClusterTargets(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
target *rpservice.Target
|
||||
wantErr string
|
||||
}{
|
||||
{name: "switch to cluster target", target: clusterTarget(), wantErr: `target_type "cluster" requires a proxy cluster with private mode enabled`},
|
||||
{name: "enable direct upstream", target: directUpstreamPeerTarget(), wantErr: "direct_upstream requires a proxy cluster with private mode enabled"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupPrivateClusterTest(t, boolPtr(false))
|
||||
|
||||
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
|
||||
require.NoError(t, err)
|
||||
|
||||
updated := newTestService("app.test.netbird.io")
|
||||
updated.ID = created.ID
|
||||
updated.AccountID = testAccountID
|
||||
updated.Targets = []*rpservice.Target{tc.target}
|
||||
|
||||
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tc.wantErr)
|
||||
|
||||
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, stored.Targets, 1)
|
||||
assert.Equal(t, rpservice.TargetTypePeer, stored.Targets[0].TargetType, "the stored target must be unchanged")
|
||||
assert.False(t, stored.Targets[0].Options.DirectUpstream, "the stored target must keep direct upstream disabled")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateService_PrivateClusterAllowsClusterTarget(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr, testStore := setupPrivateClusterTest(t, boolPtr(true))
|
||||
|
||||
created, err := mgr.CreateService(ctx, testAccountID, testUserID, newTestService("app.test.netbird.io"))
|
||||
require.NoError(t, err)
|
||||
|
||||
updated := newTestService("app.test.netbird.io")
|
||||
updated.ID = created.ID
|
||||
updated.AccountID = testAccountID
|
||||
updated.Targets = []*rpservice.Target{clusterTarget()}
|
||||
|
||||
_, err = mgr.UpdateService(ctx, testAccountID, testUserID, updated)
|
||||
require.NoError(t, err)
|
||||
|
||||
stored, err := testStore.GetServiceByID(ctx, store.LockingStrengthNone, testAccountID, created.ID)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, stored.Targets, 1)
|
||||
assert.Equal(t, rpservice.TargetTypeCluster, stored.Targets[0].TargetType, "the cluster target should be stored")
|
||||
}
|
||||
|
||||
func TestValidatePrivateClusterTargets_NoLookupWithoutPrivateTargets(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
// No ClusterAllProxiesPrivate expectation: a lookup would fail the test.
|
||||
mgr := &Manager{capabilities: proxy.NewMockManager(ctrl)}
|
||||
|
||||
targets := []*rpservice.Target{{TargetId: testPeerID, TargetType: rpservice.TargetTypePeer}}
|
||||
require.NoError(t, mgr.validatePrivateClusterTargets(context.Background(), targets, testCluster))
|
||||
}
|
||||
@@ -3,15 +3,16 @@ package manager
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"slices"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/zones"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
"github.com/netbirdio/netbird/management/server/affectedpeers"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
@@ -69,6 +70,9 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
|
||||
}
|
||||
|
||||
zone = zones.NewZone(accountID, zone.Name, zone.Domain, zone.Enabled, zone.EnableSearchDomain, zone.DistributionGroups)
|
||||
var snap *affectedpeers.Snapshot
|
||||
change := affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
existingZone, err := transaction.GetZoneByDomain(ctx, accountID, zone.Domain)
|
||||
if err != nil {
|
||||
@@ -88,7 +92,15 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
|
||||
}
|
||||
|
||||
if err = transaction.CreateZone(ctx, zone); err != nil {
|
||||
return fmt.Errorf("failed to create zone: %w", err)
|
||||
return fmt.Errorf("create zone: %w", err)
|
||||
}
|
||||
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
|
||||
return fmt.Errorf("increment network serial: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -99,6 +111,8 @@ func (m *managerImpl) CreateZone(ctx context.Context, accountID, userID string,
|
||||
|
||||
m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneCreated, zone.EventMeta())
|
||||
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return zone, nil
|
||||
}
|
||||
|
||||
@@ -111,21 +125,26 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
|
||||
return nil, status.NewPermissionDeniedError()
|
||||
}
|
||||
|
||||
zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get zone: %w", err)
|
||||
}
|
||||
|
||||
if zone.Domain != updatedZone.Domain {
|
||||
return nil, status.Errorf(status.InvalidArgument, "zone domain cannot be updated")
|
||||
}
|
||||
|
||||
zone.Name = updatedZone.Name
|
||||
zone.Enabled = updatedZone.Enabled
|
||||
zone.EnableSearchDomain = updatedZone.EnableSearchDomain
|
||||
zone.DistributionGroups = updatedZone.DistributionGroups
|
||||
var zone *zones.Zone
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, updatedZone.ID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get zone: %w", err)
|
||||
}
|
||||
|
||||
if zone.Domain != updatedZone.Domain {
|
||||
return status.Errorf(status.InvalidArgument, "zone domain cannot be updated")
|
||||
}
|
||||
|
||||
oldGroups := zone.DistributionGroups
|
||||
zone.Name = updatedZone.Name
|
||||
zone.Enabled = updatedZone.Enabled
|
||||
zone.EnableSearchDomain = updatedZone.EnableSearchDomain
|
||||
zone.DistributionGroups = updatedZone.DistributionGroups
|
||||
|
||||
for _, groupID := range zone.DistributionGroups {
|
||||
_, err = transaction.GetGroupByID(ctx, store.LockingStrengthNone, accountID, groupID)
|
||||
if err != nil {
|
||||
@@ -134,7 +153,16 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
|
||||
}
|
||||
|
||||
if err = transaction.UpdateZone(ctx, zone); err != nil {
|
||||
return fmt.Errorf("failed to update zone: %w", err)
|
||||
return fmt.Errorf("update zone: %w", err)
|
||||
}
|
||||
|
||||
change = affectedpeers.Change{DistributionGroupIDs: slices.Concat(zone.DistributionGroups, oldGroups)}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
|
||||
return fmt.Errorf("increment network serial: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -145,7 +173,7 @@ func (m *managerImpl) UpdateZone(ctx context.Context, accountID, userID string,
|
||||
|
||||
m.accountManager.StoreEvent(ctx, userID, zone.ID, accountID, activity.DNSZoneUpdated, zone.EventMeta())
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationUpdate})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return zone, nil
|
||||
}
|
||||
@@ -159,13 +187,23 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID
|
||||
return status.NewPermissionDeniedError()
|
||||
}
|
||||
|
||||
zone, err := m.store.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get zone: %w", err)
|
||||
}
|
||||
|
||||
var zone *zones.Zone
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
var eventsToStore []func()
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get zone: %w", err)
|
||||
}
|
||||
|
||||
// Load before delete: the post-delete state no longer references the groups.
|
||||
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
records, err := transaction.GetZoneDNSRecords(ctx, store.LockingStrengthNone, accountID, zoneID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get records: %w", err)
|
||||
@@ -207,7 +245,7 @@ func (m *managerImpl) DeleteZone(ctx context.Context, accountID, userID, zoneID
|
||||
event()
|
||||
}
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZone, Operation: types.UpdateOperationDelete})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -9,11 +9,11 @@ import (
|
||||
"github.com/netbirdio/netbird/management/internals/modules/zones/records"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
"github.com/netbirdio/netbird/management/server/affectedpeers"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
@@ -65,6 +65,8 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
}
|
||||
|
||||
var zone *zones.Zone
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
|
||||
record = records.NewRecord(accountID, zoneID, record.Name, record.Type, record.Content, record.TTL)
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
@@ -82,6 +84,11 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
return fmt.Errorf("failed to create dns record: %w", err)
|
||||
}
|
||||
|
||||
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
err = transaction.IncrementNetworkSerial(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to increment network serial: %w", err)
|
||||
@@ -96,7 +103,7 @@ func (m *managerImpl) CreateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
meta := record.EventMeta(zone.ID, zone.Name)
|
||||
m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordCreated, meta)
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationCreate})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return record, nil
|
||||
}
|
||||
@@ -112,6 +119,8 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
|
||||
var zone *zones.Zone
|
||||
var record *records.Record
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
|
||||
@@ -141,6 +150,11 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
return fmt.Errorf("failed to update dns record: %w", err)
|
||||
}
|
||||
|
||||
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
err = transaction.IncrementNetworkSerial(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to increment network serial: %w", err)
|
||||
@@ -155,7 +169,7 @@ func (m *managerImpl) UpdateRecord(ctx context.Context, accountID, userID, zoneI
|
||||
meta := record.EventMeta(zone.ID, zone.Name)
|
||||
m.accountManager.StoreEvent(ctx, userID, record.ID, accountID, activity.DNSRecordUpdated, meta)
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationUpdate})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return record, nil
|
||||
}
|
||||
@@ -171,6 +185,8 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
|
||||
|
||||
var record *records.Record
|
||||
var zone *zones.Zone
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
|
||||
err = m.store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
zone, err = transaction.GetZoneByID(ctx, store.LockingStrengthUpdate, accountID, zoneID)
|
||||
@@ -188,6 +204,11 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
|
||||
return fmt.Errorf("failed to delete dns record: %w", err)
|
||||
}
|
||||
|
||||
change = affectedpeers.Change{DistributionGroupIDs: zone.DistributionGroups}
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
|
||||
err = transaction.IncrementNetworkSerial(ctx, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to increment network serial: %w", err)
|
||||
@@ -202,7 +223,7 @@ func (m *managerImpl) DeleteRecord(ctx context.Context, accountID, userID, zoneI
|
||||
meta := record.EventMeta(zone.ID, zone.Name)
|
||||
m.accountManager.StoreEvent(ctx, userID, recordID, accountID, activity.DNSRecordDeleted, meta)
|
||||
|
||||
go m.accountManager.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceZoneRecord, Operation: types.UpdateOperationDelete})
|
||||
m.accountManager.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -334,6 +334,9 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco
|
||||
var groupChangesAffectPeers bool
|
||||
var reloadReverseProxy bool
|
||||
var effectiveOldNetworkRange netip.Prefix
|
||||
var ipv6Changed bool
|
||||
var ipv6Snap *affectedpeers.Snapshot
|
||||
var ipv6Change affectedpeers.Change
|
||||
|
||||
err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
var groupsUpdated bool
|
||||
@@ -379,10 +382,10 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco
|
||||
}
|
||||
|
||||
if ipv6SettingsChanged(oldSettings, newSettings) {
|
||||
if err = am.updatePeerIPv6Addresses(ctx, transaction, accountID, newSettings); err != nil {
|
||||
if ipv6Change, err = am.applyIPv6SettingsChange(ctx, transaction, accountID, oldSettings, newSettings); err != nil {
|
||||
return err
|
||||
}
|
||||
updateAccountPeers = true
|
||||
ipv6Changed = true
|
||||
}
|
||||
|
||||
if oldSettings.RoutingPeerDNSResolutionEnabled != newSettings.RoutingPeerDNSResolutionEnabled ||
|
||||
@@ -419,12 +422,20 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco
|
||||
return err
|
||||
}
|
||||
|
||||
if updateAccountPeers || groupsUpdated {
|
||||
if updateAccountPeers || groupsUpdated || ipv6Changed {
|
||||
if err = transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// A full account refresh already covers the IPv6 change, so the affected-peers
|
||||
// snapshot is only needed when nothing account-wide changed.
|
||||
if ipv6Changed && !updateAccountPeers && !groupChangesAffectPeers {
|
||||
if ipv6Snap, err = affectedpeers.Load(ctx, transaction, accountID, ipv6Change); err != nil {
|
||||
return fmt.Errorf("load affected peers: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
@@ -486,13 +497,34 @@ func (am *DefaultAccountManager) UpdateAccountSettings(ctx context.Context, acco
|
||||
}
|
||||
}
|
||||
|
||||
if updateAccountPeers || extraSettingsChanged || groupChangesAffectPeers {
|
||||
go am.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceAccountSettings, Operation: types.UpdateOperationUpdate})
|
||||
switch {
|
||||
case updateAccountPeers || extraSettingsChanged || groupChangesAffectPeers:
|
||||
go am.UpdateAccountPeers(context.WithoutCancel(ctx), accountID, types.UpdateReason{Resource: types.UpdateResourceAccountSettings, Operation: types.UpdateOperationUpdate})
|
||||
case ipv6Snap != nil:
|
||||
am.ExpandAndUpdateAffected(ctx, accountID, ipv6Snap, ipv6Change)
|
||||
}
|
||||
|
||||
return newSettings, nil
|
||||
}
|
||||
|
||||
// applyIPv6SettingsChange reconciles peer IPv6 addresses for new IPv6 settings and
|
||||
// returns the affected-peers change: peers whose address changed refresh together
|
||||
// with every peer that reaches them. On a range change every peer holding an address
|
||||
// also refreshes itself, since its interface prefix comes from the account range even
|
||||
// when its address stays inside the new one.
|
||||
func (am *DefaultAccountManager) applyIPv6SettingsChange(ctx context.Context, transaction store.Store, accountID string, oldSettings, newSettings *types.Settings) (affectedpeers.Change, error) {
|
||||
result, err := am.updatePeerIPv6Addresses(ctx, transaction, accountID, newSettings)
|
||||
if err != nil {
|
||||
return affectedpeers.Change{}, err
|
||||
}
|
||||
|
||||
change := affectedpeers.Change{ChangedPeerIDs: result.changed}
|
||||
if oldSettings.NetworkRangeV6 != newSettings.NetworkRangeV6 {
|
||||
change.OutputPeerIDs = result.withIPv6
|
||||
}
|
||||
return change, nil
|
||||
}
|
||||
|
||||
func ipv6SettingsChanged(old, updated *types.Settings) bool {
|
||||
if old.NetworkRangeV6 != updated.NetworkRangeV6 {
|
||||
return true
|
||||
@@ -1742,9 +1774,11 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth
|
||||
|
||||
change.LinkGroups = allGroupChanges
|
||||
|
||||
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges); err != nil {
|
||||
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges)
|
||||
if err != nil {
|
||||
return fmt.Errorf("reconcile IPv6 for group changes: %w", err)
|
||||
}
|
||||
change.ChangedPeerIDs = append(change.ChangedPeerIDs, ipv6Changed...)
|
||||
|
||||
if err = transaction.IncrementNetworkSerial(ctx, userAuth.AccountId); err != nil {
|
||||
return fmt.Errorf("error incrementing network serial: %w", err)
|
||||
@@ -2334,7 +2368,8 @@ func (am *DefaultAccountManager) propagateUserGroupMemberships(ctx context.Conte
|
||||
return false, false, err
|
||||
}
|
||||
|
||||
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, updatedGroups); err != nil {
|
||||
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, updatedGroups)
|
||||
if err != nil {
|
||||
return false, false, fmt.Errorf("reconcile IPv6 for group changes: %w", err)
|
||||
}
|
||||
|
||||
@@ -2343,7 +2378,7 @@ func (am *DefaultAccountManager) propagateUserGroupMemberships(ctx context.Conte
|
||||
return false, false, fmt.Errorf("error checking if group changes affect peers: %w", err)
|
||||
}
|
||||
|
||||
return len(updatedGroups) > 0, peersAffected, nil
|
||||
return len(updatedGroups) > 0, peersAffected || len(ipv6Changed) > 0, nil
|
||||
}
|
||||
|
||||
// propagateAutoGroupsForUsers adds each user's peers to their AutoGroups where not already present.
|
||||
@@ -2391,7 +2426,7 @@ func (am *DefaultAccountManager) reallocateAccountPeerIPs(ctx context.Context, t
|
||||
return err
|
||||
}
|
||||
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "")
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -2428,7 +2463,7 @@ func (am *DefaultAccountManager) reallocateAccountPeerIPs(ctx context.Context, t
|
||||
// v6 address get one allocated. When disabled, all v6 addresses are cleared.
|
||||
// When the v6 range changes, all v6 addresses are reallocated.
|
||||
func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transaction store.Store, accountID, peerID string, newIPv6 netip.Addr) error {
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "")
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "", "")
|
||||
if err != nil {
|
||||
return fmt.Errorf("get peers: %w", err)
|
||||
}
|
||||
@@ -2440,56 +2475,78 @@ func (am *DefaultAccountManager) checkIPv6Collision(ctx context.Context, transac
|
||||
return nil
|
||||
}
|
||||
|
||||
func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) error {
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "")
|
||||
// ipv6Reassignment reports the outcome of an IPv6 address reconciliation.
|
||||
type ipv6Reassignment struct {
|
||||
// changed are the peers whose IPv6 address was assigned, removed or reallocated.
|
||||
changed []string
|
||||
// withIPv6 are all peers holding an IPv6 address after the reconciliation.
|
||||
withIPv6 []string
|
||||
}
|
||||
|
||||
func (am *DefaultAccountManager) updatePeerIPv6Addresses(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings) (ipv6Reassignment, error) {
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthUpdate, accountID, "", "", "")
|
||||
if err != nil {
|
||||
return fmt.Errorf("get peers: %w", err)
|
||||
return ipv6Reassignment{}, fmt.Errorf("get peers: %w", err)
|
||||
}
|
||||
|
||||
network, err := transaction.GetAccountNetwork(ctx, store.LockingStrengthUpdate, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get network: %w", err)
|
||||
return ipv6Reassignment{}, fmt.Errorf("get network: %w", err)
|
||||
}
|
||||
|
||||
if err := am.ensureIPv6Subnet(ctx, transaction, accountID, settings, network); err != nil {
|
||||
return err
|
||||
return ipv6Reassignment{}, err
|
||||
}
|
||||
|
||||
allowedPeers, err := am.buildIPv6AllowedPeers(ctx, transaction, accountID, settings)
|
||||
if err != nil {
|
||||
return err
|
||||
return ipv6Reassignment{}, err
|
||||
}
|
||||
|
||||
v6Prefix, err := netip.ParsePrefix(network.NetV6.String())
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse IPv6 prefix: %w", err)
|
||||
return ipv6Reassignment{}, fmt.Errorf("parse IPv6 prefix: %w", err)
|
||||
}
|
||||
|
||||
if err := am.assignPeerIPv6Addresses(ctx, transaction, accountID, peers, network, allowedPeers, v6Prefix); err != nil {
|
||||
return err
|
||||
changed, err := am.assignPeerIPv6Addresses(ctx, transaction, accountID, peers, network, allowedPeers, v6Prefix)
|
||||
if err != nil {
|
||||
return ipv6Reassignment{}, err
|
||||
}
|
||||
|
||||
log.WithContext(ctx).Infof("updated IPv6 addresses for %d peers in account %s (groups=%d)",
|
||||
len(peers), accountID, len(settings.IPv6EnabledGroups))
|
||||
result := ipv6Reassignment{changed: changed}
|
||||
for _, peer := range peers {
|
||||
if peer.IPv6.IsValid() {
|
||||
result.withIPv6 = append(result.withIPv6, peer.ID)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
log.WithContext(ctx).Infof("updated IPv6 addresses for %d of %d peers in account %s (groups=%d)",
|
||||
len(changed), len(peers), accountID, len(settings.IPv6EnabledGroups))
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// reconcileIPv6ForGroupChanges checks whether the given group IDs overlap with
|
||||
// the account's IPv6EnabledGroups. If they do, it runs a full IPv6 address
|
||||
// reconciliation so that peers gaining or losing membership in an IPv6-enabled
|
||||
// group get their addresses assigned or removed.
|
||||
func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) error {
|
||||
// group get their addresses assigned or removed. It returns the peers whose IPv6
|
||||
// address changed, which callers pass as changed peers so every peer that can
|
||||
// reach them refreshes.
|
||||
func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) ([]string, error) {
|
||||
settings, err := transaction.GetAccountSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get account settings: %w", err)
|
||||
return nil, fmt.Errorf("get account settings: %w", err)
|
||||
}
|
||||
|
||||
if !ipv6ReconcileNeeded(settings, groupIDs) {
|
||||
return nil
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
return am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings)
|
||||
result, err := am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return result.changed, nil
|
||||
}
|
||||
|
||||
// ipv6ReconcileNeeded reports whether changes to the given groups trigger an IPv6
|
||||
@@ -2528,7 +2585,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses(
|
||||
ctx context.Context, transaction store.Store, accountID string,
|
||||
peers []*nbpeer.Peer, network *types.Network,
|
||||
allowedPeers map[string]struct{}, v6Prefix netip.Prefix,
|
||||
) error {
|
||||
) ([]string, error) {
|
||||
takenV6 := make(map[netip.Addr]struct{})
|
||||
for _, peer := range peers {
|
||||
if _, ok := allowedPeers[peer.ID]; ok && peer.IPv6.IsValid() && network.NetV6.Contains(peer.IPv6.AsSlice()) {
|
||||
@@ -2536,6 +2593,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses(
|
||||
}
|
||||
}
|
||||
|
||||
var changed []string
|
||||
for _, peer := range peers {
|
||||
_, allowed := allowedPeers[peer.ID]
|
||||
oldIPv6 := peer.IPv6
|
||||
@@ -2545,7 +2603,7 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses(
|
||||
} else if !peer.IPv6.IsValid() || !network.NetV6.Contains(peer.IPv6.AsSlice()) {
|
||||
newIP, err := allocateIPv6WithRetry(v6Prefix, takenV6, peer.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
peer.IPv6 = newIP
|
||||
}
|
||||
@@ -2555,10 +2613,11 @@ func (am *DefaultAccountManager) assignPeerIPv6Addresses(
|
||||
}
|
||||
|
||||
if err := transaction.SavePeer(ctx, accountID, peer); err != nil {
|
||||
return fmt.Errorf("save peer %s: %w", peer.ID, err)
|
||||
return nil, fmt.Errorf("save peer %s: %w", peer.ID, err)
|
||||
}
|
||||
changed = append(changed, peer.ID)
|
||||
}
|
||||
return nil
|
||||
return changed, nil
|
||||
}
|
||||
|
||||
func allocateIPv6WithRetry(prefix netip.Prefix, taken map[netip.Addr]struct{}, peerID string) (netip.Addr, error) {
|
||||
@@ -2602,7 +2661,7 @@ func (am *DefaultAccountManager) buildIPv6AllowedPeers(ctx context.Context, tran
|
||||
|
||||
// Embedded proxy peers sit outside regular group membership but must
|
||||
// participate in any v6-enabled overlay to reach v6-only peers.
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get peers: %w", err)
|
||||
}
|
||||
@@ -2673,7 +2732,7 @@ func (am *DefaultAccountManager) updatePeerIPInTransaction(ctx context.Context,
|
||||
return nil
|
||||
}
|
||||
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "")
|
||||
peers, err := transaction.GetAccountPeers(ctx, store.LockingStrengthShare, accountID, "", "", "")
|
||||
if err != nil {
|
||||
return fmt.Errorf("get account peers: %w", err)
|
||||
}
|
||||
|
||||
@@ -62,7 +62,7 @@ type Manager interface {
|
||||
GetUserByID(ctx context.Context, id string) (*types.User, error)
|
||||
GetUserFromUserAuth(ctx context.Context, userAuth auth.UserAuth) (*types.User, error)
|
||||
ListUsers(ctx context.Context, accountID string) ([]*types.User, error)
|
||||
GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error)
|
||||
GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error)
|
||||
MarkPeerConnected(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error
|
||||
MarkPeerDisconnected(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error
|
||||
DeletePeer(ctx context.Context, accountID, peerID, userID string) error
|
||||
|
||||
@@ -982,18 +982,18 @@ func (mr *MockManagerMockRecorder) GetPeerNetwork(ctx, peerID any) *gomock.Call
|
||||
}
|
||||
|
||||
// GetPeers mocks base method.
|
||||
func (m *MockManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*peer.Peer, error) {
|
||||
func (m *MockManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*peer.Peer, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetPeers", ctx, accountID, userID, nameFilter, ipFilter)
|
||||
ret := m.ctrl.Call(m, "GetPeers", ctx, accountID, userID, nameFilter, ipFilter, macFilter)
|
||||
ret0, _ := ret[0].([]*peer.Peer)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetPeers indicates an expected call of GetPeers.
|
||||
func (mr *MockManagerMockRecorder) GetPeers(ctx, accountID, userID, nameFilter, ipFilter any) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeers(ctx, accountID, userID, nameFilter, ipFilter, macFilter any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeers", reflect.TypeOf((*MockManager)(nil).GetPeers), ctx, accountID, userID, nameFilter, ipFilter)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeers", reflect.TypeOf((*MockManager)(nil).GetPeers), ctx, accountID, userID, nameFilter, ipFilter, macFilter)
|
||||
}
|
||||
|
||||
// GetPolicy mocks base method.
|
||||
|
||||
@@ -2557,7 +2557,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_PeerApproval(t *testing.T)
|
||||
_, err = manager.UpdateAccountSettings(ctx, accountID, userID, newSettings)
|
||||
require.NoError(t, err)
|
||||
|
||||
accountPeers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
accountPeers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, peer := range accountPeers {
|
||||
@@ -4557,7 +4557,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "")
|
||||
peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "")
|
||||
require.NoError(t, err)
|
||||
require.Len(t, peers, len(before))
|
||||
for _, p := range peers {
|
||||
@@ -4575,7 +4575,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "")
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "")
|
||||
require.NoError(t, err)
|
||||
for _, p := range peers {
|
||||
assert.Equal(t, before[p.ID], p.IP, "peer %s IP should not change for host-bit-set equivalent range", p.ID)
|
||||
@@ -4589,7 +4589,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "")
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "")
|
||||
require.NoError(t, err)
|
||||
for _, p := range peers {
|
||||
assert.Equal(t, before[p.ID], p.IP, "peer %s IP should not change when NetworkRange omitted", p.ID)
|
||||
@@ -4605,7 +4605,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_NetworkRangePreserved(t *te
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "")
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, account.Id, "", "", "")
|
||||
require.NoError(t, err)
|
||||
for _, p := range peers {
|
||||
assert.True(t, newRange.Contains(p.IP), "peer %s should be in new range %s, got %s", p.ID, newRange, p.IP)
|
||||
@@ -4623,7 +4623,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, settings.IPv6EnabledGroups, "new account should have IPv6 enabled for All group")
|
||||
|
||||
peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
peers, err := manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
require.NoError(t, err)
|
||||
for _, p := range peers {
|
||||
assert.True(t, p.IPv6.IsValid(), "peer %s should have IPv6 with All group enabled", p.ID)
|
||||
@@ -4651,7 +4651,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin
|
||||
assert.Equal(t, []string{partialGroup.ID}, updatedSettings.IPv6EnabledGroups)
|
||||
|
||||
// peer1 and peer2 should have IPv6; peer3 should not.
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
require.NoError(t, err)
|
||||
peerMap := make(map[string]*nbpeer.Peer, len(peers))
|
||||
for _, p := range peers {
|
||||
@@ -4671,7 +4671,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, updatedSettings.IPv6EnabledGroups)
|
||||
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
require.NoError(t, err)
|
||||
for _, p := range peers {
|
||||
assert.False(t, p.IPv6.IsValid(), "peer %s should have no IPv6 when groups cleared", p.ID)
|
||||
@@ -4686,7 +4686,7 @@ func TestDefaultAccountManager_UpdateAccountSettings_IPv6EnabledGroups(t *testin
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
peers, err = manager.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
require.NoError(t, err)
|
||||
peerMap = make(map[string]*nbpeer.Peer, len(peers))
|
||||
for _, p := range peers {
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
const (
|
||||
ipv6GroupA = "ipv6-grp-a"
|
||||
ipv6GroupB = "ipv6-grp-b"
|
||||
ipv6GroupC = "ipv6-grp-c"
|
||||
ipv6GroupD = "ipv6-grp-d"
|
||||
)
|
||||
|
||||
// ipv6AffectedTest holds three peers: peer1 in group A, peer2 in group B, peer3 in
|
||||
// group C, with a single A<->B policy. peer3 is unrelated to peer1 and peer2. Group D
|
||||
// is empty and referenced by nothing.
|
||||
type ipv6AffectedTest struct {
|
||||
manager *DefaultAccountManager
|
||||
accountID string
|
||||
peer1, peer2, peer3 *nbpeer.Peer
|
||||
updMsg1, updMsg2, updMsg3 <-chan *network_map.UpdateMessage
|
||||
}
|
||||
|
||||
func setupIPv6AffectedTest(t *testing.T, ipv6Groups []string) *ipv6AffectedTest {
|
||||
t.Helper()
|
||||
|
||||
manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t)
|
||||
ctx := context.Background()
|
||||
accountID := account.Id
|
||||
|
||||
policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID)
|
||||
require.NoError(t, err)
|
||||
for _, p := range policies {
|
||||
require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID))
|
||||
}
|
||||
|
||||
for _, g := range []*types.Group{
|
||||
{ID: ipv6GroupA, Name: "IPv6-A", Peers: []string{peer1.ID}},
|
||||
{ID: ipv6GroupB, Name: "IPv6-B", Peers: []string{peer2.ID}},
|
||||
{ID: ipv6GroupC, Name: "IPv6-C", Peers: []string{peer3.ID}},
|
||||
{ID: ipv6GroupD, Name: "IPv6-D"},
|
||||
} {
|
||||
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, g))
|
||||
}
|
||||
|
||||
_, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{
|
||||
Enabled: true,
|
||||
Rules: []*types.PolicyRule{{
|
||||
Enabled: true,
|
||||
Sources: []string{ipv6GroupA},
|
||||
Destinations: []string{ipv6GroupB},
|
||||
Bidirectional: true,
|
||||
Action: types.PolicyTrafficActionAccept,
|
||||
}},
|
||||
}, true)
|
||||
require.NoError(t, err)
|
||||
|
||||
// New accounts enable IPv6 for the All group; start from the requested groups.
|
||||
updateIPv6TestSettings(t, manager, accountID, func(s *types.Settings) {
|
||||
s.IPv6EnabledGroups = ipv6Groups
|
||||
})
|
||||
|
||||
tc := &ipv6AffectedTest{
|
||||
manager: manager,
|
||||
accountID: accountID,
|
||||
peer1: peer1,
|
||||
peer2: peer2,
|
||||
peer3: peer3,
|
||||
}
|
||||
tc.updMsg1 = updateManager.CreateChannel(ctx, peer1.ID)
|
||||
tc.updMsg2 = updateManager.CreateChannel(ctx, peer2.ID)
|
||||
tc.updMsg3 = updateManager.CreateChannel(ctx, peer3.ID)
|
||||
t.Cleanup(func() {
|
||||
updateManager.CloseChannel(ctx, peer1.ID)
|
||||
updateManager.CloseChannel(ctx, peer2.ID)
|
||||
updateManager.CloseChannel(ctx, peer3.ID)
|
||||
})
|
||||
|
||||
// The setup changes above dispatch asynchronously and can land after the
|
||||
// channels open, so drop them before the test acts.
|
||||
drainPeerUpdates(tc.updMsg1)
|
||||
drainPeerUpdates(tc.updMsg2)
|
||||
drainPeerUpdates(tc.updMsg3)
|
||||
|
||||
return tc
|
||||
}
|
||||
|
||||
// updateIPv6TestSettings applies mutate to a copy of the current settings, so only
|
||||
// the mutated fields differ from what is stored.
|
||||
func updateIPv6TestSettings(t *testing.T, manager *DefaultAccountManager, accountID string, mutate func(*types.Settings)) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
|
||||
current, err := manager.Store.GetAccountSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
require.NoError(t, err)
|
||||
|
||||
updated := current.Copy()
|
||||
mutate(updated)
|
||||
|
||||
_, err = manager.UpdateAccountSettings(ctx, accountID, userID, updated)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func (tc *ipv6AffectedTest) peerIPv6(t *testing.T, peerID string) netip.Addr {
|
||||
t.Helper()
|
||||
peer, err := tc.manager.Store.GetPeerByID(context.Background(), store.LockingStrengthNone, tc.accountID, peerID)
|
||||
require.NoError(t, err)
|
||||
return peer.IPv6
|
||||
}
|
||||
|
||||
func TestAffectedPeers_IPv6GroupEnabled_RefreshesOnlyReachablePeers(t *testing.T) {
|
||||
tc := setupIPv6AffectedTest(t, nil)
|
||||
|
||||
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
|
||||
s.IPv6EnabledGroups = []string{ipv6GroupA}
|
||||
})
|
||||
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address")
|
||||
|
||||
peerShouldReceiveUpdate(t, tc.updMsg1)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg2)
|
||||
peerShouldNotReceiveUpdate(t, tc.updMsg3)
|
||||
}
|
||||
|
||||
func TestAffectedPeers_IPv6GroupDisabled_RefreshesOnlyReachablePeers(t *testing.T) {
|
||||
tc := setupIPv6AffectedTest(t, []string{ipv6GroupA})
|
||||
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should start with an IPv6 address")
|
||||
|
||||
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
|
||||
s.IPv6EnabledGroups = []string{}
|
||||
})
|
||||
require.False(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should lose its IPv6 address")
|
||||
|
||||
peerShouldReceiveUpdate(t, tc.updMsg1)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg2)
|
||||
peerShouldNotReceiveUpdate(t, tc.updMsg3)
|
||||
}
|
||||
|
||||
// Widening the IPv6 range keeps peer addresses, but each holder's interface prefix
|
||||
// comes from the range, so holders refresh while peers that only reach them do not.
|
||||
func TestAffectedPeers_IPv6RangeWidened_RefreshesAddressHolders(t *testing.T) {
|
||||
tc := setupIPv6AffectedTest(t, []string{ipv6GroupA})
|
||||
oldIPv6 := tc.peerIPv6(t, tc.peer1.ID)
|
||||
require.True(t, oldIPv6.IsValid(), "peer1 should start with an IPv6 address")
|
||||
|
||||
// The range is allocated on the account network; settings may leave it empty.
|
||||
network, err := tc.manager.Store.GetAccountNetwork(context.Background(), store.LockingStrengthNone, tc.accountID)
|
||||
require.NoError(t, err)
|
||||
current := prefixFromIPNet(network.NetV6)
|
||||
require.True(t, current.IsValid(), "account should have an IPv6 range")
|
||||
widened := netip.PrefixFrom(current.Addr(), current.Bits()-8).Masked()
|
||||
|
||||
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
|
||||
s.NetworkRangeV6 = widened
|
||||
})
|
||||
require.Equal(t, oldIPv6, tc.peerIPv6(t, tc.peer1.ID), "peer1 should keep its address inside the widened range")
|
||||
|
||||
peerShouldReceiveUpdate(t, tc.updMsg1)
|
||||
peerShouldNotReceiveUpdate(t, tc.updMsg2)
|
||||
peerShouldNotReceiveUpdate(t, tc.updMsg3)
|
||||
}
|
||||
|
||||
func TestAffectedPeers_IPv4RangeChange_RefreshesWholeAccount(t *testing.T) {
|
||||
tc := setupIPv6AffectedTest(t, nil)
|
||||
|
||||
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
|
||||
s.NetworkRange = netip.MustParsePrefix("100.70.0.0/16")
|
||||
})
|
||||
|
||||
peerShouldReceiveUpdate(t, tc.updMsg1)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg2)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg3)
|
||||
}
|
||||
|
||||
func TestAffectedPeers_IPv6WithAccountWideChange_RefreshesWholeAccount(t *testing.T) {
|
||||
tc := setupIPv6AffectedTest(t, nil)
|
||||
|
||||
updateIPv6TestSettings(t, tc.manager, tc.accountID, func(s *types.Settings) {
|
||||
s.IPv6EnabledGroups = []string{ipv6GroupA}
|
||||
s.LazyConnectionEnabled = !s.LazyConnectionEnabled
|
||||
})
|
||||
|
||||
peerShouldReceiveUpdate(t, tc.updMsg1)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg2)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg3)
|
||||
}
|
||||
|
||||
// Joining an IPv6-enabled group that no policy references gives peer1 an address.
|
||||
// peer2 reaches peer1 through group A, not through the joined group, and must still
|
||||
// learn the new address.
|
||||
func TestAffectedPeers_GroupAddPeerIPv6_RefreshesPeersReachingThroughOtherGroups(t *testing.T) {
|
||||
tc := setupIPv6AffectedTest(t, []string{ipv6GroupD})
|
||||
|
||||
require.NoError(t, tc.manager.GroupAddPeer(context.Background(), tc.accountID, ipv6GroupD, tc.peer1.ID))
|
||||
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address")
|
||||
|
||||
peerShouldReceiveUpdate(t, tc.updMsg1)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg2)
|
||||
peerShouldNotReceiveUpdate(t, tc.updMsg3)
|
||||
}
|
||||
|
||||
func TestAffectedPeers_UpdateGroupIPv6_RefreshesPeersReachingThroughOtherGroups(t *testing.T) {
|
||||
tc := setupIPv6AffectedTest(t, []string{ipv6GroupD})
|
||||
|
||||
require.NoError(t, tc.manager.UpdateGroup(context.Background(), tc.accountID, userID, &types.Group{
|
||||
ID: ipv6GroupD,
|
||||
Name: "IPv6-D",
|
||||
Peers: []string{tc.peer1.ID},
|
||||
}))
|
||||
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address")
|
||||
|
||||
peerShouldReceiveUpdate(t, tc.updMsg1)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg2)
|
||||
peerShouldNotReceiveUpdate(t, tc.updMsg3)
|
||||
}
|
||||
|
||||
// Deleting an IPv6-enabled group removes its members' addresses after the
|
||||
// pre-delete snapshot was taken.
|
||||
func TestAffectedPeers_DeleteIPv6Group_RefreshesFormerMembersAndReachablePeers(t *testing.T) {
|
||||
tc := setupIPv6AffectedTest(t, []string{ipv6GroupD})
|
||||
ctx := context.Background()
|
||||
|
||||
require.NoError(t, tc.manager.GroupAddPeer(ctx, tc.accountID, ipv6GroupD, tc.peer1.ID))
|
||||
require.True(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should get an IPv6 address")
|
||||
drainPeerUpdates(tc.updMsg1)
|
||||
drainPeerUpdates(tc.updMsg2)
|
||||
drainPeerUpdates(tc.updMsg3)
|
||||
|
||||
require.NoError(t, tc.manager.DeleteGroup(ctx, tc.accountID, userID, ipv6GroupD))
|
||||
require.False(t, tc.peerIPv6(t, tc.peer1.ID).IsValid(), "peer1 should lose its IPv6 address")
|
||||
|
||||
peerShouldReceiveUpdate(t, tc.updMsg1)
|
||||
peerShouldReceiveUpdate(t, tc.updMsg2)
|
||||
peerShouldNotReceiveUpdate(t, tc.updMsg3)
|
||||
}
|
||||
@@ -108,11 +108,13 @@ func TestAffectedPeers_SaveUser_OnlyAffectedPeersUpdated(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("auto group change reassigning IPv6 refreshes the changed peers and their observers", func(t *testing.T) {
|
||||
account, err := manager.Store.GetAccount(ctx, accountID)
|
||||
require.NoError(t, err)
|
||||
account.Settings.IPv6EnabledGroups = []string{"ug-v6"}
|
||||
require.NoError(t, manager.Store.SaveAccount(ctx, account))
|
||||
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "ug-v6", Name: "ug-v6"}))
|
||||
// Apply through the settings API so the reconciliation that strips the other
|
||||
// peers' addresses happens here, leaving the target as the only peer the
|
||||
// user update reassigns.
|
||||
updateIPv6TestSettings(t, manager, accountID, func(s *types.Settings) {
|
||||
s.IPv6EnabledGroups = []string{"ug-v6"}
|
||||
})
|
||||
|
||||
drainPeerUpdates(updTarget)
|
||||
drainPeerUpdates(upd2)
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/zones"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/zones/records"
|
||||
"github.com/netbirdio/netbird/management/server/affectedpeers"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
const affectedZoneDomain = "zone.test"
|
||||
|
||||
// createAffectedZone stores a zone distributed to the given groups, optionally with
|
||||
// one A record so the network map actually ships it.
|
||||
func createAffectedZone(t *testing.T, s store.Store, accountID, domain string, enabled, withRecord bool, groups []string) *zones.Zone {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
|
||||
zone := zones.NewZone(accountID, domain, domain, enabled, false, groups)
|
||||
require.NoError(t, s.CreateZone(ctx, zone))
|
||||
|
||||
if withRecord {
|
||||
record := records.NewRecord(accountID, zone.ID, "host."+domain, records.RecordTypeA, "10.0.0.1", 300)
|
||||
require.NoError(t, s.CreateDNSRecord(ctx, record))
|
||||
}
|
||||
|
||||
return zone
|
||||
}
|
||||
|
||||
func TestCollectGroupChange_ZoneLinked(t *testing.T) {
|
||||
_, s, accountID, _, groupIDs := setupAffectedPeersTest(t)
|
||||
ctx := context.Background()
|
||||
|
||||
createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]})
|
||||
|
||||
groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0]})
|
||||
assert.Contains(t, groups, groupIDs[0], "group distributed a zone should be affected by its own change")
|
||||
|
||||
groups, _ = collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[1]})
|
||||
assert.Empty(t, groups, "group not referenced by any zone should not be affected")
|
||||
}
|
||||
|
||||
func TestCollectGroupChange_UnshippedZoneNotLinked(t *testing.T) {
|
||||
_, s, accountID, _, groupIDs := setupAffectedPeersTest(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Disabled zone and zone without records are never shipped by the network map.
|
||||
createAffectedZone(t, s, accountID, "disabled."+affectedZoneDomain, false, true, []string{groupIDs[0]})
|
||||
createAffectedZone(t, s, accountID, "empty."+affectedZoneDomain, true, false, []string{groupIDs[1]})
|
||||
|
||||
groups, _ := collectGroupChangeAffectedGroups(ctx, s, accountID, []string{groupIDs[0], groupIDs[1]})
|
||||
assert.Empty(t, groups, "groups referenced only by unshipped zones should not be affected")
|
||||
}
|
||||
|
||||
func TestResolveAffectedPeers_ZoneGroupMembershipChange(t *testing.T) {
|
||||
_, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t)
|
||||
|
||||
createAffectedZone(t, s, accountID, affectedZoneDomain, true, true, []string{groupIDs[0]})
|
||||
|
||||
// Same change shape UpdateGroup builds: the group changed as a whole and peer1
|
||||
// left it, so peer1 must refresh to drop the zone.
|
||||
change := affectedpeers.Change{
|
||||
ChangedGroupIDs: []string{groupIDs[0]},
|
||||
RemovedPeersByGroup: map[string][]string{groupIDs[0]: {peerIDs[1]}},
|
||||
}
|
||||
|
||||
result := resolveAffected(t, s, accountID, change)
|
||||
assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[1]}, result, "current and removed members of the zone group should be affected")
|
||||
}
|
||||
|
||||
func TestResolveAffectedPeers_ZoneDistributionChange(t *testing.T) {
|
||||
_, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t)
|
||||
|
||||
// Zone create/update/delete passes old and new distribution groups.
|
||||
change := affectedpeers.Change{DistributionGroupIDs: []string{groupIDs[0], groupIDs[2]}}
|
||||
|
||||
result := resolveAffected(t, s, accountID, change)
|
||||
assert.ElementsMatch(t, []string{peerIDs[0], peerIDs[2]}, result, "only members of the distribution groups should be affected")
|
||||
}
|
||||
|
||||
// TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone verifies that adding a peer
|
||||
// to a group referenced only by a zone pushes the zone to the new member and leaves
|
||||
// unrelated peers alone.
|
||||
func TestAffectedPeers_ZoneGroupUpdate_NewMemberReceivesZone(t *testing.T) {
|
||||
manager, updateManager, account, peer1, peer2, peer3 := setupNetworkMapTest(t)
|
||||
ctx := context.Background()
|
||||
accountID := account.Id
|
||||
|
||||
policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID)
|
||||
require.NoError(t, err)
|
||||
for _, p := range policies {
|
||||
require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID))
|
||||
}
|
||||
|
||||
zoneGroup := &types.Group{ID: "zone-grp", Name: "ZoneGroup", Peers: []string{peer1.ID}}
|
||||
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, zoneGroup))
|
||||
|
||||
createAffectedZone(t, manager.Store, accountID, affectedZoneDomain, true, true, []string{zoneGroup.ID})
|
||||
|
||||
updMsg1 := updateManager.CreateChannel(ctx, peer1.ID)
|
||||
updMsg2 := updateManager.CreateChannel(ctx, peer2.ID)
|
||||
updMsg3 := updateManager.CreateChannel(ctx, peer3.ID)
|
||||
t.Cleanup(func() {
|
||||
updateManager.CloseChannel(ctx, peer1.ID)
|
||||
updateManager.CloseChannel(ctx, peer2.ID)
|
||||
updateManager.CloseChannel(ctx, peer3.ID)
|
||||
})
|
||||
|
||||
zoneGroup.Peers = []string{peer1.ID, peer2.ID}
|
||||
require.NoError(t, manager.UpdateGroup(ctx, accountID, userID, zoneGroup))
|
||||
|
||||
peerShouldReceiveUpdate(t, updMsg1)
|
||||
msg := receivePeerUpdate(t, updMsg2)
|
||||
assert.True(t, syncHasCustomZone(msg, affectedZoneDomain+"."), "new zone group member should receive the zone")
|
||||
peerShouldNotReceiveUpdate(t, updMsg3)
|
||||
}
|
||||
|
||||
func receivePeerUpdate(t *testing.T, ch <-chan *network_map.UpdateMessage) *network_map.UpdateMessage {
|
||||
t.Helper()
|
||||
select {
|
||||
case msg := <-ch:
|
||||
require.NotNil(t, msg, "update message should not be nil")
|
||||
return msg
|
||||
case <-time.After(peerUpdateTimeout):
|
||||
require.FailNow(t, "timed out waiting for update message")
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func syncHasCustomZone(msg *network_map.UpdateMessage, domain string) bool {
|
||||
for _, zone := range msg.Update.GetNetworkMap().GetDNSConfig().GetCustomZones() {
|
||||
if zone.GetDomain() == domain {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -22,6 +22,7 @@ import (
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/zones"
|
||||
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
|
||||
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
|
||||
networkTypes "github.com/netbirdio/netbird/management/server/networks/types"
|
||||
@@ -50,6 +51,7 @@ type Snapshot struct {
|
||||
policies []*types.Policy
|
||||
routes []*route.Route
|
||||
nsGroups []*nbdns.NameServerGroup
|
||||
zones []*zones.Zone
|
||||
dnsSettings *types.DNSSettings
|
||||
routers []*routerTypes.NetworkRouter
|
||||
resources []*resourceTypes.NetworkResource
|
||||
@@ -127,12 +129,15 @@ func (snap *Snapshot) loadRoutesAndProxy(ctx context.Context, s store.Store, acc
|
||||
return snap.loadProxyServices(ctx, s, accountID)
|
||||
}
|
||||
|
||||
// loadDNS loads the nameserver groups and account DNS settings.
|
||||
// loadDNS loads the nameserver groups, custom DNS zones and account DNS settings.
|
||||
func (snap *Snapshot) loadDNS(ctx context.Context, s store.Store, accountID string) error {
|
||||
var err error
|
||||
if snap.nsGroups, err = s.GetAccountNameServerGroups(ctx, store.LockingStrengthNone, accountID); err != nil {
|
||||
return err
|
||||
}
|
||||
if snap.zones, err = s.GetAccountZones(ctx, store.LockingStrengthNone, accountID); err != nil {
|
||||
return err
|
||||
}
|
||||
snap.dnsSettings, err = s.GetAccountDNSSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
return err
|
||||
}
|
||||
@@ -357,7 +362,7 @@ func (s policySide) opposite() policySide {
|
||||
// - a changed router/resource/network sits on a NETWORK -> fold the SOURCE side of
|
||||
// the policies whose destination reaches it (and the routers it implies).
|
||||
//
|
||||
// Routes, nameserver groups, DNS and embedded-proxy services distribute to their own
|
||||
// Routes, nameserver groups, DNS zones, DNS and embedded-proxy services distribute to their own
|
||||
// member peers, outside the policy graph, and are folded here too.
|
||||
func (r *resolver) walk() {
|
||||
for _, policy := range r.bothSidesPolicies() {
|
||||
@@ -369,6 +374,7 @@ func (r *resolver) walk() {
|
||||
r.collectFromPolicies()
|
||||
r.collectFromRoutes()
|
||||
r.collectFromNameServers()
|
||||
r.collectFromZones()
|
||||
r.collectFromDNSSettings()
|
||||
r.collectFromNetworkRouters()
|
||||
r.collectFromProxyServices()
|
||||
@@ -829,6 +835,25 @@ func (r *resolver) collectFromNameServers() {
|
||||
}
|
||||
}
|
||||
|
||||
// collectFromZones folds the distribution groups of the custom DNS zones that
|
||||
// reference a linked group. Like nameserver groups, a zone has no opposite side, so
|
||||
// only a whole-group change folds its groups. Zones the network map does not ship
|
||||
// (disabled or without records) are skipped.
|
||||
func (r *resolver) collectFromZones() {
|
||||
if len(r.linkGroups) == 0 {
|
||||
return
|
||||
}
|
||||
for _, zone := range r.snap.zones {
|
||||
if !zone.Enabled || len(zone.Records) == 0 {
|
||||
continue
|
||||
}
|
||||
if anyInSet(zone.DistributionGroups, r.linkGroups) {
|
||||
log.WithContext(r.ctx).Tracef("collectFromZones: zone %s references a linked group -> folding its groups %v (outputGroups only)", zone.ID, zone.DistributionGroups)
|
||||
r.foldOutputGroups(zone.DistributionGroups)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// collectFromSSHAuthorizedGroups folds the destinations of the enabled SSH rules that
|
||||
// authorize a group whose user membership changed. Those destination peers carry the
|
||||
// group -> user mapping for the groups they authorize, so they refresh even when no
|
||||
|
||||
+33
-13
@@ -166,9 +166,11 @@ func (am *DefaultAccountManager) UpdateGroup(ctx context.Context, accountID, use
|
||||
return err
|
||||
}
|
||||
|
||||
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}); err != nil {
|
||||
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
change.ChangedPeerIDs = ipv6Changed
|
||||
|
||||
// A membership change does not alter which entities reference the group, so
|
||||
// the dependency walk runs once against the post-change snapshot. The new
|
||||
@@ -321,7 +323,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us
|
||||
var globalErr error
|
||||
for _, newGroup := range groups {
|
||||
change := affectedpeers.Change{ChangedGroupIDs: []string{newGroup.ID}}
|
||||
events, snap, err := am.updateSingleGroup(ctx, accountID, userID, newGroup, change)
|
||||
events, snap, change, err := am.updateSingleGroup(ctx, accountID, userID, newGroup, change)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to update group %s: %v", newGroup.ID, err)
|
||||
if len(groups) == 1 {
|
||||
@@ -344,7 +346,7 @@ func (am *DefaultAccountManager) UpdateGroups(ctx context.Context, accountID, us
|
||||
return globalErr
|
||||
}
|
||||
|
||||
func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountID, userID string, newGroup *types.Group, change affectedpeers.Change) ([]func(), *affectedpeers.Snapshot, error) {
|
||||
func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountID, userID string, newGroup *types.Group, change affectedpeers.Change) ([]func(), *affectedpeers.Snapshot, affectedpeers.Change, error) {
|
||||
var events []func()
|
||||
var snap *affectedpeers.Snapshot
|
||||
err := am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
|
||||
@@ -364,9 +366,11 @@ func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountI
|
||||
return err
|
||||
}
|
||||
|
||||
if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID}); err != nil {
|
||||
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{newGroup.ID})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
change.ChangedPeerIDs = ipv6Changed
|
||||
|
||||
if err := transaction.IncrementNetworkSerial(ctx, accountID); err != nil {
|
||||
return err
|
||||
@@ -377,7 +381,7 @@ func (am *DefaultAccountManager) updateSingleGroup(ctx context.Context, accountI
|
||||
snap, err = affectedpeers.Load(ctx, transaction, accountID, change)
|
||||
return err
|
||||
})
|
||||
return events, snap, err
|
||||
return events, snap, change, err
|
||||
}
|
||||
|
||||
// prepareGroupEvents prepares a list of event functions to be stored.
|
||||
@@ -480,8 +484,8 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us
|
||||
var allErrors error
|
||||
var groupIDsToDelete []string
|
||||
var deletedGroups []*types.Group
|
||||
var snap *affectedpeers.Snapshot
|
||||
var change affectedpeers.Change
|
||||
var snap, ipv6Snap *affectedpeers.Snapshot
|
||||
var change, ipv6Change affectedpeers.Change
|
||||
|
||||
extraSettings, err := am.settingsManager.GetExtraSettings(ctx, accountID)
|
||||
if err != nil {
|
||||
@@ -510,10 +514,20 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us
|
||||
return err
|
||||
}
|
||||
|
||||
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, groupIDsToDelete); err != nil {
|
||||
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, groupIDsToDelete)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Members of a deleted IPv6-enabled group lose their address, which the
|
||||
// pre-delete snapshot cannot see, so they are resolved post-delete.
|
||||
if len(ipv6Changed) > 0 {
|
||||
ipv6Change = affectedpeers.Change{ChangedPeerIDs: ipv6Changed}
|
||||
if ipv6Snap, err = affectedpeers.Load(ctx, transaction, accountID, ipv6Change); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return transaction.IncrementNetworkSerial(ctx, accountID)
|
||||
})
|
||||
if err != nil {
|
||||
@@ -524,7 +538,7 @@ func (am *DefaultAccountManager) DeleteGroups(ctx context.Context, accountID, us
|
||||
am.StoreEvent(ctx, userID, group.ID, accountID, activity.GroupDeleted, group.EventMeta())
|
||||
}
|
||||
|
||||
am.ExpandAndUpdateAffected(ctx, accountID, snap, change)
|
||||
go am.dispatchAffected(ctx, accountID, []*affectedpeers.Snapshot{snap, ipv6Snap}, []affectedpeers.Change{change, ipv6Change})
|
||||
|
||||
return allErrors
|
||||
}
|
||||
@@ -564,11 +578,14 @@ func (am *DefaultAccountManager) GroupAddPeer(ctx context.Context, accountID, gr
|
||||
return err
|
||||
}
|
||||
|
||||
if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}); err != nil {
|
||||
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// A peer whose IPv6 address changed is visible to every peer that reaches it
|
||||
// through any of its groups, not only through this one.
|
||||
change.ChangedPeerIDs = ipv6Changed
|
||||
|
||||
var err error
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -634,11 +651,14 @@ func (am *DefaultAccountManager) GroupDeletePeer(ctx context.Context, accountID,
|
||||
return err
|
||||
}
|
||||
|
||||
if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID}); err != nil {
|
||||
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, []string{groupID})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// A peer whose IPv6 address changed is visible to every peer that reaches it
|
||||
// through any of its groups, not only through this one.
|
||||
change.ChangedPeerIDs = ipv6Changed
|
||||
|
||||
var err error
|
||||
if snap, err = affectedpeers.Load(ctx, transaction, accountID, change); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -127,7 +127,7 @@ func (h *handler) validateNetworkRange(ctx context.Context, accountID, userID st
|
||||
}
|
||||
|
||||
func (h *handler) validateCapacity(ctx context.Context, accountID, userID string, prefix netip.Prefix) error {
|
||||
peers, err := h.accountManager.GetPeers(ctx, accountID, userID, "", "")
|
||||
peers, err := h.accountManager.GetPeers(ctx, accountID, userID, "", "", "")
|
||||
if err != nil {
|
||||
return status.Errorf(status.Internal, "get peer count: %v", err)
|
||||
}
|
||||
|
||||
@@ -58,7 +58,7 @@ func (h *handler) getAllGroups(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
@@ -77,7 +77,7 @@ func (h *handler) getAllGroups(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
@@ -169,7 +169,7 @@ func (h *handler) updateGroup(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
@@ -226,7 +226,7 @@ func (h *handler) createGroup(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
@@ -287,7 +287,7 @@ func (h *handler) getGroup(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "")
|
||||
accountPeers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, "", "", "")
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
|
||||
@@ -78,7 +78,7 @@ func initGroupTestData(initGroups ...*types.Group) *handler {
|
||||
|
||||
return nil, status.Errorf(status.NotFound, "unknown group name")
|
||||
},
|
||||
GetPeersFunc: func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
|
||||
GetPeersFunc: func(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
|
||||
return maps.Values(TestPeers), nil
|
||||
},
|
||||
DeleteGroupFunc: func(_ context.Context, accountID, userId, groupID string) error {
|
||||
|
||||
@@ -317,10 +317,11 @@ func (h *Handler) GetAllPeers(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
nameFilter := r.URL.Query().Get("name")
|
||||
ipFilter := r.URL.Query().Get("ip")
|
||||
macFilter := r.URL.Query().Get("mac")
|
||||
|
||||
accountID, userID := userAuth.AccountId, userAuth.UserId
|
||||
|
||||
peers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, nameFilter, ipFilter)
|
||||
peers, err := h.accountManager.GetPeers(r.Context(), accountID, userID, nameFilter, ipFilter, macFilter)
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
@@ -571,6 +572,17 @@ func peerToAccessiblePeer(peer *nbpeer.Peer, dnsDomain string) api.AccessiblePee
|
||||
}
|
||||
}
|
||||
|
||||
func toNetworkAddresses(addrs []nbpeer.NetworkAddress) *[]api.NetworkAddress {
|
||||
if len(addrs) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]api.NetworkAddress, 0, len(addrs))
|
||||
for _, a := range addrs {
|
||||
out = append(out, api.NetworkAddress{NetIp: a.NetIP.String(), Mac: a.Mac})
|
||||
}
|
||||
return &out
|
||||
}
|
||||
|
||||
func toSinglePeerResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dnsDomain string, approved bool, reason string) *api.Peer {
|
||||
osVersion := peer.Meta.OSVersion
|
||||
if osVersion == "" {
|
||||
@@ -583,6 +595,7 @@ func toSinglePeerResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dnsD
|
||||
Name: peer.Name,
|
||||
Ip: peer.IP.String(),
|
||||
Ipv6: peerIPv6String(peer),
|
||||
NetworkAddresses: toNetworkAddresses(peer.Meta.NetworkAddresses),
|
||||
ConnectionIp: peer.Location.ConnectionIP.String(),
|
||||
Connected: peer.Status.Connected,
|
||||
LastSeen: peer.Status.LastSeen,
|
||||
@@ -639,6 +652,7 @@ func toPeerListItemResponse(peer *nbpeer.Peer, groupsInfo []api.GroupMinimum, dn
|
||||
Name: peer.Name,
|
||||
Ip: peer.IP.String(),
|
||||
Ipv6: peerIPv6String(peer),
|
||||
NetworkAddresses: toNetworkAddresses(peer.Meta.NetworkAddresses),
|
||||
ConnectionIp: peer.Location.ConnectionIP.String(),
|
||||
Connected: peer.Status.Connected,
|
||||
LastSeen: peer.Status.LastSeen,
|
||||
|
||||
@@ -173,7 +173,7 @@ func initTestMetaData(t *testing.T, peers ...*nbpeer.Peer) *Handler {
|
||||
return nil, fmt.Errorf("user not found")
|
||||
}
|
||||
},
|
||||
GetPeersFunc: func(_ context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
|
||||
GetPeersFunc: func(_ context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
|
||||
return peers, nil
|
||||
},
|
||||
GetPeerGroupsFunc: func(ctx context.Context, accountID, peerID string) ([]*types.Group, error) {
|
||||
@@ -364,6 +364,50 @@ func TestGetPeers(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPeerResponseNetworkAddresses(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
addresses []nbpeer.NetworkAddress
|
||||
wantJSON string
|
||||
}{
|
||||
{name: "not reported"},
|
||||
{name: "empty", addresses: []nbpeer.NetworkAddress{}},
|
||||
{
|
||||
name: "multiple interfaces",
|
||||
addresses: []nbpeer.NetworkAddress{
|
||||
{NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"},
|
||||
{NetIP: netip.MustParsePrefix("2001:db8::123/64"), Mac: "00:93:37:bd:83:10"},
|
||||
},
|
||||
wantJSON: `[{"net_ip":"192.168.0.11/24","mac":"00:93:37:bd:83:0f"},{"net_ip":"2001:db8::123/64","mac":"00:93:37:bd:83:10"}]`,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
peer := &nbpeer.Peer{
|
||||
Status: &nbpeer.PeerStatus{},
|
||||
Meta: nbpeer.PeerSystemMeta{NetworkAddresses: tt.addresses},
|
||||
}
|
||||
responses := map[string]any{
|
||||
"single peer": toSinglePeerResponse(peer, nil, "example.com", true, ""),
|
||||
"peer list": toPeerListItemResponse(peer, nil, "example.com", 0),
|
||||
}
|
||||
for name, response := range responses {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
body, err := json.Marshal(response)
|
||||
require.NoError(t, err)
|
||||
var fields map[string]json.RawMessage
|
||||
require.NoError(t, json.Unmarshal(body, &fields))
|
||||
if tt.wantJSON == "" {
|
||||
assert.NotContains(t, fields, "network_addresses", "unreported interfaces should be omitted")
|
||||
return
|
||||
}
|
||||
assert.JSONEq(t, tt.wantJSON, string(fields["network_addresses"]), "response should preserve interface addresses and MACs")
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetAccessiblePeers(t *testing.T) {
|
||||
peer1 := &nbpeer.Peer{
|
||||
ID: "peer1",
|
||||
|
||||
@@ -125,9 +125,9 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
|
||||
http.Error(w, "Failed to create session", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
query.Set("session_code", code)
|
||||
query.Set(auth.SessionCodeQueryParam, code)
|
||||
} else {
|
||||
query.Set("session_token", sessionToken)
|
||||
query.Set(auth.SessionTokenQueryParam, sessionToken)
|
||||
}
|
||||
redirectURL.RawQuery = query.Encode()
|
||||
|
||||
|
||||
@@ -532,8 +532,8 @@ func TestAuthCallback_UserAllowedToLogin(t *testing.T) {
|
||||
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"},
|
||||
{name: "legacy proxy", manager: testSessionCodeManager{}, wantParam: "session_token", absentParam: "nb_session_code"},
|
||||
{name: "compatible proxy", manager: testSessionCodeManager{supported: true}, wantParam: "nb_session_code", absentParam: "session_token"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -555,8 +555,8 @@ func TestAuthCallback_UserAllowedToLogin(t *testing.T) {
|
||||
require.Empty(t, location.Query().Get(tt.absentParam))
|
||||
require.Empty(t, location.Query().Get("error"))
|
||||
|
||||
if tt.wantParam == "session_code" {
|
||||
code := location.Query().Get("session_code")
|
||||
if tt.wantParam == "nb_session_code" {
|
||||
code := location.Query().Get("nb_session_code")
|
||||
response, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
Domain: location.Hostname(),
|
||||
SessionCode: code,
|
||||
|
||||
@@ -100,7 +100,7 @@ func (am *DefaultAccountManager) GetValidatedPeers(ctx context.Context, accountI
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
peers, err = am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "")
|
||||
peers, err = am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, "", "", "")
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
@@ -39,7 +39,7 @@ type MockAccountManager struct {
|
||||
GetAccountIDByUserIdFunc func(ctx context.Context, userAuth auth.UserAuth) (string, error)
|
||||
GetUserFromUserAuthFunc func(ctx context.Context, userAuth auth.UserAuth) (*types.User, error)
|
||||
ListUsersFunc func(ctx context.Context, accountID string) ([]*types.User, error)
|
||||
GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error)
|
||||
GetPeersFunc func(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error)
|
||||
MarkPeerConnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64, nmap *types.NetworkMap) error
|
||||
MarkPeerDisconnectedFunc func(ctx context.Context, peerKey string, accountID string, sessionStartedAt int64) error
|
||||
SyncAndMarkPeerFunc func(ctx context.Context, accountID string, peerPubKey string, meta nbpeer.PeerSystemMeta, realIP net.IP, syncTime time.Time) (*nbpeer.Peer, *types.NetworkMap, []*nmdata.PostureChecks, int64, error)
|
||||
@@ -807,9 +807,9 @@ func (am *MockAccountManager) GetAccountIDFromUserAuth(ctx context.Context, user
|
||||
}
|
||||
|
||||
// GetPeers mocks GetPeers of the AccountManager interface
|
||||
func (am *MockAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
|
||||
func (am *MockAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
|
||||
if am.GetPeersFunc != nil {
|
||||
return am.GetPeersFunc(ctx, accountID, userID, nameFilter, ipFilter)
|
||||
return am.GetPeersFunc(ctx, accountID, userID, nameFilter, ipFilter, macFilter)
|
||||
}
|
||||
return nil, status.Errorf(codes.Unimplemented, "method GetPeers is not implemented")
|
||||
}
|
||||
|
||||
@@ -47,7 +47,7 @@ const (
|
||||
|
||||
// GetPeers returns peers visible to the user within an account.
|
||||
// Users with "peers:read" see all peers. Otherwise, users see only their own peers, or none if restricted by account settings.
|
||||
func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
|
||||
func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
|
||||
user, err := am.Store.GetUserByUserID(ctx, store.LockingStrengthNone, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -59,7 +59,7 @@ func (am *DefaultAccountManager) GetPeers(ctx context.Context, accountID, userID
|
||||
}
|
||||
|
||||
if allowed {
|
||||
return am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, nameFilter, ipFilter)
|
||||
return am.Store.GetAccountPeers(ctx, store.LockingStrengthNone, accountID, nameFilter, ipFilter, macFilter)
|
||||
}
|
||||
|
||||
settings, err := am.Store.GetAccountSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
|
||||
@@ -4,10 +4,14 @@ import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
b64 "encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"runtime"
|
||||
"strconv"
|
||||
@@ -33,12 +37,15 @@ import (
|
||||
"github.com/netbirdio/netbird/management/internals/server/config"
|
||||
"github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
peershandler "github.com/netbirdio/netbird/management/server/http/handlers/peers"
|
||||
"github.com/netbirdio/netbird/management/server/http/testing/testing_tools"
|
||||
"github.com/netbirdio/netbird/management/server/integrations/port_forwarding"
|
||||
"github.com/netbirdio/netbird/management/server/job"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/settings"
|
||||
"github.com/netbirdio/netbird/shared/auth"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/util"
|
||||
@@ -718,7 +725,7 @@ func TestDefaultAccountManager_GetPeers(t *testing.T) {
|
||||
return
|
||||
}
|
||||
|
||||
peers, err := manager.GetPeers(context.Background(), accountID, someUser, "", "")
|
||||
peers, err := manager.GetPeers(context.Background(), accountID, someUser, "", "", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
return
|
||||
@@ -731,6 +738,71 @@ func TestDefaultAccountManager_GetPeers(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_GetPeers_FilterByMac(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
manager, _, err := createManager(t)
|
||||
require.NoError(t, err)
|
||||
account := newAccountWithId(ctx, "mac-account", "mac-admin", "", "", "", false)
|
||||
account.Peers["matching"] = &nbpeer.Peer{
|
||||
ID: "matching", Key: "matching-key", Name: "laptop", DNSLabel: "laptop",
|
||||
IP: netip.MustParseAddr("100.64.0.10"), Status: &nbpeer.PeerStatus{LastSeen: time.Now()},
|
||||
Meta: nbpeer.PeerSystemMeta{NetworkAddresses: []nbpeer.NetworkAddress{
|
||||
{NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"},
|
||||
{NetIP: netip.MustParsePrefix("192.168.1.11/24"), Mac: "aa:bb:cc:dd:ee:ff"},
|
||||
}},
|
||||
}
|
||||
account.Peers["other"] = &nbpeer.Peer{
|
||||
ID: "other", Key: "other-key", Name: "desktop", DNSLabel: "desktop",
|
||||
IP: netip.MustParseAddr("100.64.0.20"), Status: &nbpeer.PeerStatus{LastSeen: time.Now()},
|
||||
}
|
||||
require.NoError(t, manager.Store.SaveAccount(ctx, account))
|
||||
otherAccount := newAccountWithId(ctx, "other-account", "other-admin", "", "", "", false)
|
||||
otherPeer := account.Peers["matching"].Copy()
|
||||
otherPeer.ID, otherPeer.Key = "outside-account", "outside-key"
|
||||
otherAccount.Peers[otherPeer.ID] = otherPeer
|
||||
require.NoError(t, manager.Store.SaveAccount(ctx, otherAccount))
|
||||
handler := peershandler.NewHandler(manager, manager.networkMapController, manager.permissionsManager)
|
||||
|
||||
tests := []struct {
|
||||
name, nameFilter, ipFilter, macFilter string
|
||||
wantIDs []string
|
||||
}{
|
||||
{name: "no filter", wantIDs: []string{"matching", "other"}},
|
||||
{name: "full MAC", macFilter: "00:93:37:bd:83:0f", wantIDs: []string{"matching"}},
|
||||
{name: "partial MAC", macFilter: "93:37:bd", wantIDs: []string{"matching"}},
|
||||
{name: "second interface", macFilter: "aa:bb:cc:dd:ee:ff", wantIDs: []string{"matching"}},
|
||||
{name: "unknown MAC", macFilter: "11:22:33:44:55:66"},
|
||||
{name: "combined filters", nameFilter: "laptop", ipFilter: "100.64.0.10", macFilter: "00:93:37", wantIDs: []string{"matching"}},
|
||||
{name: "name mismatch", nameFilter: "desktop", macFilter: "00:93:37"},
|
||||
{name: "IP mismatch", ipFilter: "100.64.0.20", macFilter: "00:93:37"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
peers, err := manager.GetPeers(ctx, account.Id, "mac-admin", tt.nameFilter, tt.ipFilter, tt.macFilter)
|
||||
require.NoError(t, err)
|
||||
ids := make([]string, 0, len(peers))
|
||||
for _, peer := range peers {
|
||||
ids = append(ids, peer.ID)
|
||||
}
|
||||
assert.ElementsMatch(t, tt.wantIDs, ids, "filters should return only matching peers in the account")
|
||||
|
||||
query := url.Values{"name": {tt.nameFilter}, "ip": {tt.ipFilter}, "mac": {tt.macFilter}}
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/peers?"+query.Encode(), nil)
|
||||
req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{AccountId: account.Id, UserId: "mac-admin"})
|
||||
recorder := httptest.NewRecorder()
|
||||
handler.GetAllPeers(recorder, req)
|
||||
require.Equal(t, http.StatusOK, recorder.Code, "peer listing should succeed: %s", recorder.Body.String())
|
||||
var response []api.PeerBatch
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response))
|
||||
responseIDs := make([]string, 0, len(response))
|
||||
for _, peer := range response {
|
||||
responseIDs = append(responseIDs, peer.Id)
|
||||
}
|
||||
assert.ElementsMatch(t, tt.wantIDs, responseIDs, "HTTP query filters should reach the store")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func setupTestAccountManager(b testing.TB, peers int, groups int) (*DefaultAccountManager, *update_channel.PeersUpdateManager, string, string, error) {
|
||||
b.Helper()
|
||||
|
||||
@@ -934,7 +1006,7 @@ func BenchmarkGetPeers(b *testing.B) {
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_, err := manager.GetPeers(context.Background(), accountID, userID, "", "")
|
||||
_, err := manager.GetPeers(context.Background(), accountID, userID, "", "", "")
|
||||
if err != nil {
|
||||
b.Fatalf("GetPeers failed: %v", err)
|
||||
}
|
||||
|
||||
@@ -492,7 +492,7 @@ func (s *SqlStore) GetPeerByPeerPubKey(ctx context.Context, lockStrength Locking
|
||||
}
|
||||
|
||||
// GetAccountPeers retrieves peers for an account.
|
||||
func (s *SqlStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error) {
|
||||
func (s *SqlStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error) {
|
||||
var peers []*nbpeer.Peer
|
||||
tx := s.db
|
||||
if lockStrength != LockingStrengthNone {
|
||||
@@ -506,6 +506,11 @@ func (s *SqlStore) GetAccountPeers(ctx context.Context, lockStrength LockingStre
|
||||
if ipFilter != "" {
|
||||
query = query.Where("ip LIKE ? OR ipv6 LIKE ?", "%"+ipFilter+"%", "%"+ipFilter+"%")
|
||||
}
|
||||
// MAC addresses live in the JSON-serialized meta_network_addresses column,
|
||||
// so we match the raw JSON text rather than a dedicated column.
|
||||
if macFilter != "" {
|
||||
query = query.Where("meta_network_addresses LIKE ?", "%"+macFilter+"%")
|
||||
}
|
||||
|
||||
if err := query.Find(&peers).Error; err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to get peers from the store: %s", err)
|
||||
|
||||
@@ -512,7 +512,7 @@ func TestSqlStore_GetAccountPeers(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
peers, err := store.GetAccountPeers(context.Background(), LockingStrengthNone, tt.accountID, tt.nameFilter, tt.ipFilter)
|
||||
peers, err := store.GetAccountPeers(context.Background(), LockingStrengthNone, tt.accountID, tt.nameFilter, tt.ipFilter, "")
|
||||
require.NoError(t, err)
|
||||
require.Len(t, peers, tt.expectedCount)
|
||||
})
|
||||
@@ -520,6 +520,48 @@ func TestSqlStore_GetAccountPeers(t *testing.T) {
|
||||
|
||||
}
|
||||
|
||||
func TestSqlStore_GetAccountPeers_FilterByMac(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store, cleanup, err := NewTestStoreFromSQL(ctx, "", t.TempDir())
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(cleanup)
|
||||
|
||||
accountID := "test-account-mac"
|
||||
userID := "test-user-mac"
|
||||
account := newAccountWithId(ctx, accountID, userID, "example.com")
|
||||
account.Peers["peer-mac-1"] = &nbpeer.Peer{
|
||||
ID: "peer-mac-1",
|
||||
AccountID: accountID,
|
||||
Key: "peer-mac-key-1",
|
||||
Name: "macpeer",
|
||||
IP: netip.MustParseAddr("100.64.0.10"),
|
||||
Meta: nbpeer.PeerSystemMeta{
|
||||
NetworkAddresses: []nbpeer.NetworkAddress{
|
||||
{NetIP: netip.MustParsePrefix("192.168.0.11/24"), Mac: "00:93:37:bd:83:0f"},
|
||||
},
|
||||
},
|
||||
}
|
||||
require.NoError(t, store.SaveAccount(ctx, account))
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
macFilter string
|
||||
expectedCount int
|
||||
}{
|
||||
{name: "full mac matches", macFilter: "00:93:37:bd:83:0f", expectedCount: 1},
|
||||
{name: "mac prefix matches", macFilter: "00:93:37", expectedCount: 1},
|
||||
{name: "unknown mac does not match", macFilter: "11:22:33:44:55:66", expectedCount: 0},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
peers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "", tt.macFilter)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, peers, tt.expectedCount)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSqlStore_GetAccountPeersWithExpiration(t *testing.T) {
|
||||
store, cleanup, err := NewTestStoreFromSQL(context.Background(), "../testdata/store_with_expired_peers.sql", t.TempDir())
|
||||
t.Cleanup(cleanup)
|
||||
@@ -878,7 +920,7 @@ func TestSqlStore_ApproveAccountPeers(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 2, count)
|
||||
|
||||
allPeers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "")
|
||||
allPeers, err := store.GetAccountPeers(ctx, LockingStrengthNone, accountID, "", "", "")
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, peer := range allPeers {
|
||||
|
||||
@@ -358,6 +358,14 @@ func (s *SqlStore) GetClusterSupportsPrivate(ctx context.Context, clusterAddr st
|
||||
return s.getClusterCapability(ctx, clusterAddr, "private")
|
||||
}
|
||||
|
||||
// GetClusterAllProxiesPrivate reports whether every active proxy in the cluster
|
||||
// has the private capability. Returns nil when no proxy reported the capability.
|
||||
// Use it where any proxy in the cluster may serve the result, since a single
|
||||
// non-private proxy would serve it without the private guarantees.
|
||||
func (s *SqlStore) GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
|
||||
return s.getClusterUnanimousCapability(ctx, clusterAddr, "private")
|
||||
}
|
||||
|
||||
// GetClusterSupportsCrowdSec returns whether all active proxies in the cluster
|
||||
// have CrowdSec configured. Returns nil when no proxy reported the capability.
|
||||
// Unlike other capabilities that use ANY-true (for rolling upgrades), CrowdSec
|
||||
|
||||
@@ -160,7 +160,7 @@ type Store interface {
|
||||
RemoveResourceFromGroup(ctx context.Context, accountId string, groupID string, resourceID string) error
|
||||
AddPeerToAccount(ctx context.Context, peer *nbpeer.Peer) error
|
||||
GetPeerByPeerPubKey(ctx context.Context, lockStrength LockingStrength, peerKey string) (*nbpeer.Peer, error)
|
||||
GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*nbpeer.Peer, error)
|
||||
GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*nbpeer.Peer, error)
|
||||
GetUserPeers(ctx context.Context, lockStrength LockingStrength, accountID, userID string) ([]*nbpeer.Peer, error)
|
||||
GetPeerByID(ctx context.Context, lockStrength LockingStrength, accountID string, peerID string) (*nbpeer.Peer, error)
|
||||
GetPeersByIDs(ctx context.Context, lockStrength LockingStrength, accountID string, peerIDs []string) (map[string]*nbpeer.Peer, error)
|
||||
@@ -336,6 +336,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
|
||||
GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error)
|
||||
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error)
|
||||
|
||||
@@ -1330,18 +1330,18 @@ func (mr *MockStoreMockRecorder) GetAccountOwner(ctx, lockStrength, accountID an
|
||||
}
|
||||
|
||||
// GetAccountPeers mocks base method.
|
||||
func (m *MockStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter string) ([]*peer.Peer, error) {
|
||||
func (m *MockStore) GetAccountPeers(ctx context.Context, lockStrength LockingStrength, accountID, nameFilter, ipFilter, macFilter string) ([]*peer.Peer, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetAccountPeers", ctx, lockStrength, accountID, nameFilter, ipFilter)
|
||||
ret := m.ctrl.Call(m, "GetAccountPeers", ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter)
|
||||
ret0, _ := ret[0].([]*peer.Peer)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetAccountPeers indicates an expected call of GetAccountPeers.
|
||||
func (mr *MockStoreMockRecorder) GetAccountPeers(ctx, lockStrength, accountID, nameFilter, ipFilter any) *gomock.Call {
|
||||
func (mr *MockStoreMockRecorder) GetAccountPeers(ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountPeers", reflect.TypeOf((*MockStore)(nil).GetAccountPeers), ctx, lockStrength, accountID, nameFilter, ipFilter)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountPeers", reflect.TypeOf((*MockStore)(nil).GetAccountPeers), ctx, lockStrength, accountID, nameFilter, ipFilter, macFilter)
|
||||
}
|
||||
|
||||
// GetAccountPeersWithExpiration mocks base method.
|
||||
@@ -1870,6 +1870,20 @@ func (mr *MockStoreMockRecorder) GetAnyAccountID(ctx any) *gomock.Call {
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAnyAccountID", reflect.TypeOf((*MockStore)(nil).GetAnyAccountID), ctx)
|
||||
}
|
||||
|
||||
// GetClusterAllProxiesPrivate mocks base method.
|
||||
func (m *MockStore) GetClusterAllProxiesPrivate(ctx context.Context, clusterAddr string) *bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetClusterAllProxiesPrivate", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].(*bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// GetClusterAllProxiesPrivate indicates an expected call of GetClusterAllProxiesPrivate.
|
||||
func (mr *MockStoreMockRecorder) GetClusterAllProxiesPrivate(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusterAllProxiesPrivate", reflect.TypeOf((*MockStore)(nil).GetClusterAllProxiesPrivate), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// GetClusterRequireSubdomain mocks base method.
|
||||
func (m *MockStore) GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -861,9 +861,11 @@ func (am *DefaultAccountManager) processUserUpdate(ctx context.Context, transact
|
||||
allGroupChanges := slices.Concat(removedGroups, addedGroups)
|
||||
change.LinkGroups = allGroupChanges
|
||||
|
||||
if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, allGroupChanges); err != nil {
|
||||
ipv6Changed, err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, allGroupChanges)
|
||||
if err != nil {
|
||||
return change, nil, nil, nil, fmt.Errorf("reconcile IPv6 for group changes: %w", err)
|
||||
}
|
||||
change.ChangedPeerIDs = append(change.ChangedPeerIDs, ipv6Changed...)
|
||||
}
|
||||
|
||||
userEventsToAdd := am.prepareUserUpdateEvents(ctx, updatedUser.AccountID, initiatorUserId, oldUser, updatedUser, transferredOwnerRole, isNewUser, removedGroups, addedGroups, transaction)
|
||||
|
||||
@@ -30,6 +30,14 @@ const (
|
||||
SessionJWTIssuer = "netbird-management"
|
||||
)
|
||||
|
||||
// Query parameters management uses to hand the OIDC session to the proxy. The
|
||||
// proxy strips them before forwarding, so they must not collide with names the
|
||||
// proxied service uses itself.
|
||||
const (
|
||||
SessionCodeQueryParam = "nb_session_code"
|
||||
SessionTokenQueryParam = "session_token"
|
||||
)
|
||||
|
||||
// HeaderUserID is the synthetic user id recorded for header-authenticated
|
||||
// requests. Header auth validates a per-service secret and resolves no user
|
||||
// record, so proxy access logs and management-minted session tokens both
|
||||
|
||||
@@ -583,7 +583,7 @@ func (mw *Middleware) authenticateWithSchemes(w http.ResponseWriter, r *http.Req
|
||||
// handleAuthenticatedToken validates the token, handles denied access, and on
|
||||
// success sets a session cookie and redirects to the original URL.
|
||||
func (mw *Middleware) handleAuthenticatedToken(w http.ResponseWriter, r *http.Request, host, token string, config DomainConfig, scheme Scheme) {
|
||||
isCode := scheme.Type() == auth.MethodOIDC && r.URL.Query().Get("session_code") != ""
|
||||
isCode := scheme.Type() == auth.MethodOIDC && r.URL.Query().Get(auth.SessionCodeQueryParam) != ""
|
||||
result, err := mw.validateSessionToken(r.Context(), host, token, isCode, config.SessionPublicKey, scheme.Type())
|
||||
if err != nil {
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
@@ -661,7 +661,7 @@ func wasCredentialSubmitted(r *http.Request, method auth.Method) bool {
|
||||
case auth.MethodPassword:
|
||||
return credentialFormValue(r, passwordFormId) != ""
|
||||
case auth.MethodOIDC:
|
||||
return r.URL.Query().Get("session_token") != "" || r.URL.Query().Get("session_code") != ""
|
||||
return r.URL.Query().Get(auth.SessionTokenQueryParam) != "" || r.URL.Query().Get(auth.SessionCodeQueryParam) != ""
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -806,11 +806,11 @@ func sessionGroupsAllowed(allowed map[string]struct{}, method auth.Method, group
|
||||
// or history.
|
||||
func stripSessionTokenParam(u *url.URL) string {
|
||||
q := u.Query()
|
||||
if !q.Has("session_token") && !q.Has("session_code") {
|
||||
if !q.Has(auth.SessionTokenQueryParam) && !q.Has(auth.SessionCodeQueryParam) {
|
||||
return u.RequestURI()
|
||||
}
|
||||
q.Del("session_token")
|
||||
q.Del("session_code")
|
||||
q.Del(auth.SessionTokenQueryParam)
|
||||
q.Del(auth.SessionCodeQueryParam)
|
||||
clean := *u
|
||||
clean.RawQuery = q.Encode()
|
||||
return clean.RequestURI()
|
||||
|
||||
@@ -786,9 +786,15 @@ func TestWasCredentialSubmitted(t *testing.T) {
|
||||
{
|
||||
name: "OIDC code in query",
|
||||
method: auth.MethodOIDC,
|
||||
query: url.Values{"session_code": {"abc123"}},
|
||||
query: url.Values{"nb_session_code": {"abc123"}},
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "OIDC backend session_code in query",
|
||||
method: auth.MethodOIDC,
|
||||
query: url.Values{"session_code": {"abc123"}},
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "OIDC token not in query",
|
||||
method: auth.MethodOIDC,
|
||||
@@ -1585,8 +1591,9 @@ func TestStripSessionTokenParam(t *testing.T) {
|
||||
want string
|
||||
}{
|
||||
{"strips session_token", "https://ex.com/p?a=1&session_token=tok", "/p?a=1"},
|
||||
{"strips session_code", "https://ex.com/p?a=1&session_code=code", "/p?a=1"},
|
||||
{"strips both", "https://ex.com/p?session_token=tok&session_code=code&a=1", "/p?a=1"},
|
||||
{"strips nb_session_code", "https://ex.com/p?a=1&nb_session_code=code", "/p?a=1"},
|
||||
{"strips both", "https://ex.com/p?session_token=tok&nb_session_code=code&a=1", "/p?a=1"},
|
||||
{"keeps backend session_code", "https://ex.com/p?a=1&session_code=backend", "/p?a=1&session_code=backend"},
|
||||
{"no-op when absent", "https://ex.com/p?a=1", "/p?a=1"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
|
||||
@@ -43,12 +43,12 @@ func (o OIDC) Authenticate(r *http.Request) (string, string, error) {
|
||||
// Check for the session credential returned by the OIDC callback. The management
|
||||
// server passes it in the URL because it cannot set a cookie for the proxy's
|
||||
// domain (cookies are domain-scoped per RFC 6265). The current flow uses a
|
||||
// single-use session_code to keep the durable token out of the URL.
|
||||
// single-use session code to keep the durable token out of the URL.
|
||||
// session_token remains supported for backward compatibility.
|
||||
if code := r.URL.Query().Get("session_code"); code != "" {
|
||||
if code := r.URL.Query().Get(auth.SessionCodeQueryParam); code != "" {
|
||||
return code, "", nil
|
||||
}
|
||||
if token := r.URL.Query().Get("session_token"); token != "" {
|
||||
if token := r.URL.Query().Get(auth.SessionTokenQueryParam); token != "" {
|
||||
return token, "", nil
|
||||
}
|
||||
|
||||
|
||||
@@ -725,9 +725,9 @@ func stripSessionCookie(r *httputil.ProxyRequest) {
|
||||
// from the outgoing URL to prevent credential leakage to backends.
|
||||
func stripSessionTokenQuery(r *httputil.ProxyRequest) {
|
||||
q := r.Out.URL.Query()
|
||||
if q.Has("session_token") || q.Has("session_code") {
|
||||
q.Del("session_token")
|
||||
q.Del("session_code")
|
||||
if q.Has(auth.SessionTokenQueryParam) || q.Has(auth.SessionCodeQueryParam) {
|
||||
q.Del(auth.SessionTokenQueryParam)
|
||||
q.Del(auth.SessionCodeQueryParam)
|
||||
r.Out.URL.RawQuery = q.Encode()
|
||||
}
|
||||
}
|
||||
@@ -809,6 +809,12 @@ func classifyProxyError(err error) (title, message string, code int, status web.
|
||||
http.StatusBadGateway,
|
||||
web.ErrorStatus{Proxy: false, Destination: false}
|
||||
|
||||
case errors.Is(err, roundtrip.ErrDirectUpstreamBlocked):
|
||||
return "Destination Not Allowed",
|
||||
"This proxy does not connect to private or internal addresses. Please contact your administrator.",
|
||||
http.StatusBadGateway,
|
||||
web.ErrorStatus{Proxy: false, Destination: false}
|
||||
|
||||
case errors.Is(err, roundtrip.ErrTooManyInflight):
|
||||
return "Service Overloaded",
|
||||
"The service is currently handling too many requests. Please try again shortly.",
|
||||
|
||||
@@ -236,6 +236,17 @@ func TestRewriteFunc_SessionTokenQueryStripping(t *testing.T) {
|
||||
"other query parameters must be preserved")
|
||||
})
|
||||
|
||||
t.Run("strips nb_session_code query parameter", func(t *testing.T) {
|
||||
pr := newProxyRequest(t, "http://example.com/callback?nb_session_code=code123&other=keep", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
|
||||
assert.Empty(t, pr.Out.URL.Query().Get("nb_session_code"),
|
||||
"OIDC session code must be stripped from backend request")
|
||||
assert.Equal(t, "keep", pr.Out.URL.Query().Get("other"),
|
||||
"other query parameters must be preserved")
|
||||
})
|
||||
|
||||
t.Run("preserves query when no session_token present", func(t *testing.T) {
|
||||
pr := newProxyRequest(t, "http://example.com/api?foo=bar&baz=qux", "1.2.3.4:5000")
|
||||
|
||||
@@ -1053,6 +1064,17 @@ func TestClassifyProxyError(t *testing.T) {
|
||||
wantCode: http.StatusBadGateway,
|
||||
wantStatus: web.ErrorStatus{Proxy: true, Destination: false},
|
||||
},
|
||||
{
|
||||
name: "direct upstream blocked by dial guard",
|
||||
err: &net.OpError{
|
||||
Op: "dial",
|
||||
Net: "tcp",
|
||||
Err: roundtrip.ErrDirectUpstreamBlocked,
|
||||
},
|
||||
wantTitle: "Destination Not Allowed",
|
||||
wantCode: http.StatusBadGateway,
|
||||
wantStatus: web.ErrorStatus{Proxy: false, Destination: false},
|
||||
},
|
||||
{
|
||||
name: "unknown error falls to default",
|
||||
err: errors.New("something unexpected"),
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
package roundtrip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/netip"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// ErrDirectUpstreamBlocked is returned when a direct-upstream dial targets
|
||||
// an address that is not globally reachable while
|
||||
// NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE is set.
|
||||
var ErrDirectUpstreamBlocked = errors.New("direct upstream address is not allowed")
|
||||
|
||||
// blockedUpstreamPrefixes are the ranges that reach the proxy host, its
|
||||
// cluster or its cloud provider rather than the public internet. NAT64
|
||||
// and 6to4 addresses are matched by the IPv4 address they embed.
|
||||
var blockedUpstreamPrefixes = []netip.Prefix{
|
||||
// IPv4
|
||||
netip.MustParsePrefix("0.0.0.0/8"), // "this network", including 0.0.0.0
|
||||
netip.MustParsePrefix("10.0.0.0/8"), // RFC1918
|
||||
netip.MustParsePrefix("100.64.0.0/10"), // CGNAT
|
||||
netip.MustParsePrefix("127.0.0.0/8"), // loopback
|
||||
netip.MustParsePrefix("169.254.0.0/16"), // link-local, cloud metadata services
|
||||
netip.MustParsePrefix("172.16.0.0/12"), // RFC1918
|
||||
netip.MustParsePrefix("192.0.0.0/24"), // IETF protocol assignments
|
||||
netip.MustParsePrefix("192.0.2.0/24"), // documentation
|
||||
netip.MustParsePrefix("192.88.99.0/24"), // 6to4 relay anycast (deprecated)
|
||||
netip.MustParsePrefix("192.168.0.0/16"), // RFC1918
|
||||
netip.MustParsePrefix("198.18.0.0/15"), // benchmarking
|
||||
netip.MustParsePrefix("198.51.100.0/24"), // documentation
|
||||
netip.MustParsePrefix("203.0.113.0/24"), // documentation
|
||||
netip.MustParsePrefix("224.0.0.0/4"), // multicast
|
||||
netip.MustParsePrefix("240.0.0.0/4"), // reserved, including broadcast
|
||||
|
||||
// IPv6
|
||||
netip.MustParsePrefix("::/96"), // unspecified, loopback, IPv4-compatible
|
||||
netip.MustParsePrefix("64:ff9b:1::/48"), // local-use NAT64
|
||||
netip.MustParsePrefix("100::/64"), // discard-only
|
||||
netip.MustParsePrefix("2001::/32"), // Teredo
|
||||
netip.MustParsePrefix("2001:2::/48"), // benchmarking
|
||||
netip.MustParsePrefix("2001:db8::/32"), // documentation
|
||||
netip.MustParsePrefix("3fff::/20"), // documentation
|
||||
netip.MustParsePrefix("5f00::/16"), // SRv6 SIDs
|
||||
netip.MustParsePrefix("fc00::/7"), // unique local, including AWS IMDS fd00:ec2::254
|
||||
netip.MustParsePrefix("fe80::/10"), // link-local
|
||||
netip.MustParsePrefix("fec0::/10"), // site-local (deprecated)
|
||||
netip.MustParsePrefix("ff00::/8"), // multicast
|
||||
}
|
||||
|
||||
var (
|
||||
nat64Prefix = netip.MustParsePrefix("64:ff9b::/96")
|
||||
sixToFour = netip.MustParsePrefix("2002::/16")
|
||||
)
|
||||
|
||||
// isBlockedUpstreamAddr reports whether a guarded direct-upstream dial
|
||||
// must refuse addr.
|
||||
func isBlockedUpstreamAddr(addr netip.Addr) bool {
|
||||
addr = addr.Unmap().WithZone("")
|
||||
if !addr.IsValid() {
|
||||
return true
|
||||
}
|
||||
|
||||
if nat64Prefix.Contains(addr) {
|
||||
b := addr.As16()
|
||||
return isBlockedUpstreamAddr(netip.AddrFrom4([4]byte(b[12:16])))
|
||||
}
|
||||
if sixToFour.Contains(addr) {
|
||||
b := addr.As16()
|
||||
return isBlockedUpstreamAddr(netip.AddrFrom4([4]byte(b[2:6])))
|
||||
}
|
||||
|
||||
for _, p := range blockedUpstreamPrefixes {
|
||||
if p.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// guardUpstreamDial is a net.Dialer ControlContext that refuses blocked
|
||||
// addresses. It sees the resolved address of each socket just before
|
||||
// connect, so DNS rebinding cannot swap the target after the check.
|
||||
func guardUpstreamDial(_ context.Context, _, address string, _ syscall.RawConn) error {
|
||||
ap, err := netip.ParseAddrPort(address)
|
||||
if err != nil || isBlockedUpstreamAddr(ap.Addr()) {
|
||||
return ErrDirectUpstreamBlocked
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,189 @@
|
||||
package roundtrip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestIsBlockedUpstreamAddr(t *testing.T) {
|
||||
blocked := []string{
|
||||
"0.0.0.0",
|
||||
"0.1.2.3",
|
||||
"10.1.2.3",
|
||||
"100.64.0.1",
|
||||
"100.127.255.254",
|
||||
"127.0.0.1",
|
||||
"127.255.255.255",
|
||||
"169.254.169.254",
|
||||
"172.16.0.1",
|
||||
"172.31.255.255",
|
||||
"192.0.0.170",
|
||||
"192.168.1.1",
|
||||
"192.88.99.1",
|
||||
"198.18.0.1",
|
||||
"224.0.0.1",
|
||||
"255.255.255.255",
|
||||
"::",
|
||||
"::1",
|
||||
"::169.254.169.254",
|
||||
"::ffff:127.0.0.1",
|
||||
"::ffff:169.254.169.254",
|
||||
"::ffff:10.0.0.1",
|
||||
"64:ff9b::a9fe:a9fe", // NAT64 of 169.254.169.254
|
||||
"64:ff9b::a00:1", // NAT64 of 10.0.0.1
|
||||
"64:ff9b:1::1",
|
||||
"2001::1",
|
||||
"2001:0:4136:e378:8000:63bf:3fff:fdd2",
|
||||
"2001:2::1",
|
||||
"3fff::1",
|
||||
"5f00::1",
|
||||
"2002:a9fe:a9fe::1", // 6to4 of 169.254.169.254
|
||||
"2002:7f00:1::", // 6to4 of 127.0.0.1
|
||||
"fc00::1",
|
||||
"fd00:ec2::254",
|
||||
"fe80::1",
|
||||
"fe80::1%eth0",
|
||||
"fec0::1",
|
||||
"ff02::1",
|
||||
}
|
||||
for _, s := range blocked {
|
||||
t.Run("blocks "+s, func(t *testing.T) {
|
||||
assert.True(t, isBlockedUpstreamAddr(netip.MustParseAddr(s)))
|
||||
})
|
||||
}
|
||||
|
||||
allowed := []string{
|
||||
"1.1.1.1",
|
||||
"8.8.8.8",
|
||||
"100.63.255.255",
|
||||
"100.128.0.0",
|
||||
"172.15.255.255",
|
||||
"172.32.0.0",
|
||||
"169.253.255.255",
|
||||
"2606:4700:4700::1111",
|
||||
"2001:4860:4860::8888",
|
||||
"2001:1::1",
|
||||
"4000::1",
|
||||
"::ffff:8.8.8.8",
|
||||
"64:ff9b::808:808", // NAT64 of 8.8.8.8
|
||||
"2002:808:808::1", // 6to4 of 8.8.8.8
|
||||
}
|
||||
for _, s := range allowed {
|
||||
t.Run("allows "+s, func(t *testing.T) {
|
||||
assert.False(t, isBlockedUpstreamAddr(netip.MustParseAddr(s)))
|
||||
})
|
||||
}
|
||||
|
||||
assert.True(t, isBlockedUpstreamAddr(netip.Addr{}), "the zero Addr must be refused")
|
||||
}
|
||||
|
||||
func TestGuardUpstreamDial_RejectsUnparsableAddress(t *testing.T) {
|
||||
err := guardUpstreamDial(context.Background(), "tcp", "not-an-address", nil)
|
||||
assert.ErrorIs(t, err, ErrDirectUpstreamBlocked, "an address the guard cannot parse must fail closed")
|
||||
}
|
||||
|
||||
// TestMultiTransport_BlockPrivateUpstreams exercises the guard end to end
|
||||
// against a loopback test server: by IP literal and by a hostname that
|
||||
// resolves to loopback, on both direct branches, and confirms the
|
||||
// embedded branch is not affected.
|
||||
func TestMultiTransport_BlockPrivateUpstreams(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = io.WriteString(w, "reached")
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
_, port, err := net.SplitHostPort(srv.Listener.Addr().String())
|
||||
require.NoError(t, err)
|
||||
byName := (&url.URL{Scheme: "http", Host: net.JoinHostPort("localhost", port)}).String()
|
||||
|
||||
directCtx := WithDirectUpstream(context.Background())
|
||||
insecureCtx := WithSkipTLSVerify(directCtx)
|
||||
|
||||
// roundTrip returns the response body, so callers never hold one open.
|
||||
roundTrip := func(t *testing.T, mt *MultiTransport, ctx context.Context, target string) (string, error) {
|
||||
t.Helper()
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil)
|
||||
require.NoError(t, err)
|
||||
resp, err := mt.RoundTrip(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
require.NoError(t, err)
|
||||
return string(body), nil
|
||||
}
|
||||
|
||||
t.Run("enabled refuses loopback", func(t *testing.T) {
|
||||
t.Setenv(EnvDirectUpstreamBlockPrivate, "true")
|
||||
mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
target string
|
||||
}{
|
||||
{"direct by IP", directCtx, srv.URL},
|
||||
{"direct by hostname", directCtx, byName},
|
||||
{"insecure by IP", insecureCtx, srv.URL},
|
||||
{"insecure by hostname", insecureCtx, byName},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := roundTrip(t, mt, tc.ctx, tc.target)
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, ErrDirectUpstreamBlocked)
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid value enables the guard", func(t *testing.T) {
|
||||
t.Setenv(EnvDirectUpstreamBlockPrivate, "yes please")
|
||||
mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil)
|
||||
|
||||
_, err := roundTrip(t, mt, directCtx, srv.URL)
|
||||
assert.ErrorIs(t, err, ErrDirectUpstreamBlocked, "a value that does not parse must fail closed")
|
||||
})
|
||||
|
||||
t.Run("explicit false disables the guard", func(t *testing.T) {
|
||||
t.Setenv(EnvDirectUpstreamBlockPrivate, "false")
|
||||
mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil)
|
||||
|
||||
body, err := roundTrip(t, mt, directCtx, srv.URL)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "reached", body)
|
||||
})
|
||||
|
||||
t.Run("enabled leaves embedded branch alone", func(t *testing.T) {
|
||||
t.Setenv(EnvDirectUpstreamBlockPrivate, "true")
|
||||
embedded := &stubRoundTripper{body: "embedded"}
|
||||
mt := NewMultiTransport(embedded, nil)
|
||||
|
||||
body, err := roundTrip(t, mt, context.Background(), srv.URL)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "embedded", body)
|
||||
assert.True(t, embedded.called, "the guard must not change dispatch to the embedded transport")
|
||||
})
|
||||
|
||||
t.Run("disabled by default", func(t *testing.T) {
|
||||
// Register the restore first so an exported value comes back after
|
||||
// the test, then exercise a genuinely absent variable.
|
||||
t.Setenv(EnvDirectUpstreamBlockPrivate, "")
|
||||
require.NoError(t, os.Unsetenv(EnvDirectUpstreamBlockPrivate))
|
||||
mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil)
|
||||
|
||||
body, err := roundTrip(t, mt, directCtx, srv.URL)
|
||||
require.NoError(t, err, "private and self-hosted proxies must keep reaching local upstreams")
|
||||
assert.Equal(t, "reached", body)
|
||||
})
|
||||
}
|
||||
@@ -41,7 +41,9 @@ var errNoEmbeddedTransport = errors.New("multitransport: embedded roundtripper n
|
||||
// MultiTransport that only ever uses the direct branch. The direct
|
||||
// branches honour the same NB_PROXY_* tuning env vars as the embedded
|
||||
// transport (see loadTransportConfig) plus a dial-timeout wrapper that
|
||||
// respects types.WithDialTimeout.
|
||||
// respects types.WithDialTimeout. With NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE
|
||||
// set, the direct branches refuse addresses that are not globally reachable
|
||||
// (see guardUpstreamDial).
|
||||
func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTransport {
|
||||
if logger == nil {
|
||||
logger = log.StandardLogger()
|
||||
@@ -51,6 +53,9 @@ func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTra
|
||||
Timeout: 30 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
if cfg.blockPrivateUpstreams {
|
||||
dialer.ControlContext = guardUpstreamDial
|
||||
}
|
||||
direct := &http.Transport{
|
||||
DialContext: dialWithTimeout(dialer.DialContext),
|
||||
MaxIdleConns: cfg.maxIdleConns,
|
||||
|
||||
@@ -25,6 +25,12 @@ const (
|
||||
EnvDisableCompression = "NB_PROXY_DISABLE_COMPRESSION"
|
||||
EnvMaxInflight = "NB_PROXY_MAX_INFLIGHT"
|
||||
EnvUpstreamHTTPVersion = "NB_PROXY_UPSTREAM_HTTP_VERSION"
|
||||
// EnvDirectUpstreamBlockPrivate refuses direct-upstream dials to
|
||||
// addresses that are not globally reachable (loopback, private,
|
||||
// link-local, CGNAT, ...). Off by default: private and self-hosted
|
||||
// proxies use direct_upstream to reach LAN and localhost services.
|
||||
// Proxies that serve untrusted accounts must turn it on.
|
||||
EnvDirectUpstreamBlockPrivate = "NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE"
|
||||
)
|
||||
|
||||
// upstreamHTTPVersion selects the HTTP version the proxy uses towards an
|
||||
@@ -69,6 +75,9 @@ type transportConfig struct {
|
||||
// explicit values are for backends whose advertised h2 support is
|
||||
// unusable and whose failure mode the negotiation cannot see.
|
||||
upstreamHTTPVersion upstreamHTTPVersion
|
||||
// blockPrivateUpstreams guards the direct branches' dialer with
|
||||
// guardUpstreamDial. It has no effect on the embedded branch.
|
||||
blockPrivateUpstreams bool
|
||||
}
|
||||
|
||||
func defaultTransportConfig() transportConfig {
|
||||
@@ -122,6 +131,7 @@ func loadTransportConfig(logger *log.Logger) transportConfig {
|
||||
if v, ok := envUpstreamHTTPVersion(EnvUpstreamHTTPVersion, logger); ok {
|
||||
cfg.upstreamHTTPVersion = v
|
||||
}
|
||||
cfg.blockPrivateUpstreams = envGuardBool(EnvDirectUpstreamBlockPrivate, logger)
|
||||
|
||||
logger.WithFields(log.Fields{
|
||||
"max_idle_conns": cfg.maxIdleConns,
|
||||
@@ -136,6 +146,7 @@ func loadTransportConfig(logger *log.Logger) transportConfig {
|
||||
"disable_compression": cfg.disableCompression,
|
||||
"max_inflight": cfg.maxInflight,
|
||||
"upstream_http_version": cfg.upstreamHTTPVersion,
|
||||
"block_private_upstreams": cfg.blockPrivateUpstreams,
|
||||
}).Debug("backend transport configuration")
|
||||
|
||||
return cfg
|
||||
@@ -246,6 +257,22 @@ func envDuration(key string, logger *log.Logger) (time.Duration, bool) {
|
||||
return v, true
|
||||
}
|
||||
|
||||
// envGuardBool reads a bool that turns a security guard on. Unset means
|
||||
// off, but a value that does not parse turns the guard on: a typo must not
|
||||
// leave a proxy that was meant to be guarded without the guard.
|
||||
func envGuardBool(key string, logger *log.Logger) bool {
|
||||
s := os.Getenv(key)
|
||||
if s == "" {
|
||||
return false
|
||||
}
|
||||
v, err := strconv.ParseBool(s)
|
||||
if err != nil {
|
||||
logger.Warnf("failed to parse %s=%q as bool, enabling it: %v", key, s, err)
|
||||
return true
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func envBool(key string, logger *log.Logger) (bool, bool) {
|
||||
s := os.Getenv(key)
|
||||
if s == "" {
|
||||
|
||||
@@ -246,6 +246,10 @@ func (m *testProxyManager) ClusterSupportsPrivate(_ context.Context, _ string) *
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *testProxyManager) ClusterAllProxiesPrivate(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *testProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
|
||||
return m.supportsSessionCode
|
||||
}
|
||||
|
||||
+4
-2
@@ -10,6 +10,8 @@ import (
|
||||
"net/url"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
)
|
||||
|
||||
// PathPrefix is the unique URL prefix for serving the proxy's own web assets.
|
||||
@@ -180,8 +182,8 @@ func ServeAccessDeniedPage(w http.ResponseWriter, r *http.Request, code int, tit
|
||||
// stripAuthParams returns the request URI with auth-related query parameters removed.
|
||||
func stripAuthParams(u *url.URL) string {
|
||||
q := u.Query()
|
||||
q.Del("session_token")
|
||||
q.Del("session_code")
|
||||
q.Del(auth.SessionTokenQueryParam)
|
||||
q.Del(auth.SessionCodeQueryParam)
|
||||
q.Del("error")
|
||||
q.Del("error_description")
|
||||
clean := *u
|
||||
|
||||
@@ -826,6 +826,20 @@ components:
|
||||
- ssh_enabled
|
||||
- login_expiration_enabled
|
||||
- inactivity_expiration_enabled
|
||||
NetworkAddress:
|
||||
type: object
|
||||
properties:
|
||||
net_ip:
|
||||
description: IP address with CIDR of the interface
|
||||
type: string
|
||||
example: 192.168.0.11/24
|
||||
mac:
|
||||
description: MAC address of the interface
|
||||
type: string
|
||||
example: "00:93:37:bd:83:0f"
|
||||
required:
|
||||
- net_ip
|
||||
- mac
|
||||
Peer:
|
||||
allOf:
|
||||
- $ref: '#/components/schemas/PeerMinimum'
|
||||
@@ -845,6 +859,11 @@ components:
|
||||
type: string
|
||||
format: ipv6
|
||||
example: "fd00:4e42:ab12::1"
|
||||
network_addresses:
|
||||
description: Network interfaces (IP + MAC) reported by the peer
|
||||
type: array
|
||||
items:
|
||||
$ref: '#/components/schemas/NetworkAddress'
|
||||
connection_ip:
|
||||
description: Peer's public connection IP address
|
||||
type: string
|
||||
@@ -7516,6 +7535,11 @@ paths:
|
||||
schema:
|
||||
type: string
|
||||
description: Filter peers by IP address
|
||||
- in: query
|
||||
name: mac
|
||||
schema:
|
||||
type: string
|
||||
description: Filter peers by MAC address of a network interface
|
||||
security:
|
||||
- BearerAuth: [ ]
|
||||
- TokenAuth: [ ]
|
||||
|
||||
@@ -3829,6 +3829,15 @@ type Network struct {
|
||||
RoutingPeersCount int `json:"routing_peers_count"`
|
||||
}
|
||||
|
||||
// NetworkAddress defines model for NetworkAddress.
|
||||
type NetworkAddress struct {
|
||||
// Mac MAC address of the interface
|
||||
Mac string `json:"mac"`
|
||||
|
||||
// NetIp IP address with CIDR of the interface
|
||||
NetIp string `json:"net_ip"`
|
||||
}
|
||||
|
||||
// NetworkRequest defines model for NetworkRequest.
|
||||
type NetworkRequest struct {
|
||||
// Description Network description
|
||||
@@ -4278,6 +4287,9 @@ type Peer struct {
|
||||
// Name Peer's hostname
|
||||
Name string `json:"name"`
|
||||
|
||||
// NetworkAddresses Network interfaces (IP + MAC) reported by the peer
|
||||
NetworkAddresses *[]NetworkAddress `json:"network_addresses,omitempty"`
|
||||
|
||||
// Os Peer's operating system and version
|
||||
Os string `json:"os"`
|
||||
|
||||
@@ -4372,6 +4384,9 @@ type PeerBatch struct {
|
||||
// Name Peer's hostname
|
||||
Name string `json:"name"`
|
||||
|
||||
// NetworkAddresses Network interfaces (IP + MAC) reported by the peer
|
||||
NetworkAddresses *[]NetworkAddress `json:"network_addresses,omitempty"`
|
||||
|
||||
// Os Peer's operating system and version
|
||||
Os string `json:"os"`
|
||||
|
||||
@@ -6294,6 +6309,9 @@ type GetApiPeersParams struct {
|
||||
|
||||
// Ip Filter peers by IP address
|
||||
Ip *string `form:"ip,omitempty" json:"ip,omitempty"`
|
||||
|
||||
// Mac Filter peers by MAC address of a network interface
|
||||
Mac *string `form:"mac,omitempty" json:"mac,omitempty"`
|
||||
}
|
||||
|
||||
// GetApiPeersPeerIdIngressPortsParams defines parameters for GetApiPeersPeerIdIngressPorts.
|
||||
|
||||
Reference in New Issue
Block a user