diff --git a/.github/workflows/golang-test-freebsd.yml b/.github/workflows/golang-test-freebsd.yml index 65c39147a..7bd48e3d0 100644 --- a/.github/workflows/golang-test-freebsd.yml +++ b/.github/workflows/golang-test-freebsd.yml @@ -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 diff --git a/.github/workflows/redhat-certify.yml b/.github/workflows/redhat-certify.yml new file mode 100644 index 000000000..e592dabc2 --- /dev/null +++ b/.github/workflows/redhat-certify.yml @@ -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_ +# 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 diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 673fcc281..dee4d398d 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -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 -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 diff --git a/base62/base62.go b/base62/base62.go index efafbc768..1a02e98e2 100644 --- a/base62/base62.go +++ b/base62/base62.go @@ -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) -} diff --git a/base62/base62_test.go b/base62/base62_test.go index 00da2124a..f2ad06d6f 100644 --- a/base62/base62_test.go +++ b/base62/base62_test.go @@ -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) + } +} diff --git a/client/android/client.go b/client/android/client.go index e47a1c13d..6f5eaacf3 100644 --- a/client/android/client.go +++ b/client/android/client.go @@ -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 } } diff --git a/client/iface/iface_test.go b/client/iface/iface_test.go index fff0d4e30..cb50ca4a1 100644 --- a/client/iface/iface_test.go +++ b/client/iface/iface_test.go @@ -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, diff --git a/client/iface/udpmux/mux.go b/client/iface/udpmux/mux.go index c5d2de4a5..68cecc953 100644 --- a/client/iface/udpmux/mux.go +++ b/client/iface/udpmux/mux.go @@ -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) diff --git a/client/internal/debug/debug.go b/client/internal/debug/debug.go index 3503881ef..a7e8e6a24 100644 --- a/client/internal/debug/debug.go +++ b/client/internal/debug/debug.go @@ -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) + } +} diff --git a/client/internal/debug/debug_test.go b/client/internal/debug/debug_test.go index 17d520358..6a810bccc 100644 --- a/client/internal/debug/debug_test.go +++ b/client/internal/debug/debug_test.go @@ -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") +} diff --git a/client/internal/dns/server_privileged_test.go b/client/internal/dns/server_privileged_test.go index a17044cf5..270e3bf91 100644 --- a/client/internal/dns/server_privileged_test.go +++ b/client/internal/dns/server_privileged_test.go @@ -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{ diff --git a/client/internal/dns/server_test.go b/client/internal/dns/server_test.go index 0144a4a8b..414890158 100644 --- a/client/internal/dns/server_test.go +++ b/client/internal/dns/server_test.go @@ -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() diff --git a/client/internal/engine.go b/client/internal/engine.go index ed39f697e..a52b561ec 100644 --- a/client/internal/engine.go +++ b/client/internal/engine.go @@ -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, diff --git a/client/internal/engine_privileged_test.go b/client/internal/engine_privileged_test.go index 1b047e017..2db0cd5ed 100644 --- a/client/internal/engine_privileged_test.go +++ b/client/internal/engine_privileged_test.go @@ -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) diff --git a/client/internal/engine_stdnet.go b/client/internal/engine_stdnet.go index 1ebb5779c..86f6d297a 100644 --- a/client/internal/engine_stdnet.go +++ b/client/internal/engine_stdnet.go @@ -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) } diff --git a/client/internal/engine_stdnet_android.go b/client/internal/engine_stdnet_android.go index de3c80bcf..b14deeadf 100644 --- a/client/internal/engine_stdnet_android.go +++ b/client/internal/engine_stdnet_android.go @@ -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) } diff --git a/client/internal/engine_test.go b/client/internal/engine_test.go index 2a7ecd652..14076c051 100644 --- a/client/internal/engine_test.go +++ b/client/internal/engine_test.go @@ -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) +} diff --git a/client/internal/peer/ice/agent.go b/client/internal/peer/ice/agent.go index c74b46d10..6cd8c48de 100644 --- a/client/internal/peer/ice/agent.go +++ b/client/internal/peer/ice/agent.go @@ -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() diff --git a/client/internal/peer/ice/stdnet.go b/client/internal/peer/ice/stdnet.go index 685ed0363..0c819ff66 100644 --- a/client/internal/peer/ice/stdnet.go +++ b/client/internal/peer/ice/stdnet.go @@ -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) } diff --git a/client/internal/peer/ice/stdnet_android.go b/client/internal/peer/ice/stdnet_android.go index 5033ec1b9..2962ecf66 100644 --- a/client/internal/peer/ice/stdnet_android.go +++ b/client/internal/peer/ice/stdnet_android.go @@ -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) } diff --git a/client/internal/peer/status.go b/client/internal/peer/status.go index 70c2689ac..fd42b48d2 100644 --- a/client/internal/peer/status.go +++ b/client/internal/peer/status.go @@ -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) { diff --git a/client/internal/peer/status_test.go b/client/internal/peer/status_test.go index 82dff0d6f..b3f01b217 100644 --- a/client/internal/peer/status_test.go +++ b/client/internal/peer/status_test.go @@ -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]) +} diff --git a/client/internal/relay/relay.go b/client/internal/relay/relay.go index 051717608..f0c65301e 100644 --- a/client/internal/relay/relay.go +++ b/client/internal/relay/relay.go @@ -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, diff --git a/client/internal/routemanager/client/client.go b/client/internal/routemanager/client/client.go index c691c54f8..973cf1ab8 100644 --- a/client/internal/routemanager/client/client.go +++ b/client/internal/routemanager/client/client.go @@ -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) } diff --git a/client/internal/routemanager/manager_test.go b/client/internal/routemanager/manager_test.go index 18b44820a..a1624cf46 100644 --- a/client/internal/routemanager/manager_test.go +++ b/client/internal/routemanager/manager_test.go @@ -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"), diff --git a/client/internal/routemanager/systemops/systemops_generic_test.go b/client/internal/routemanager/systemops/systemops_generic_test.go index c4f739c30..5b569ebd6 100644 --- a/client/internal/routemanager/systemops/systemops_generic_test.go +++ b/client/internal/routemanager/systemops/systemops_generic_test.go @@ -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, diff --git a/client/internal/stdnet/stdnet.go b/client/internal/stdnet/stdnet.go index 381886ac6..c3a9d3d97 100644 --- a/client/internal/stdnet/stdnet.go +++ b/client/internal/stdnet/stdnet.go @@ -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 diff --git a/client/internal/stdnet/stdnet_test.go b/client/internal/stdnet/stdnet_test.go new file mode 100644 index 000000000..822972f39 --- /dev/null +++ b/client/internal/stdnet/stdnet_test.go @@ -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) +} diff --git a/client/ui/frontend/src/components/ReadySignal.tsx b/client/ui/frontend/src/components/ReadySignal.tsx index 0d040cabc..6a98b30ea 100644 --- a/client/ui/frontend/src/components/ReadySignal.tsx +++ b/client/ui/frontend/src/components/ReadySignal.tsx @@ -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(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; }; diff --git a/client/ui/frontend/src/hooks/useAutoSizeWindow.ts b/client/ui/frontend/src/hooks/useAutoSizeWindow.ts index d4f4d80b2..6623e72c2 100644 --- a/client/ui/frontend/src/hooks/useAutoSizeWindow.ts +++ b/client/ui/frontend/src/hooks/useAutoSizeWindow.ts @@ -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(width: number, ready: boolean = true) { const ref = useRef(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(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(width: number, ready: b cancelAnimationFrame(raf2); i18next.off("languageChanged", scheduleApply); }; - }, [width, ready]); + }, [width, ready, generation]); return ref; } diff --git a/client/ui/frontend/src/modules/login/LoginWaitingForBrowserDialog.tsx b/client/ui/frontend/src/modules/login/LoginWaitingForBrowserDialog.tsx index efbd1ee84..f03751c4d 100644 --- a/client/ui/frontend/src/modules/login/LoginWaitingForBrowserDialog.tsx +++ b/client/ui/frontend/src/modules/login/LoginWaitingForBrowserDialog.tsx @@ -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(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); diff --git a/client/ui/services/connection.go b/client/ui/services/connection.go index f78ce4c0f..f6a8eca72 100644 --- a/client/ui/services/connection.go +++ b/client/ui/services/connection.go @@ -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") + } +} diff --git a/client/ui/services/windowmanager.go b/client/ui/services/windowmanager.go index 24319dae0..af6d726a3 100644 --- a/client/ui/services/windowmanager.go +++ b/client/ui/services/windowmanager.go @@ -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 } diff --git a/client/ui/services/windowmanager_test.go b/client/ui/services/windowmanager_test.go index 13c8548ab..890fba24f 100644 --- a/client/ui/services/windowmanager_test.go +++ b/client/ui/services/windowmanager_test.go @@ -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)) + }) + } +} diff --git a/infrastructure_files/configure.sh b/infrastructure_files/configure.sh index 92252d0b3..ce1a041e6 100755 --- a/infrastructure_files/configure.sh +++ b/infrastructure_files/configure.sh @@ -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 diff --git a/management/internals/controllers/network_map/controller/repository.go b/management/internals/controllers/network_map/controller/repository.go index 5c3195f16..c11af0b69 100644 --- a/management/internals/controllers/network_map/controller/repository.go +++ b/management/internals/controllers/network_map/controller/repository.go @@ -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) { diff --git a/management/internals/modules/peers/manager.go b/management/internals/modules/peers/manager.go index 3274ec524..e944be291 100644 --- a/management/internals/modules/peers/manager.go +++ b/management/internals/modules/peers/manager.go @@ -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) { diff --git a/management/internals/modules/reverseproxy/proxy/manager.go b/management/internals/modules/reverseproxy/proxy/manager.go index c0b8435ec..9350ad9b9 100644 --- a/management/internals/modules/reverseproxy/proxy/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager.go @@ -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) diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager.go b/management/internals/modules/reverseproxy/proxy/manager/manager.go index 7ddb66eec..5a95ea94a 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager.go @@ -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) diff --git a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go index 66ddb95bd..56806613a 100644 --- a/management/internals/modules/reverseproxy/proxy/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/proxy/manager/manager_test.go @@ -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) diff --git a/management/internals/modules/reverseproxy/proxy/manager_mock.go b/management/internals/modules/reverseproxy/proxy/manager_mock.go index d6f7197d7..5f3404096 100644 --- a/management/internals/modules/reverseproxy/proxy/manager_mock.go +++ b/management/internals/modules/reverseproxy/proxy/manager_mock.go @@ -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() diff --git a/management/internals/modules/reverseproxy/service/manager/manager.go b/management/internals/modules/reverseproxy/service/manager/manager.go index 62897c9ae..900b7759f 100644 --- a/management/internals/modules/reverseproxy/service/manager/manager.go +++ b/management/internals/modules/reverseproxy/service/manager/manager.go @@ -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 diff --git a/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go b/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go new file mode 100644 index 000000000..1f507294e --- /dev/null +++ b/management/internals/modules/reverseproxy/service/manager/private_cluster_test.go @@ -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)) +} diff --git a/management/internals/modules/zones/manager/manager.go b/management/internals/modules/zones/manager/manager.go index d5348d3d0..6f6ba6c40 100644 --- a/management/internals/modules/zones/manager/manager.go +++ b/management/internals/modules/zones/manager/manager.go @@ -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 } diff --git a/management/internals/modules/zones/records/manager/manager.go b/management/internals/modules/zones/records/manager/manager.go index b041aca30..16839c1b4 100644 --- a/management/internals/modules/zones/records/manager/manager.go +++ b/management/internals/modules/zones/records/manager/manager.go @@ -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 } diff --git a/management/server/account.go b/management/server/account.go index 340bcc84b..038c5d8db 100644 --- a/management/server/account.go +++ b/management/server/account.go @@ -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) } diff --git a/management/server/account/manager.go b/management/server/account/manager.go index 154c9ab18..2ac8584f4 100644 --- a/management/server/account/manager.go +++ b/management/server/account/manager.go @@ -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 diff --git a/management/server/account/manager_mock.go b/management/server/account/manager_mock.go index f31f63d0e..60075b169 100644 --- a/management/server/account/manager_mock.go +++ b/management/server/account/manager_mock.go @@ -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. diff --git a/management/server/account_test.go b/management/server/account_test.go index 8c735b28e..881ad19d7 100644 --- a/management/server/account_test.go +++ b/management/server/account_test.go @@ -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 { diff --git a/management/server/affected_peers_ipv6_test.go b/management/server/affected_peers_ipv6_test.go new file mode 100644 index 000000000..c64360016 --- /dev/null +++ b/management/server/affected_peers_ipv6_test.go @@ -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) +} diff --git a/management/server/affected_peers_user_test.go b/management/server/affected_peers_user_test.go index c0dbbb84f..3d73bbed0 100644 --- a/management/server/affected_peers_user_test.go +++ b/management/server/affected_peers_user_test.go @@ -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) diff --git a/management/server/affected_peers_zone_test.go b/management/server/affected_peers_zone_test.go new file mode 100644 index 000000000..4d622325c --- /dev/null +++ b/management/server/affected_peers_zone_test.go @@ -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 +} diff --git a/management/server/affectedpeers/resolver.go b/management/server/affectedpeers/resolver.go index cb2063ac9..895e4fd36 100644 --- a/management/server/affectedpeers/resolver.go +++ b/management/server/affectedpeers/resolver.go @@ -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 diff --git a/management/server/group.go b/management/server/group.go index 88295e2f6..8d91df3ab 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -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 } diff --git a/management/server/http/handlers/accounts/accounts_handler.go b/management/server/http/handlers/accounts/accounts_handler.go index c4cba5962..795214c31 100644 --- a/management/server/http/handlers/accounts/accounts_handler.go +++ b/management/server/http/handlers/accounts/accounts_handler.go @@ -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) } diff --git a/management/server/http/handlers/groups/groups_handler.go b/management/server/http/handlers/groups/groups_handler.go index ed01e7c3d..1a7753a57 100644 --- a/management/server/http/handlers/groups/groups_handler.go +++ b/management/server/http/handlers/groups/groups_handler.go @@ -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 diff --git a/management/server/http/handlers/groups/groups_handler_test.go b/management/server/http/handlers/groups/groups_handler_test.go index 78e4a2578..3e322db4e 100644 --- a/management/server/http/handlers/groups/groups_handler_test.go +++ b/management/server/http/handlers/groups/groups_handler_test.go @@ -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 { diff --git a/management/server/http/handlers/peers/peers_handler.go b/management/server/http/handlers/peers/peers_handler.go index 773b640e0..8a9bf1f70 100644 --- a/management/server/http/handlers/peers/peers_handler.go +++ b/management/server/http/handlers/peers/peers_handler.go @@ -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, diff --git a/management/server/http/handlers/peers/peers_handler_test.go b/management/server/http/handlers/peers/peers_handler_test.go index 592d64d1a..7054082cc 100644 --- a/management/server/http/handlers/peers/peers_handler_test.go +++ b/management/server/http/handlers/peers/peers_handler_test.go @@ -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", diff --git a/management/server/http/handlers/proxy/auth.go b/management/server/http/handlers/proxy/auth.go index 298fb503e..133236401 100644 --- a/management/server/http/handlers/proxy/auth.go +++ b/management/server/http/handlers/proxy/auth.go @@ -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() diff --git a/management/server/http/handlers/proxy/auth_callback_integration_test.go b/management/server/http/handlers/proxy/auth_callback_integration_test.go index 964841a63..862d5d5f2 100644 --- a/management/server/http/handlers/proxy/auth_callback_integration_test.go +++ b/management/server/http/handlers/proxy/auth_callback_integration_test.go @@ -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, diff --git a/management/server/integrated_validator.go b/management/server/integrated_validator.go index 9ec1f491e..5928a8ed2 100644 --- a/management/server/integrated_validator.go +++ b/management/server/integrated_validator.go @@ -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 } diff --git a/management/server/mock_server/account_mock.go b/management/server/mock_server/account_mock.go index 2f871c3e2..3313bf99c 100644 --- a/management/server/mock_server/account_mock.go +++ b/management/server/mock_server/account_mock.go @@ -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") } diff --git a/management/server/peer.go b/management/server/peer.go index 9f5572252..5d5863fa7 100644 --- a/management/server/peer.go +++ b/management/server/peer.go @@ -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) diff --git a/management/server/peer_test.go b/management/server/peer_test.go index 22f2b9b6f..5c3e02af5 100644 --- a/management/server/peer_test.go +++ b/management/server/peer_test.go @@ -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) } diff --git a/management/server/store/sql_store_peer.go b/management/server/store/sql_store_peer.go index e5086b6db..1b0e23cec 100644 --- a/management/server/store/sql_store_peer.go +++ b/management/server/store/sql_store_peer.go @@ -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) diff --git a/management/server/store/sql_store_peer_test.go b/management/server/store/sql_store_peer_test.go index b49e04f2f..1432b5d96 100644 --- a/management/server/store/sql_store_peer_test.go +++ b/management/server/store/sql_store_peer_test.go @@ -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 { diff --git a/management/server/store/sql_store_proxy.go b/management/server/store/sql_store_proxy.go index 58fa86468..bdccd282c 100644 --- a/management/server/store/sql_store_proxy.go +++ b/management/server/store/sql_store_proxy.go @@ -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 diff --git a/management/server/store/store.go b/management/server/store/store.go index 465f84413..01aaf4892 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -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) diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index cd9e7334d..956cac4b8 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -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() diff --git a/management/server/user.go b/management/server/user.go index 3510a624b..5f29f4df7 100644 --- a/management/server/user.go +++ b/management/server/user.go @@ -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) diff --git a/proxy/auth/auth.go b/proxy/auth/auth.go index 084046c49..605780959 100644 --- a/proxy/auth/auth.go +++ b/proxy/auth/auth.go @@ -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 diff --git a/proxy/internal/auth/middleware.go b/proxy/internal/auth/middleware.go index 672286748..647741139 100644 --- a/proxy/internal/auth/middleware.go +++ b/proxy/internal/auth/middleware.go @@ -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() diff --git a/proxy/internal/auth/middleware_test.go b/proxy/internal/auth/middleware_test.go index 88c900f97..cce35ae35 100644 --- a/proxy/internal/auth/middleware_test.go +++ b/proxy/internal/auth/middleware_test.go @@ -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 { diff --git a/proxy/internal/auth/oidc.go b/proxy/internal/auth/oidc.go index 739777924..0215fddc3 100644 --- a/proxy/internal/auth/oidc.go +++ b/proxy/internal/auth/oidc.go @@ -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 } diff --git a/proxy/internal/proxy/reverseproxy.go b/proxy/internal/proxy/reverseproxy.go index 7b0acd1b3..a3987fe5a 100644 --- a/proxy/internal/proxy/reverseproxy.go +++ b/proxy/internal/proxy/reverseproxy.go @@ -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.", diff --git a/proxy/internal/proxy/reverseproxy_test.go b/proxy/internal/proxy/reverseproxy_test.go index 83afee387..b26ca1f9f 100644 --- a/proxy/internal/proxy/reverseproxy_test.go +++ b/proxy/internal/proxy/reverseproxy_test.go @@ -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"), diff --git a/proxy/internal/roundtrip/dialguard.go b/proxy/internal/roundtrip/dialguard.go new file mode 100644 index 000000000..ac01263b4 --- /dev/null +++ b/proxy/internal/roundtrip/dialguard.go @@ -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 +} diff --git a/proxy/internal/roundtrip/dialguard_test.go b/proxy/internal/roundtrip/dialguard_test.go new file mode 100644 index 000000000..79453d8c3 --- /dev/null +++ b/proxy/internal/roundtrip/dialguard_test.go @@ -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) + }) +} diff --git a/proxy/internal/roundtrip/multi.go b/proxy/internal/roundtrip/multi.go index d50ad1fc9..a430d45bd 100644 --- a/proxy/internal/roundtrip/multi.go +++ b/proxy/internal/roundtrip/multi.go @@ -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, diff --git a/proxy/internal/roundtrip/transport.go b/proxy/internal/roundtrip/transport.go index 9e872e447..6383079c4 100644 --- a/proxy/internal/roundtrip/transport.go +++ b/proxy/internal/roundtrip/transport.go @@ -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 == "" { diff --git a/proxy/management_integration_test.go b/proxy/management_integration_test.go index 000d8ce72..03a9855de 100644 --- a/proxy/management_integration_test.go +++ b/proxy/management_integration_test.go @@ -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 } diff --git a/proxy/web/web.go b/proxy/web/web.go index a45fc8730..de3e4771a 100644 --- a/proxy/web/web.go +++ b/proxy/web/web.go @@ -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 diff --git a/shared/management/http/api/openapi.yml b/shared/management/http/api/openapi.yml index 3dd9f41f1..4b7077cac 100644 --- a/shared/management/http/api/openapi.yml +++ b/shared/management/http/api/openapi.yml @@ -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: [ ] diff --git a/shared/management/http/api/types.gen.go b/shared/management/http/api/types.gen.go index 9a90a72d3..009a9a7a7 100644 --- a/shared/management/http/api/types.gen.go +++ b/shared/management/http/api/types.gen.go @@ -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.