Merge remote-tracking branch 'origin/main' into feat-post_quantum_ml_kem

This commit is contained in:
riccardom
2026-10-05 16:03:28 +02:00
85 changed files with 3037 additions and 624 deletions
+1 -1
View File
@@ -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
+199
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
}
}
+30 -9
View File
@@ -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
View File
@@ -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,
+1 -4
View File
@@ -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)
+53 -1
View File
@@ -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)
}
}
+50
View File
@@ -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")
}
+7 -10
View File
@@ -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{
+1 -5
View File
@@ -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()
+1 -4
View File
@@ -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,
+15 -11
View File
@@ -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)
+1 -1
View File
@@ -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)
}
+1 -1
View File
@@ -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)
}
+20 -9
View File
@@ -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)
}
+1 -4
View File
@@ -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()
+1 -1
View File
@@ -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)
}
+1 -1
View File
@@ -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)
}
+20
View File
@@ -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) {
+23
View File
@@ -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])
}
+2 -10
View File
@@ -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)
}
+2 -4
View File
@@ -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,
+47 -38
View File
@@ -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
+136
View File
@@ -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);
+17 -13
View File
@@ -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
View File
@@ -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 }
+291 -2
View File
@@ -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))
})
}
}
+4 -4
View File
@@ -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
}
+92 -33
View File
@@ -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)
}
+1 -1
View File
@@ -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
+4 -4
View File
@@ -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.
+9 -9
View File
@@ -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
}
+27 -2
View File
@@ -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
View File
@@ -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,
+1 -1
View File
@@ -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")
}
+2 -2
View File
@@ -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)
+74 -2
View File
@@ -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)
}
+6 -1
View File
@@ -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)
+44 -2
View File
@@ -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
+2 -1
View File
@@ -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)
+18 -4
View File
@@ -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()
+3 -1
View File
@@ -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)
+8
View File
@@ -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
+5 -5
View File
@@ -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()
+10 -3
View File
@@ -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 {
+3 -3
View File
@@ -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
}
+9 -3
View File
@@ -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.",
+22
View File
@@ -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"),
+90
View File
@@ -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
}
+189
View File
@@ -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)
})
}
+6 -1
View File
@@ -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,
+27
View File
@@ -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 == "" {
+4
View File
@@ -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
View File
@@ -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
+24
View File
@@ -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: [ ]
+18
View File
@@ -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.