Compare commits

..

114 Commits

Author SHA1 Message Date
pascal
28074405e1 change wording from failer -> failing to pass codespell 2026-08-06 16:02:02 +02:00
pascal
fd23461b03 add network map compute test 2026-08-06 15:57:43 +02:00
Dmitri Dolguikh
b19cf7405d Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-06 15:18:13 +02:00
Dmitri Dolguikh
90e0c5bd0c added GetPolicies test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-06 15:17:39 +02:00
pascal
6ccb2bb239 handle nil values in policy, nameserver and resources 2026-08-06 15:12:12 +02:00
Dmitri Dolguikh
b160cb43b1 Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-06 13:03:28 +02:00
Dmitri Dolguikh
1902777509 added a test for GetPeers
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-06 13:03:00 +02:00
pascal
cff46eee67 introduce peerGroupsIdx to avoid looping through the groups to figure out peer membership 2026-08-06 12:49:36 +02:00
pascal
0513109c35 improve looping on connected peers filtering 2026-08-06 12:31:26 +02:00
Dmitri Dolguikh
6b8393c6b0 add GetAppliedZoneCandidatesViaPgxConnection test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 18:43:20 +02:00
Dmitri Dolguikh
316e82337f cleanup query execution in tests
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 18:09:28 +02:00
Dmitri Dolguikh
bacacb8f77 deleted group_test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 16:45:25 +02:00
Dmitri Dolguikh
994e84a686 Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types 2026-08-05 16:35:08 +02:00
Dmitri Dolguikh
07809fe923 added test for GetAllowedUsersViaPgxConnection
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 16:34:44 +02:00
pascal
ec3e37456e use Exec instead of Query for inserts and assert every insert error 2026-08-05 16:33:48 +02:00
pascal
c81103dfa6 fix linter comments 2026-08-05 16:21:01 +02:00
Dmitri Dolguikh
ba574dc739 added tests for GetPrivateServicesViaPgxConnection and GetPrivateServicesViaPgxConnection
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 16:06:43 +02:00
Dmitri Dolguikh
067982f772 Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 15:26:38 +02:00
Dmitri Dolguikh
ce4f0a7821 added GetAccountSettingsViaPgxConnection test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 15:26:06 +02:00
pascal
3218e2f744 fix client tests 2026-08-05 15:18:19 +02:00
pascal
7e24c27ef1 Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types 2026-08-05 15:09:41 +02:00
pascal
06817271bf fix embedded client test 2026-08-05 15:09:29 +02:00
Dmitri Dolguikh
1530739b54 Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 14:58:36 +02:00
Dmitri Dolguikh
e745aedc92 added GetRoutesViaPgxConnection test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 14:58:09 +02:00
pascal
4cf3903c83 fix management <-> shared dependencies 2026-08-05 14:48:27 +02:00
pascal
1e14b554a6 fix management <-> shared dependencies 2026-08-05 14:48:18 +02:00
Dmitri Dolguikh
11a3627899 Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 14:29:43 +02:00
Dmitri Dolguikh
2442bd6198 tests for GetNameServerGroupsViaPgxConnection and GetPostureChecksViaPgxConnection
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 14:29:07 +02:00
pascal
006cee000f Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types 2026-08-05 13:56:46 +02:00
pascal
c93aa03c0e merge main 2026-08-05 13:56:35 +02:00
Dmitri Dolguikh
efe2eaeb09 added GetDomains test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 12:37:58 +02:00
Dmitri Dolguikh
1926e983fb added GetDnsSettings test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 12:03:36 +02:00
Dmitri Dolguikh
e9b8175915 added GetNetworks test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 11:53:40 +02:00
Dmitri Dolguikh
88f1930450 added account network test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 11:40:23 +02:00
Dmitri Dolguikh
1774a3d9dd add network resource query test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 10:52:39 +02:00
Dmitri Dolguikh
795e06ce83 Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types 2026-08-05 10:16:54 +02:00
Dmitri Dolguikh
748f6b3fbb fixes + tests
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-05 10:16:21 +02:00
pascal
d5adfc799f handle policy with no rules to avoid index panic 2026-08-04 20:57:57 +02:00
pascal
60fe56e442 fix assert equal ordering 2026-08-04 20:56:16 +02:00
pascal
327fa4c0d6 fix group clone 2026-08-04 20:52:16 +02:00
pascal
ed015a7972 skip resource from proto failure 2026-08-04 20:46:10 +02:00
pascal
b43708e31b silently skip resource to proto failure 2026-08-04 20:42:43 +02:00
pascal
ce55b8406b handle not found peer on posture vali 2026-08-04 20:34:30 +02:00
pascal
0456b64ea1 silent skip on non existing groups references 2026-08-04 20:28:06 +02:00
Maycon Santos
6526fc2bec [management] Prevent deleting groups referenced by reverse proxy services (#7062)
## Describe your changes

A group could be deleted while a reverse proxy service still referenced
it, silently breaking the service's access control: private services
list groups in `access_groups` as the peer allowlist, and SSO bearer
auth distributes tokens to `distribution_groups`.

Group deletion now runs through the same linkage validation as routes,
policies, and agent network policies: deleting a group that backs a
private service allowlist or an enabled bearer-auth distribution list
fails with a `GroupLinkError` naming the service domain. Disabled bearer
configs and stale `access_groups` on non-private services are inert and
do not block deletion.

Tests cover both linked cases in single and bulk deletion, and pin the
non-blocking cases. The test account seeds decoy services ahead of the
linked ones so the check is proven to scan the full service list.
2026-08-05 03:24:20 +09:00
pascal
942ee81ec0 add benchmark 2026-08-04 20:19:17 +02:00
Dmitri Dolguikh
33a0e1bc2b adding integration tests
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-04 18:58:29 +02:00
Misha Bragin
2afa69b622 [management] prevent dangling group refs in agent-network ACLs. (#7060)
Block deleting a group referenced as a source group by an agent network
   policy, and drop unresolvable groups from synthesised private-service
ACLs. A deleted group survived in agent_network_policies.source_groups
   and was carried into the injected in-memory policy, where network-map
   assembly resolved it to a nil group and panicked on every proxy peer
   sync.
2026-08-04 18:04:22 +02:00
Zoltan Papp
2a61eac047 [client] Fix Linux tray right-click opening the main window (#7039)
## Describe your changes

On SNI hosts that report icon clicks via Activate (KDE Plasma, Waybar),
Wails fired the left-click handler on every dbusmenu 'opened' event, so
a right click opened the tray menu and immediately raised the main
window, which stole focus and closed the menu.

	Host                           Left click              Right click
KDE Plasma, Waybar main window (Activate) menu (host-rendered)
	GNOME Shell + AppIndicator     menu only               menu only
	Minimal WMs via XEmbed host    main window (Activate)  XEmbed GTK popup

Point the wails/v3 replace at the netbirdio fork (v3.0.0-beta.3 plus the
fix): once the host has sent Activate/SecondaryActivate, a menu open no
longer fires the click handler, while AppIndicator-only hosts that
signal clicks solely via 'opened' keep the old behavior.


## Issue ticket number and link

<!--
Required for anything that changes behavior. Link the issue (or the
validated
discussion it came from) that the NetBird team already agreed on. See

https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second
-->

## Stack

<!-- branch-stack -->

### Checklist
- [x] Is it a bug fix
- [ ] Is a typo/documentation fix
- [ ] Is a feature enhancement
- [ ] It is a refactor
- [ ] Created tests that fail without the change (if possible)
- [ ] I ran and tested this change locally — I did not rely on CI to
find out whether it works
- [ ] This PR has a single purpose (not a fix + refactor + feature in
one)
- [ ] This change is a trivial fix, **OR** it links an issue the NetBird
team agreed on beforehand. Changes to the public API, gRPC protocols,
functionality behavior, CLI / service flags, or new features always need
that agreement first. See
[CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second).

> By submitting this pull request, you confirm that you have read and
agree to the terms of the [Contributor License
Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md).

## Documentation
Select exactly one:

- [ ] I added/updated documentation for this change
- [x] Documentation is **not needed** for this change (explain why)

### Docs PR URL (required if "docs added" is checked)
Paste the PR link from https://github.com/netbirdio/docs here:

https://github.com/netbirdio/docs/pull/__


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

## Chores

* Updated underlying application components to support compatibility and
ongoing maintenance.
* Clarified Linux tray interaction documentation, including
platform-specific left- and right-click behavior and menu activation.
* No user-facing features, workflow changes, or visual updates are
included.
* Existing Linux tray behavior remains unchanged.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
2026-08-04 17:36:04 +02:00
Dmitri Dolguikh
dbfdd04c7b adding network-routers integration query test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-04 16:43:17 +02:00
Zoltan Papp
f2d13b884a [client] Fix session expired relogin (#7055)
## Describe your changes

After the SSO session expires, the daemon tears the engine down
permanently
(management returns `PermissionDenied` → `runCancel()` → the retry loop
exits
for good). The "Session expired" dialog's Login button still drove the
extend-session flow, which requires a live engine: the user completed
the full
browser SSO + 2FA round trip only to get
`Failed to extend the session — engine is not initialised`, with no way
out
other than quitting and relaunching the client.

Reproduce:
1. Log in on a desktop client with session expiration enabled (e.g. 16h
TTL).
2. Let the session expire (e.g. leave the machine asleep overnight).
3. Wake it, click **Login** on the "Session expired" dialog, complete
SSO + 2FA.
4. The error dialog appears and every retry fails the same way.

Changes:
- The expired branch of the session-expiration dialog now emits
`trigger-login`, driving the full `Login → SSO → Up` sequence that
rebuilds
the client, instead of the extend flow (an expired session can no longer
be
  extended).
- `RequestExtendAuthSession` fails fast when the engine is already gone,
so the
  browser/2FA round trip is not wasted on a doomed extend.
- The expired tray row navigated the main window to `/#/login`, a route
that
does not exist and fell through to the main page without starting a
login;
  it now emits `trigger-login` as well.


## Issue ticket number and link

<!--
Required for anything that changes behavior. Link the issue (or the
validated
discussion it came from) that the NetBird team already agreed on. See

https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second
-->

## Stack

<!-- branch-stack -->

### Checklist
- [x] Is it a bug fix
- [ ] Is a typo/documentation fix
- [ ] Is a feature enhancement
- [ ] It is a refactor
- [ ] Created tests that fail without the change (if possible)
- [ ] I ran and tested this change locally — I did not rely on CI to
find out whether it works
- [ ] This PR has a single purpose (not a fix + refactor + feature in
one)
- [ ] This change is a trivial fix, **OR** it links an issue the NetBird
team agreed on beforehand. Changes to the public API, gRPC protocols,
functionality behavior, CLI / service flags, or new features always need
that agreement first. See
[CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second).

> By submitting this pull request, you confirm that you have read and
agree to the terms of the [Contributor License
Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md).

## Documentation
Select exactly one:

- [ ] I added/updated documentation for this change
- [x] Documentation is **not needed** for this change (explain why)

### Docs PR URL (required if "docs added" is checked)
Paste the PR link from https://github.com/netbirdio/docs here:

https://github.com/netbirdio/docs/pull/__


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

- **Bug Fixes**
- Improved session extension handling when the client engine is
unavailable by prompting users to log in again.
- Updated expired-session behavior to trigger the standard login flow,
providing a more consistent sign-in experience.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
2026-08-04 16:02:31 +02:00
Zoltan Papp
564595d283 [client, android] Fix profile account path test on Windows (#7057)
## Describe your changes

The expected paths were hardcoded with Unix separators while
profileAccountPathFor builds the result with filepath.Join, so the
comparison failed on Windows. Derive the expectations with
filepath.FromSlash to keep the test platform-independent.


## Issue ticket number and link

<!--
Required for anything that changes behavior. Link the issue (or the
validated
discussion it came from) that the NetBird team already agreed on. See

https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second
-->

## Stack

<!-- branch-stack -->

### Checklist
- [x] Is it a bug fix
- [ ] Is a typo/documentation fix
- [ ] Is a feature enhancement
- [ ] It is a refactor
- [ ] Created tests that fail without the change (if possible)
- [ ] I ran and tested this change locally — I did not rely on CI to
find out whether it works
- [ ] This PR has a single purpose (not a fix + refactor + feature in
one)
- [ ] This change is a trivial fix, **OR** it links an issue the NetBird
team agreed on beforehand. Changes to the public API, gRPC protocols,
functionality behavior, CLI / service flags, or new features always need
that agreement first. See
[CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second).

> By submitting this pull request, you confirm that you have read and
agree to the terms of the [Contributor License
Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md).

## Documentation
Select exactly one:

- [ ] I added/updated documentation for this change
- [x] Documentation is **not needed** for this change (explain why)

### Docs PR URL (required if "docs added" is checked)
Paste the PR link from https://github.com/netbirdio/docs here:

https://github.com/netbirdio/docs/pull/__


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->

## Summary by CodeRabbit

* **Tests**
* Updated profile account path test expectations to use
platform-appropriate path separators, improving test reliability across
operating systems.

<!-- end of auto-generated comment: release notes by coderabbit.ai -->
2026-08-04 16:02:09 +02:00
Dmitri Dolguikh
4d010f60ce fix another import
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-04 15:40:28 +02:00
Zoltan Papp
78c1c2fc32 [client] Probe the daemon login with IsLoginRequired (#7052)
## Describe your changes

Probe the daemon login with IsLoginRequired 

The Login probe attemptLogin(ctx, "", "") on an unregistered peer ends
in registerPeer with no setup key and no JWT, which fails locally with
InvalidArgument before reaching Management. Since #6983 classified that
as StatusLoginFailed and returned early, every setup-key enrolment and
every expired-session SSO re-login aborted before using its credentials,
breaking all netbird-cloud e2e runs from commit e90be36cd.

IsLoginRequired asks the question the probe actually means - is the
peer's key alone still accepted - and reports Management's refusal as a
decision (needsLogin) instead of an error, the same pattern
foregroundLogin, Android and iOS already use.

## Issue ticket number and link

<!--
Required for anything that changes behavior. Link the issue (or the
validated
discussion it came from) that the NetBird team already agreed on. See

https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second
-->

## Stack

<!-- branch-stack -->

### Checklist
- [x] Is it a bug fix
- [ ] Is a typo/documentation fix
- [ ] Is a feature enhancement
- [ ] It is a refactor
- [ ] Created tests that fail without the change (if possible)
- [ ] I ran and tested this change locally — I did not rely on CI to
find out whether it works
- [ ] This PR has a single purpose (not a fix + refactor + feature in
one)
- [ ] This change is a trivial fix, **OR** it links an issue the NetBird
team agreed on beforehand. Changes to the public API, gRPC protocols,
functionality behavior, CLI / service flags, or new features always need
that agreement first. See
[CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second).

> By submitting this pull request, you confirm that you have read and
agree to the terms of the [Contributor License
Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md).

## Documentation
Select exactly one:

- [ ] I added/updated documentation for this change
- [x] Documentation is **not needed** for this change (explain why)

### Docs PR URL (required if "docs added" is checked)
Paste the PR link from https://github.com/netbirdio/docs here:

https://github.com/netbirdio/docs/pull/__


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->

## Summary by CodeRabbit

* **Bug Fixes**
  * Improved login handling when Management connectivity checks fail.
* Prevented unnecessary SSO prompts for already-authenticated sessions.
* Preserved setup-key login behavior while ensuring authentication
attempts proceed correctly.
* Login failures now return a clear failure status when authentication
state cannot be verified.

<!-- end of auto-generated comment: release notes by coderabbit.ai -->
2026-08-04 14:27:47 +02:00
Dmitri Dolguikh
3272058e56 fix extra setting manager package
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-04 13:34:25 +02:00
Dmitri Dolguikh
70c3feb05b Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-04 13:04:51 +02:00
Dmitri Dolguikh
a0fe80cd99 wire up validated peers
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-04 12:52:09 +02:00
Maycon Santos
bc7a15ab71 [management] Align agent-network API contracts for API clients (#7026)
## Describe your changes

Work on the Terraform provider (terraform-provider-netbird #177–#183)
surfaced places where the agent-network API broke its own contracts or
deviated from the conventions the rest of the management API follows,
forcing client-side workarounds.

Settings reads now follow the settings-endpoint convention: GET always
answers with a JSON object. Before bootstrap it returns the defaults
with an empty cluster/subdomain/endpoint (previously 200 with a JSON
`null` body, while the spec said 404). The settings PUT can bootstrap
the account by carrying a `cluster` — previously the row could only come
into existence through the first provider create, and a settings-first
setup was impossible; a differing cluster on a bootstrapped account is
rejected instead of silently ignored. PUT remains full-state.

The provider PUT schema promised omit-preserves semantics for several
operator-editable fields that the handler never delivered (it builds the
row from the request, like every other update handler). The schema
wording now matches the shipped full-state behavior; only the api_key
(secret) and session keys stay preserved by the manager. Identity
headers are always present in provider responses so an explicitly
cleared value round-trips as an empty string.

The Go REST client gains the full agent-network surface (catalog,
providers, policies, guardrails, budget rules, settings), including a
shim translating the legacy 200+`null` settings body from older servers
into an `IsNotFound` error.

Note for reviewers: the dashboard special-cased the `null` settings
body; it needs a small follow-up for the new defaults response (in
progress).
2026-08-04 01:59:09 +02:00
pascal
6fadf8f24a fix nmdata store expanding router peer groups alongside static peer 2026-08-03 23:58:47 +02:00
pascal
ef0032685b fix nmdata store shipping zones for disabled/public services 2026-08-03 23:54:22 +02:00
Zoltan Papp
530021aec6 [client] Update wails to v3.0.0-beta.3 (#7038)
## Describe your changes

Update wails to v3.0.0-beta.3

## Issue ticket number and link

<!--
Required for anything that changes behavior. Link the issue (or the
validated
discussion it came from) that the NetBird team already agreed on. See

https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second
-->

## Stack

<!-- branch-stack -->

### Checklist
- [ ] Is it a bug fix
- [ ] Is a typo/documentation fix
- [x] Is a feature enhancement
- [ ] It is a refactor
- [ ] Created tests that fail without the change (if possible)
- [ ] I ran and tested this change locally — I did not rely on CI to
find out whether it works
- [ ] This PR has a single purpose (not a fix + refactor + feature in
one)
- [ ] This change is a trivial fix, **OR** it links an issue the NetBird
team agreed on beforehand. Changes to the public API, gRPC protocols,
functionality behavior, CLI / service flags, or new features always need
that agreement first. See
[CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#ticket-first-pr-second).

> By submitting this pull request, you confirm that you have read and
agree to the terms of the [Contributor License
Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md).

## Documentation
Select exactly one:

- [ ] I added/updated documentation for this change
- [x] Documentation is **not needed** for this change (explain why)

### Docs PR URL (required if "docs added" is checked)
Paste the PR link from https://github.com/netbirdio/docs here:

https://github.com/netbirdio/docs/pull/__


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->

## Summary by CodeRabbit

* **Chores**
* Updated the application framework dependency to a newer beta release
for improved compatibility and stability.

<!-- end of auto-generated comment: release notes by coderabbit.ai -->
2026-08-03 23:19:59 +02:00
pascal
5ed38569f7 fix PrivateServiceCandidates nil assumption 2026-08-03 23:19:55 +02:00
pascal
d4e1c8978e add all group to user group lookup 2026-08-03 23:08:43 +02:00
Zoltan Papp
1bedb4e59d [client, android] Reuse the profile's account for Android SSO logins (#6988)
The Android binding never recorded which account a profile belongs to,
so every interactive login and every session extend went to the IdP with
no login_hint. With nothing to go on the IdP picks an account itself,
which on a session extend means re-authenticating an account the profile
is already signed in with.

Store the email the PKCE flow already parses out of the ID token, and
pass it back as the hint on later flows. An empty hint stays meaningful:
a fresh profile, or one that was logged out, deliberately leaves the
choice to the IdP, which is how a profile changes accounts. Logout
clears the stored email for that reason — while it is on disk it would
steer the next login straight back into the account just logged out of.

The email is keyed off the profile's config path rather than the active
profile: Auth.login runs in a goroutine, so the active profile can
change under a flow already in flight. It lands in
<profile>.account.json, not the <profile>.state.json desktop uses for
the same data — there the email and the engine's state manager sit in
different directories, but on Android both resolve under files/, and the
state manager rewrites the whole file from its own keys.

## Describe your changes

## Issue ticket number and link

## Stack

<!-- branch-stack -->

### Checklist
- [x] Is it a bug fix
- [ ] Is a typo/documentation fix
- [ ] Is a feature enhancement
- [ ] It is a refactor
- [ ] Created tests that fail without the change (if possible)
- [ ] This change does **not** modify the public API, gRPC protocols,
functionality behavior, CLI / service flags, or introduce a new feature
— **OR** I have discussed it with the NetBird team beforehand (link the
issue / Slack thread in the description). See
[CONTRIBUTING.md](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTING.md#discuss-changes-with-the-netbird-team-first).

> By submitting this pull request, you confirm that you have read and
agree to the terms of the [Contributor License
Agreement](https://github.com/netbirdio/netbird/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT.md).

## Documentation
Select exactly one:

- [ ] I added/updated documentation for this change
- [x] Documentation is **not needed** for this change (explain why)

### Docs PR URL (required if "docs added" is checked)
Paste the PR link from https://github.com/netbirdio/docs here:

https://github.com/netbirdio/docs/pull/__


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Added account email and active-status details to Android profile
information.
* Improved SSO sign-in and session renewal by restoring the previously
used account as a login hint.
* Added Android-specific profile email persistence with automatic
cleanup on logout.

* **Bug Fixes**
* Profile email persistence failures now generate warnings without
blocking login or logout.
* Improved handling of missing or unreadable account data and repeated
logout cleanup.

* **Tests**
* Added coverage for account-file naming, email persistence, and logout
behavior.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
2026-08-03 22:50:57 +02:00
pascal
d5ac70d806 fix missing ForceRoutingPeerDNSResolution on the nmdata store path 2026-08-03 22:50:06 +02:00
pascal
23579e4dd4 Skip the left-join NULL row instead of failing the account load + drop unsupported record types silently to match buildAppliedZoneCandidates 2026-08-03 22:38:12 +02:00
pascal
a00c5164a8 continue on resource policy loop + exclude disabled policies form index 2026-08-03 22:29:47 +02:00
pascal
02138dffdd fix group id mapping 2026-08-03 22:24:28 +02:00
pascal
02ec1f5dcb fix equivalence test 2026-08-03 22:24:11 +02:00
pascal
07d0440e34 fix slice initialization to avoid nil pointer dereference 2026-08-03 21:51:48 +02:00
pascal
2c85d94c6c hookup validated peers 2026-08-03 21:47:41 +02:00
Dmitri Dolguikh
1f190a50cf Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 18:22:34 +02:00
Dmitri Dolguikh
52fdfd5bdc add an posture-check-id to public-id index
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 18:21:07 +02:00
Dmitri Dolguikh
e35a0f3318 added PrivateServiceCandidates
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 18:04:11 +02:00
pascal
56c411babd fixed networks query and extended error wrapping 2026-08-03 17:52:33 +02:00
Dmitri Dolguikh
164b2baa4d Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 17:29:42 +02:00
Dmitri Dolguikh
0b29c6ed1a adding buildPrivateServiceCandidates
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 17:29:09 +02:00
pascal
724b61440b use new path on initial sync 2026-08-03 17:27:58 +02:00
pascal
22590ad66e hookup network map store 2026-08-03 17:27:24 +02:00
Dmitri Dolguikh
b8e004ea89 also build proxy-cluster to peer idx
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 14:58:53 +02:00
Dmitri Dolguikh
e2797360f4 support for applied zone candidates
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 14:32:54 +02:00
Dmitri Dolguikh
2f399f1e6e wire up dnssettings
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 14:12:27 +02:00
Dmitri Dolguikh
82c1e18264 renamed allowed_user_ids.go to user.got 2026-08-03 13:54:22 +02:00
Dmitri Dolguikh
ce9023bd27 wire up accountsettings
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 13:51:25 +02:00
Dmitri Dolguikh
7d33356776 support for group to user ids
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 13:47:14 +02:00
Dmitri Dolguikh
993291149c Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 12:20:18 +02:00
Dmitri Dolguikh
98f7ea40a1 added allowed_user_ids call
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 12:19:50 +02:00
pascal
76db9ab94f update integrated validator 2026-08-03 12:11:35 +02:00
Dmitri Dolguikh
9653b15d78 automatically fill empty public ids for view types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 11:34:22 +02:00
Dmitri Dolguikh
206bb1676b build resourcePolicies map
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-03 10:53:54 +02:00
Dmitri Dolguikh
600b0c752b also return a net resource to groups index from GetGroups call
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-01 11:56:17 +02:00
Dmitri Dolguikh
b02736adc3 cleanup network routers retrieval
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-08-01 11:54:14 +02:00
Dmitri Dolguikh
5d6117d2c0 adding network_map_data: compute Routers field
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-31 17:56:52 +02:00
Dmitri Dolguikh
c2f8360b00 adding support for bulding of network_map_data
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-31 15:58:40 +02:00
Dmitri Dolguikh
ea1b4d56e8 added support for queries via pgx connection
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-31 15:03:15 +02:00
Dmitri Dolguikh
acbe22b831 small cleanups
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-31 13:41:34 +02:00
Dmitri Dolguikh
3b5c8e2298 started moving pg-specific nmap tests to integration_tests/management/network_map_db/pgsql
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-31 12:29:35 +02:00
Dmitri Dolguikh
2baeb4bc0d clean up sql types conversion test
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-30 19:35:11 +02:00
Dmitri Dolguikh
4bf75fdc97 support for posture-checks
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-30 18:00:58 +02:00
Dmitri Dolguikh
799f3a3c62 Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-30 16:03:46 +02:00
Dmitri Dolguikh
129736ad61 support for account settings
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-30 16:03:19 +02:00
pascal
f7be9c4347 legacy network mal and equivalence test 2026-07-30 15:45:15 +02:00
Dmitri Dolguikh
e05cb5264d support for dns custom zones
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-30 15:34:09 +02:00
Dmitri Dolguikh
5a10561ca1 fix handling of string slices
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-30 11:31:53 +02:00
Dmitri Dolguikh
42ce83a8f3 added support for account network
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-29 18:57:08 +02:00
Dmitri Dolguikh
e620c86cd4 added support for networkrouters
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-29 18:32:51 +02:00
Dmitri Dolguikh
9dee2d60b9 added support for networkresources
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-29 18:19:34 +02:00
Dmitri Dolguikh
4525014632 support for nameservergroups
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-29 17:59:44 +02:00
Dmitri Dolguikh
23a5c0de4b added support for routes
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-29 17:37:00 +02:00
Dmitri Dolguikh
25e882004f added retrieval of policies
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-28 18:06:28 +02:00
Dmitri Dolguikh
15003258d2 networkmap read-only interface for pgsql
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-28 16:46:45 +02:00
Dmitri Dolguikh
2af3a5fba5 do not send resource policies map over the wire
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
2026-07-27 19:00:30 +02:00
pascal
a6603a2e0a have explicit network map data type and do calculation from there 2026-07-27 15:17:30 +02:00
pascal
7ed3737cda revert component types 2026-07-24 15:16:55 +02:00
186 changed files with 15752 additions and 4082 deletions

View File

@@ -82,6 +82,8 @@ type Client struct {
connectClient *internal.ConnectClient
config *profilemanager.Config
cacheDir string
// Identifies the running profile for the SSO login hint; see profile_state.go.
cfgPath string
stateChangeMu sync.Mutex
stateChangeSubID string
@@ -102,11 +104,12 @@ type Client struct {
extendCancel context.CancelFunc
}
func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cc *internal.ConnectClient) {
func (c *Client) setState(cfg *profilemanager.Config, cacheDir string, cfgPath string, cc *internal.ConnectClient) {
c.stateMu.Lock()
defer c.stateMu.Unlock()
c.config = cfg
c.cacheDir = cacheDir
c.cfgPath = cfgPath
c.connectClient = cc
}
@@ -116,6 +119,16 @@ func (c *Client) stateSnapshot() (*profilemanager.Config, string, *internal.Conn
return c.config, c.cacheDir, c.connectClient
}
// authSnapshot returns the config together with the path it was loaded from, in
// one lock: the path identifies the profile whose account email backs the login
// hint, so reading it separately could pair one profile's config with another's
// hint when a profile switch lands in between.
func (c *Client) authSnapshot() (*profilemanager.Config, string, *internal.ConnectClient) {
c.stateMu.RLock()
defer c.stateMu.RUnlock()
return c.config, c.cfgPath, c.connectClient
}
func (c *Client) getConnectClient() *internal.ConnectClient {
c.stateMu.RLock()
defer c.stateMu.RUnlock()
@@ -168,7 +181,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
defer c.ctxCancel()
c.ctxCancelLock.Unlock()
auth := NewAuthWithConfig(ctx, cfg)
auth := NewAuthWithConfig(ctx, cfg, cfgFile)
err = auth.login(urlOpener, isAndroidTV)
if err != nil {
return err
@@ -176,7 +189,7 @@ func (c *Client) Run(platformFiles PlatformFiles, urlOpener URLOpener, isAndroid
// todo do not throw error in case of cancelled context
ctx = internal.CtxInitState(ctx)
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder)
c.setState(cfg, cacheDir, connectClient)
c.setState(cfg, cacheDir, cfgFile, connectClient)
// 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()
@@ -217,7 +230,7 @@ func (c *Client) RunWithoutLogin(platformFiles PlatformFiles, dns *DNSList, dnsR
// todo do not throw error in case of cancelled context
ctx = internal.CtxInitState(ctx)
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder)
c.setState(cfg, cacheDir, connectClient)
c.setState(cfg, cacheDir, cfgFile, connectClient)
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
}

View File

@@ -4,6 +4,8 @@ import (
"context"
"fmt"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/system"
@@ -61,11 +63,14 @@ func NewAuth(cfgPath string, mgmURL string) (*Auth, error) {
}, nil
}
// NewAuthWithConfig instantiate Auth based on existing config
func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config) *Auth {
// NewAuthWithConfig instantiate Auth based on existing config. cfgPath is the
// file the config was loaded from; it identifies the profile whose account email
// backs the login_hint.
func NewAuthWithConfig(ctx context.Context, config *profilemanager.Config, cfgPath string) *Auth {
return &Auth{
ctx: ctx,
config: config,
ctx: ctx,
config: config,
cfgPath: cfgPath,
}
}
@@ -158,12 +163,14 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
}
jwtToken := ""
email := ""
if needsLogin {
tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV)
if err != nil {
return fmt.Errorf("interactive sso login failed: %v", err)
}
jwtToken = tokenInfo.GetTokenToUse()
email = tokenInfo.Email
}
err, _ = authClient.Login(a.ctx, "", jwtToken)
@@ -171,17 +178,42 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
return fmt.Errorf("login failed: %v", err)
}
// Stored after Login, not before: a rejected token must not leave a hint
// pointing at an account that cannot be used.
if email != "" && a.cfgPath != "" {
if err := writeProfileEmail(a.cfgPath, email); err != nil {
log.Warnf("failed to store profile account email: %v", err)
}
}
go urlOpener.OnLoginSuccess()
return nil
}
// loginHintSetter is implemented by both concrete flows (PKCE and device code)
// but absent from the OAuthFlow interface, hence the assertion below — the same
// way internal/auth wires it in authenticateWithPKCEFlow.
type loginHintSetter interface {
SetLoginHint(hint string)
}
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, isAndroidTV bool) (*auth.TokenInfo, error) {
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV)
if err != nil {
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
}
// An empty hint is deliberate, not a fallback: a fresh or logged-out profile
// leaves the choice to the IdP, which is how accounts get switched.
if a.cfgPath != "" {
if hint := readProfileEmail(a.cfgPath); hint != "" {
if setter, ok := oAuthFlow.(loginHintSetter); ok {
setter.SetLoginHint(hint)
}
}
}
flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO())
if err != nil {
return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err)

View File

@@ -13,18 +13,17 @@ import (
)
const (
// Android-specific config filename (different from desktop default.json)
defaultConfigFilename = "netbird.cfg"
// Subdirectory for non-default profiles (must match Java Preferences.java)
profilesSubdir = "profiles"
// Android uses a single user context per app (non-empty username required by ServiceManager)
androidUsername = "android"
)
// Profile represents a profile for gomobile
type Profile struct {
ID string
Name string
ID string
Name string
// Email is the account this profile last logged in with, "" if it never
// completed an SSO login or was logged out. See profile_state.go.
Email string
IsActive bool
}
@@ -101,6 +100,7 @@ func (pm *ProfileManager) ListProfiles() (*ProfileArray, error) {
profiles = append(profiles, &Profile{
ID: p.ID.String(),
Name: p.Name,
Email: pm.profileEmail(p.ID.String()),
IsActive: p.IsActive,
})
}
@@ -123,7 +123,22 @@ func (pm *ProfileManager) GetActiveProfile() (*Profile, error) {
if err != nil {
return nil, fmt.Errorf("failed to resolve active profile %q: %w", activeState.ID, err)
}
return &Profile{ID: prof.ID.String(), Name: prof.Name, IsActive: true}, nil
return &Profile{
ID: prof.ID.String(),
Name: prof.Name,
Email: pm.profileEmail(prof.ID.String()),
IsActive: true,
}, nil
}
// profileEmail returns the account email recorded for a profile. Display-only, so
// an unresolvable path degrades to "" rather than an error.
func (pm *ProfileManager) profileEmail(id string) string {
configPath, err := pm.getProfileConfigPath(id)
if err != nil {
return ""
}
return readProfileEmail(configPath)
}
// SwitchProfile switches to a different profile
@@ -185,6 +200,11 @@ func (pm *ProfileManager) LogoutProfile(id string) error {
return fmt.Errorf("failed to save config: %w", err)
}
// Not fatal: a stale hint costs an account switch, not the logout itself.
if err := removeProfileEmail(configPath); err != nil {
log.Warnf("failed to clear stored account email for profile %s: %v", id, err)
}
log.Infof("logged out from profile: %s", id)
return nil
}

View File

@@ -0,0 +1,108 @@
package android
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/util"
)
const (
// Android-specific config filename (different from desktop default.json)
defaultConfigFilename = "netbird.cfg"
// Subdirectory for non-default profiles (must match Java Preferences.java)
profilesSubdir = "profiles"
// profileAccountSuffix names the file holding the profile's account email.
// Deliberately not ".state.json", which desktop uses for the same data:
// there the email and the engine's state manager live in different
// directories, but on Android both resolve under files/, so sharing the name
// would have the two overwrite each other — the state manager rewrites the
// whole file from its own keys (see statemanager.Manager.PersistState), and
// this package's writer does the same in reverse.
profileAccountSuffix = ".account.json"
)
// profileAccountPathFor derives the account file path from a profile's config
// path: netbird.cfg -> netbird.account.json, <id>.json -> <id>.account.json.
//
// Deriving from the config path rather than resolving the active profile keeps
// the write on the profile the login actually ran for: Auth.login runs in a
// goroutine, so the active profile can change under a flow already in flight.
func profileAccountPathFor(configPath string) (string, error) {
if configPath == "" {
return "", fmt.Errorf("empty config path")
}
base := filepath.Base(configPath)
stem := strings.TrimSuffix(base, filepath.Ext(base))
if stem == "" || stem == "." {
return "", fmt.Errorf("config path %q has no filename stem", configPath)
}
return filepath.Join(filepath.Dir(configPath), stem+profileAccountSuffix), nil
}
// readProfileEmail returns the account email stored for the profile whose config
// lives at configPath. A missing or unreadable file yields "", which leaves the
// account choice to the IdP.
func readProfileEmail(configPath string) string {
accountPath, err := profileAccountPathFor(configPath)
if err != nil {
log.Debugf("no profile account path for login hint: %v", err)
return ""
}
var state profilemanager.ProfileState
if _, err := util.ReadJson(accountPath, &state); err != nil {
if !os.IsNotExist(err) {
log.Debugf("failed to read profile account for login hint: %v", err)
}
return ""
}
return state.Email
}
// writeProfileEmail records the account email for the profile whose config lives
// at configPath, so later logins can pass it as an OIDC login_hint. An empty
// email is ignored rather than blanking what is already stored.
func writeProfileEmail(configPath string, email string) error {
if email == "" {
return nil
}
accountPath, err := profileAccountPathFor(configPath)
if err != nil {
return fmt.Errorf("resolve profile account path: %w", err)
}
state := profilemanager.ProfileState{Email: email}
if err := util.WriteJsonWithRestrictedPermission(context.Background(), accountPath, state); err != nil {
return fmt.Errorf("write profile account: %w", err)
}
return nil
}
// removeProfileEmail drops the stored account email. Called on logout: while the
// email is on disk it goes out as a login_hint, which would steer the next login
// straight back into the account just logged out of. Mirrors the desktop UI's
// RemoveProfileState call.
func removeProfileEmail(configPath string) error {
accountPath, err := profileAccountPathFor(configPath)
if err != nil {
return fmt.Errorf("resolve profile account path: %w", err)
}
if err := os.Remove(accountPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("remove profile account: %w", err)
}
return nil
}

View File

@@ -0,0 +1,161 @@
package android
import (
"os"
"path/filepath"
"testing"
)
func TestProfileAccountPathFor(t *testing.T) {
tests := []struct {
name string
configPath string
want string
wantErr bool
}{
{
name: "default profile",
configPath: "/data/data/io.netbird.client/files/netbird.cfg",
want: filepath.FromSlash("/data/data/io.netbird.client/files/netbird.account.json"),
},
{
name: "id profile",
configPath: "/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.json",
want: filepath.FromSlash("/data/data/io.netbird.client/files/profiles/4c5f5c8198c3989cffb5b5394f5a7ae0.account.json"),
},
{
name: "legacy name-keyed profile is handled the same way",
configPath: "/data/data/io.netbird.client/files/profiles/work.json",
want: filepath.FromSlash("/data/data/io.netbird.client/files/profiles/work.account.json"),
},
{
name: "empty path is rejected",
configPath: "",
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := profileAccountPathFor(tt.configPath)
if tt.wantErr {
if err == nil {
t.Fatalf("expected an error, got path %q", got)
}
return
}
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != tt.want {
t.Errorf("got %q, want %q", got, tt.want)
}
})
}
}
func TestProfileAccountPathForDefaultDoesNotCollide(t *testing.T) {
root := "/data/data/io.netbird.client/files"
defaultAccount, err := profileAccountPathFor(filepath.Join(root, defaultConfigFilename))
if err != nil {
t.Fatalf("default profile: %v", err)
}
idAccount, err := profileAccountPathFor(filepath.Join(root, profilesSubdir, "abc123.json"))
if err != nil {
t.Fatalf("id profile: %v", err)
}
if defaultAccount == idAccount {
t.Fatalf("default and id profile share an account file: %q", defaultAccount)
}
}
// The account file must never land on the engine state file: on Android both
// resolve under files/, and the state manager rewrites the whole file from its
// own keys, so sharing a path would have the two overwrite each other. The
// expected names here mirror ProfileManager.GetStateFilePath.
func TestProfileAccountPathAvoidsEngineStateFile(t *testing.T) {
root := "/data/data/io.netbird.client/files"
cases := []struct {
configPath string
engineState string
}{
{
configPath: filepath.Join(root, defaultConfigFilename),
engineState: filepath.Join(root, "state.json"),
},
{
configPath: filepath.Join(root, profilesSubdir, "abc123.json"),
engineState: filepath.Join(root, profilesSubdir, "abc123.state.json"),
},
}
for _, c := range cases {
account, err := profileAccountPathFor(c.configPath)
if err != nil {
t.Fatalf("%s: %v", c.configPath, err)
}
if account == c.engineState {
t.Errorf("account file collides with the engine state file: %q", account)
}
}
}
func TestWriteThenReadProfileEmail(t *testing.T) {
configPath := filepath.Join(t.TempDir(), "profiles", "abc123.json")
if err := ensureDirFor(t, configPath); err != nil {
t.Fatalf("prepare dir: %v", err)
}
if got := readProfileEmail(configPath); got != "" {
t.Errorf("expected no email before a login, got %q", got)
}
const email = "user@example.com"
if err := writeProfileEmail(configPath, email); err != nil {
t.Fatalf("write: %v", err)
}
if got := readProfileEmail(configPath); got != email {
t.Errorf("got %q, want %q", got, email)
}
if err := removeProfileEmail(configPath); err != nil {
t.Fatalf("remove: %v", err)
}
if got := readProfileEmail(configPath); got != "" {
t.Errorf("expected no email after logout, got %q", got)
}
// Logout may run on a never-logged-in profile, so a second remove must pass.
if err := removeProfileEmail(configPath); err != nil {
t.Fatalf("second remove should be a no-op: %v", err)
}
}
func TestWriteProfileEmailIgnoresEmpty(t *testing.T) {
configPath := filepath.Join(t.TempDir(), "profiles", "abc123.json")
if err := ensureDirFor(t, configPath); err != nil {
t.Fatalf("prepare dir: %v", err)
}
const email = "user@example.com"
if err := writeProfileEmail(configPath, email); err != nil {
t.Fatalf("write: %v", err)
}
if err := writeProfileEmail(configPath, ""); err != nil {
t.Fatalf("write empty: %v", err)
}
if got := readProfileEmail(configPath); got != email {
t.Errorf("empty write clobbered the stored email: got %q, want %q", got, email)
}
}
func ensureDirFor(t *testing.T, path string) error {
t.Helper()
return os.MkdirAll(filepath.Dir(path), 0o700)
}

View File

@@ -278,7 +278,7 @@ func (c *Client) endExtend() {
}
func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isAndroidTV bool) error {
cfg, _, cc := c.stateSnapshot()
cfg, cfgPath, cc := c.authSnapshot()
if cfg == nil || cc == nil {
return fmt.Errorf("engine is not running")
}
@@ -293,7 +293,10 @@ func (c *Client) extendAuthSession(ctx context.Context, urlOpener URLOpener, isA
}
defer authClient.Close()
a := &Auth{ctx: ctx, config: cfg}
// Passing the config path makes the flow pick up the login_hint: an extend
// renews the session of the account already signed in, so it must not stop to
// offer a choice.
a := NewAuthWithConfig(ctx, cfg, cfgPath)
tokenInfo, err := a.foregroundGetTokenInfo(authClient, urlOpener, isAndroidTV)
if err != nil {
return fmt.Errorf("interactive sso login failed: %v", err)

View File

@@ -27,8 +27,8 @@ import (
const errCloseConnection = "Failed to close connection: %v"
var (
logFileCount uint32
systemInfoFlag bool
logFileCount uint32
systemInfoFlag bool
uploadBundleFlag bool
uploadBundleURLFlag string
uploadBundleInsecureFlag bool

View File

@@ -124,7 +124,7 @@ func startManagement(t *testing.T, config *config.Config, testFile string) (*grp
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := mgmt.NewAccountRequestBuffer(ctx, store)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersmanager), config)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersmanager), config, nil)
accountManager, err := mgmt.BuildManager(ctx, config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManagerMock, false, cacheStore)
if err != nil {

View File

@@ -21,8 +21,8 @@ import (
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/proto"
nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/server"
"github.com/netbirdio/netbird/client/system"
"github.com/netbirdio/netbird/shared/management/domain"

View File

@@ -146,7 +146,7 @@ func startManagement(t *testing.T, signalAddr string) string {
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := mgmt.NewAccountRequestBuffer(context.Background(), testStore)
networkMapController := controller.NewController(context.Background(), testStore, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(testStore, peersManager), cfg)
networkMapController := controller.NewController(context.Background(), testStore, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(testStore, peersManager), cfg, nil)
accountManager, err := mgmt.BuildManager(context.Background(), cfg, testStore, networkMapController, jobManager, nil, "", eventStore, nil, false, iv, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
require.NoError(t, err)

View File

@@ -519,7 +519,7 @@ func startManagement(t *testing.T, dataDir, testFile string) (*grpc.Server, stri
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := server.NewAccountRequestBuffer(context.Background(), store)
networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config)
networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config, nil)
accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
if err != nil {
return nil, "", err

View File

@@ -8,8 +8,6 @@ import (
"testing"
"github.com/stretchr/testify/require"
"google.golang.org/grpc/codes"
gstatus "google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/proto"
@@ -27,9 +25,9 @@ func TestLogin_ManagementUnreachableIsReturnedInsteadOfDemandingSSO(t *testing.T
unreachable := errors.New("create connection: dial context: context deadline exceeded")
attempts := 0
s.loginAttemptFn = func(context.Context, string, string) (internal.StatusType, error) {
s.isLoginRequiredFn = func(context.Context) (bool, error) {
attempts++
return internal.StatusLoginFailed, unreachable
return false, unreachable
}
resp, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username})
@@ -55,15 +53,12 @@ func TestLogin_AuthRefusalStartsSSOFlow(t *testing.T) {
s.rootCtx = internal.CtxInitState(context.Background())
breakProfilePrivateKey(t, cfgPath)
refused := gstatus.Error(codes.PermissionDenied, "peer is not registered")
s.loginAttemptFn = func(context.Context, string, string) (internal.StatusType, error) {
return internal.StatusNeedsLogin, refused
s.isLoginRequiredFn = func(context.Context) (bool, error) {
return true, nil
}
_, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username})
require.Error(t, err)
require.NotErrorIs(t, err, refused,
"the refusal was handed back to the caller instead of starting the SSO flow")
status, stateErr := internal.CtxGetState(s.rootCtx).Status()
require.NoError(t, stateErr)
@@ -71,6 +66,32 @@ func TestLogin_AuthRefusalStartsSSOFlow(t *testing.T) {
"the SSO flow setup was never reached with the broken key")
}
func TestLogin_SetupKeyStillRunsWhenPeerNeedsLogin(t *testing.T) {
s, _, _, username, _ := setupServerWithProfile(t)
s.rootCtx = internal.CtxInitState(context.Background())
s.isLoginRequiredFn = func(context.Context) (bool, error) {
return true, nil
}
var keysTried []string
s.loginAttemptFn = func(_ context.Context, setupKey, _ string) (internal.StatusType, error) {
keysTried = append(keysTried, setupKey)
return "", nil
}
setupKey := "A2C8E32F-AEB2-4B45-8FD3-8A0C1B2D3E4F"
resp, err := s.Login(userCtx(), &proto.LoginRequest{Username: &username, SetupKey: setupKey})
require.NoError(t, err, "the probe's outcome leaked out as the login result")
require.NotNil(t, resp)
require.Equal(t, []string{setupKey}, keysTried, "the setup key never reached the login attempt")
require.Nil(t, s.oauthAuthFlow.flow, "a setup-key login started an SSO flow")
status, err := internal.CtxGetState(s.rootCtx).Status()
require.NoError(t, err)
require.Equal(t, internal.StatusIdle, status)
}
// breakProfilePrivateKey replaces the profile's private key with an unparseable
// one, which makes any attempt to build a Management client fail on the spot.
func breakProfilePrivateKey(t *testing.T, cfgPath string) {

View File

@@ -232,4 +232,3 @@ func toNetIDs(routes []string) []route.NetID {
}
return netIDs
}

View File

@@ -140,6 +140,8 @@ type Server struct {
// it to drive the login outcomes that need a server on the other end;
// production leaves it nil, and every login goes through loginAttempt.
loginAttemptFn func(ctx context.Context, setupKey, jwtToken string) (internal.StatusType, error)
isLoginRequiredFn func(ctx context.Context) (bool, error)
}
type oauthAuthFlow struct {
@@ -384,6 +386,21 @@ func (s *Server) attemptLogin(ctx context.Context, setupKey, jwtToken string) (i
return s.loginAttempt(ctx, setupKey, jwtToken)
}
func (s *Server) isLoginRequired(ctx context.Context) (bool, error) {
if s.isLoginRequiredFn != nil {
return s.isLoginRequiredFn(ctx)
}
authClient, err := auth.NewAuth(ctx, s.config.PrivateKey, s.config.ManagementURL, s.config)
if err != nil {
log.Errorf("failed to create auth client: %v", err)
return false, err
}
defer authClient.Close()
return authClient.IsLoginRequired(ctx)
}
// loginAttempt attempts to login using the provided information. It returns
// StatusNeedsLogin when Management refused the peer's credentials and
// StatusLoginFailed for every other failure, so callers can tell an
@@ -640,22 +657,22 @@ func (s *Server) Login(callerCtx context.Context, msg *proto.LoginRequest) (*pro
s.config = config
s.mutex.Unlock()
loginStatus, err := s.attemptLogin(ctx, "", "")
if err == nil {
state.Set(internal.StatusIdle)
return &proto.LoginResponse{}, nil
}
// Only an authentication refusal means the peer has to (re-)authenticate.
// Any other failure leaves the login undecided: Management unreachable, a
// A probe that errors leaves the login undecided: Management unreachable, a
// restart mid-request, an internal error. Those are returned for the caller
// to retry, because turning them into an SSO prompt asks the user to solve
// something that is not theirs to solve, and a browser login cannot succeed
// while Management is unreachable anyway.
if loginStatus != internal.StatusNeedsLogin {
state.Set(loginStatus)
// while Management is unreachable anyway. Only Management refusing the
// peer's key is a decision, and IsLoginRequired reports that as
// needsLogin=true rather than an error.
needsLogin, err := s.isLoginRequired(ctx)
if err != nil {
state.Set(internal.StatusLoginFailed)
return nil, err
}
if !needsLogin {
state.Set(internal.StatusIdle)
return &proto.LoginResponse{}, nil
}
if msg.SetupKey == "" {
hint := ""
@@ -1798,6 +1815,9 @@ func (s *Server) RequestExtendAuthSession(
if connectClient == nil {
return nil, gstatus.Errorf(codes.FailedPrecondition, "client is not running")
}
if connectClient.Engine() == nil {
return nil, gstatus.Errorf(codes.FailedPrecondition, "session can no longer be extended, log in again to reconnect")
}
hint := ""
if msg.Hint != nil {

View File

@@ -200,7 +200,7 @@ func startManagement(t *testing.T, signalAddr string, counter *int) (*grpc.Serve
requestBuffer := server.NewAccountRequestBuffer(context.Background(), store)
peersUpdateManager := update_channel.NewPeersUpdateManager(metrics)
networkMapController := controller.NewController(context.Background(), store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config)
networkMapController := controller.NewController(context.Background(), store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config, nil)
accountManager, err := server.BuildManager(context.Background(), config, store, networkMapController, jobManager, nil, "", eventStore, nil, false, ia, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManagerMock, false, cacheStore)
if err != nil {
return nil, "", err

View File

@@ -11,7 +11,7 @@ import { DialogHeading } from "@/components/dialog/DialogHeading";
import { SquareIcon } from "@/components/SquareIcon";
import { Connection, Profiles as ProfilesSvc, Session, WindowManager } from "@bindings/services";
import { useAutoSizeWindow } from "@/hooks/useAutoSizeWindow";
import { EVENT_BROWSER_LOGIN_CANCEL } from "@/lib/connection";
import { EVENT_BROWSER_LOGIN_CANCEL, EVENT_TRIGGER_LOGIN } from "@/lib/connection";
import { errorDialog, formatErrorMessage } from "@/lib/errors.ts";
import { formatRemaining } from "@/lib/formatters";
@@ -131,6 +131,21 @@ export default function SessionExpirationDialog() {
}
}, [busy, t]);
const authenticate = useCallback(async () => {
if (busy) return;
setBusy(true);
try {
await Events.Emit(EVENT_TRIGGER_LOGIN);
await WindowManager.CloseSessionExpiration();
} catch (e) {
setBusy(false);
await errorDialog({
Title: t("connect.error.loginTitle"),
Message: formatErrorMessage(e),
});
}
}, [busy, t]);
const logout = useCallback(async () => {
if (busy) return;
setBusy(true);
@@ -185,7 +200,7 @@ export default function SessionExpirationDialog() {
variant={"primary"}
size={"md"}
className={"w-full"}
onClick={stay}
onClick={expired ? authenticate : stay}
disabled={busy}
>
{expired ? t("sessionExpiration.authenticate") : t("sessionExpiration.stay")}

View File

@@ -95,10 +95,6 @@ func main() {
}
})
// Debug patch, not for release: dumps heap/goroutine profiles and the
// process tree to /tmp/nbgui for the memory consumption investigation.
startMemProfiler(app)
profiles := services.NewProfiles(conn)
// updater.Holder owns the typed update State; DaemonFeed feeds it and the
// Update service is a thin Wails-bound facade over it plus the install RPCs.

View File

@@ -1,319 +0,0 @@
//go:build !android && !ios && !freebsd && !js
package main
import (
"bufio"
"fmt"
"os"
"path/filepath"
"runtime"
"runtime/pprof"
"strconv"
"strings"
"time"
"github.com/shirou/gopsutil/v4/process"
log "github.com/sirupsen/logrus"
"github.com/wailsapp/wails/v3/pkg/application"
"github.com/wailsapp/wails/v3/pkg/events"
)
// memProfOffsets are the snapshot times measured from application startup.
var memProfOffsets = []time.Duration{0, 2 * time.Minute, 5 * time.Minute}
// memProfMaxDepth bounds the child walk so a cycle in the reported parent links
// cannot spin forever.
const memProfMaxDepth = 4
type memProfileSpec struct {
profile string
file string
debug int
}
var memProfileSpecs = []memProfileSpec{
{profile: "heap", file: "heap.pprof", debug: 0},
{profile: "heap", file: "heap.txt", debug: 1},
{profile: "goroutine", file: "goroutine.txt", debug: 1},
{profile: "threadcreate", file: "threadcreate.txt", debug: 1},
}
var memProfStart = time.Now()
// startMemProfiler dumps a profile snapshot at every memProfOffsets mark, each
// into its own timestamped directory under memProfBaseDir. The first runs once
// the application is up so the window inventory sees the eagerly created
// windows. Every failure is logged and never stops the GUI.
func startMemProfiler(app *application.App) {
log.Infof("memory profiler enabled, writing to %s (snapshots at %v after startup)", memProfBaseDir(), memProfOffsets)
app.Event.OnApplicationEvent(events.Common.ApplicationStarted, func(*application.ApplicationEvent) {
go func() {
started := time.Now()
for _, offset := range memProfOffsets {
if wait := time.Until(started.Add(offset)); wait > 0 {
time.Sleep(wait)
}
writeMemProfile(app)
}
}()
})
}
// memProfBaseDir returns the directory holding the snapshot directories.
func memProfBaseDir() string {
if runtime.GOOS == "windows" {
return filepath.Join(os.TempDir(), "nbgui")
}
return "/tmp/nbgui"
}
// writeMemProfile creates a <timestamp>-<pid> directory and fills it with the
// runtime profiles, the memory statistics summary and the process tree.
func writeMemProfile(app *application.App) {
name := fmt.Sprintf("%s-%d", time.Now().Format("20060102-150405"), os.Getpid())
dir := filepath.Join(memProfBaseDir(), name)
if err := os.MkdirAll(dir, 0o755); err != nil {
log.Warnf("create memory profile dir %s: %v", dir, err)
return
}
// The heap profile reports live objects as of the last collection, so force
// one to keep inuse_space from counting garbage that is already unreachable.
runtime.GC()
if err := writeMemStats(filepath.Join(dir, "memstats.txt"), app); err != nil {
log.Warnf("write memory statistics: %v", err)
}
if err := writeProcTree(filepath.Join(dir, "proctree.txt")); err != nil {
log.Warnf("write process tree: %v", err)
}
for _, spec := range memProfileSpecs {
if err := writeMemProfileFile(spec, filepath.Join(dir, spec.file)); err != nil {
log.Warnf("write %s profile: %v", spec.profile, err)
}
}
log.Infof("memory profile written to %s", dir)
}
// writeMemProfileFile writes a single runtime profile to path.
func writeMemProfileFile(spec memProfileSpec, path string) error {
p := pprof.Lookup(spec.profile)
if p == nil {
return fmt.Errorf("unknown profile %q", spec.profile)
}
f, err := os.Create(path)
if err != nil {
return fmt.Errorf("create %s: %w", path, err)
}
defer func() {
if err := f.Close(); err != nil {
log.Debugf("close %s: %v", path, err)
}
}()
if err := p.WriteTo(f, spec.debug); err != nil {
return fmt.Errorf("write %s: %w", path, err)
}
return nil
}
// writeMemStats dumps the runtime memory statistics next to the process
// resident set size. A resident set much larger than Sys means the memory sits
// outside the Go heap (webview, GTK, other cgo allocations), where the pprof
// profiles cannot see it.
func writeMemStats(path string, app *application.App) error {
var m runtime.MemStats
runtime.ReadMemStats(&m)
var b strings.Builder
fmt.Fprintf(&b, "time: %s\n", time.Now().Format(time.RFC3339))
fmt.Fprintf(&b, "uptime: %s\n", time.Since(memProfStart).Round(time.Second))
fmt.Fprintf(&b, "pid: %d\n", os.Getpid())
fmt.Fprintf(&b, "\n")
rss, vms := processMemory()
fmt.Fprintf(&b, "process_rss: %s\n", rss)
fmt.Fprintf(&b, "process_vms: %s\n", vms)
fmt.Fprintf(&b, "\n")
fmt.Fprintf(&b, "sys: %s\n", formatMemBytes(m.Sys))
fmt.Fprintf(&b, "heap_alloc: %s\n", formatMemBytes(m.HeapAlloc))
fmt.Fprintf(&b, "heap_sys: %s\n", formatMemBytes(m.HeapSys))
fmt.Fprintf(&b, "heap_inuse: %s\n", formatMemBytes(m.HeapInuse))
fmt.Fprintf(&b, "heap_idle: %s\n", formatMemBytes(m.HeapIdle))
fmt.Fprintf(&b, "heap_released: %s\n", formatMemBytes(m.HeapReleased))
fmt.Fprintf(&b, "heap_objects: %d\n", m.HeapObjects)
fmt.Fprintf(&b, "stack_inuse: %s\n", formatMemBytes(m.StackInuse))
fmt.Fprintf(&b, "stack_sys: %s\n", formatMemBytes(m.StackSys))
fmt.Fprintf(&b, "mspan_sys: %s\n", formatMemBytes(m.MSpanSys))
fmt.Fprintf(&b, "mcache_sys: %s\n", formatMemBytes(m.MCacheSys))
fmt.Fprintf(&b, "gc_sys: %s\n", formatMemBytes(m.GCSys))
fmt.Fprintf(&b, "other_sys: %s\n", formatMemBytes(m.OtherSys))
fmt.Fprintf(&b, "next_gc: %s\n", formatMemBytes(m.NextGC))
fmt.Fprintf(&b, "num_gc: %d\n", m.NumGC)
fmt.Fprintf(&b, "\n")
fmt.Fprintf(&b, "goroutines: %d\n", runtime.NumGoroutine())
fmt.Fprintf(&b, "cgo_calls: %d\n", runtime.NumCgoCall())
fmt.Fprintf(&b, "gomaxprocs: %d\n", runtime.GOMAXPROCS(0))
fmt.Fprintf(&b, "\n")
writeWindowInventory(&b, app)
if err := os.WriteFile(path, []byte(b.String()), 0o644); err != nil {
return fmt.Errorf("write %s: %w", path, err)
}
return nil
}
// writeWindowInventory lists the live Wails windows. A window that exists holds
// a webview process even while hidden, so this tells apart a leaked window (the
// count grows) from windows whose content grew (the count stays put).
func writeWindowInventory(b *strings.Builder, app *application.App) {
windows := app.Window.GetAll()
fmt.Fprintf(b, "windows: %d\n", len(windows))
for _, w := range windows {
visible := "unknown"
if ww, ok := w.(*application.WebviewWindow); ok {
visible = strconv.FormatBool(ww.IsVisible())
}
fmt.Fprintf(b, " id=%-3d name=%-20q visible=%-7s minimised=%-5t focused=%t\n",
w.ID(), w.Name(), visible, w.IsMinimised(), w.IsFocused())
}
}
// writeProcTree dumps this process and its descendants with their memory
// footprint. The webview runs in child processes whose memory the Go runtime
// profiles cannot see, so this is what attributes a footprint to a component.
func writeProcTree(path string) error {
self, err := process.NewProcess(int32(os.Getpid()))
if err != nil {
return fmt.Errorf("open own process: %w", err)
}
var b strings.Builder
fmt.Fprintf(&b, "time: %s\n", time.Now().Format(time.RFC3339))
fmt.Fprintf(&b, "uptime: %s\n\n", time.Since(memProfStart).Round(time.Second))
fmt.Fprintf(&b, "%-8s %-8s %-28s %12s %12s %12s %12s\n", "PID", "PPID", "NAME", "RSS", "VMS", "PSS", "PRIV_DIRTY")
var totalRSS, totalPSS, totalPrivate uint64
walkProcTree(&b, self, 0, &totalRSS, &totalPSS, &totalPrivate)
fmt.Fprintf(&b, "\n%-8s %-8s %-28s %12s %12s %12s %12s\n", "", "", "TOTAL",
formatKB(totalRSS), "", formatKB(totalPSS), formatKB(totalPrivate))
fmt.Fprintf(&b, "\nPSS and PRIV_DIRTY come from /proc/<pid>/smaps_rollup and are Linux only.\n")
if err := os.WriteFile(path, []byte(b.String()), 0o644); err != nil {
return fmt.Errorf("write %s: %w", path, err)
}
return nil
}
// walkProcTree appends one line per process, depth-first, accumulating totals.
func walkProcTree(b *strings.Builder, p *process.Process, depth int, totalRSS, totalPSS, totalPrivate *uint64) {
name, err := p.Name()
if err != nil {
name = "unknown"
}
var rss, vms uint64
if info, err := p.MemoryInfo(); err == nil {
rss, vms = info.RSS, info.VMS
}
pss, private := smapsRollup(p.Pid)
*totalRSS += rss
*totalPSS += pss
*totalPrivate += private
ppid, err := p.Ppid()
if err != nil {
ppid = -1
}
fmt.Fprintf(b, "%-8d %-8d %-28s %12s %12s %12s %12s\n", p.Pid, ppid,
strings.Repeat(" ", depth)+name, formatKB(rss), formatKB(vms), formatKB(pss), formatKB(private))
if depth >= memProfMaxDepth {
return
}
children, err := p.Children()
if err != nil {
return
}
for _, child := range children {
walkProcTree(b, child, depth+1, totalRSS, totalPSS, totalPrivate)
}
}
// smapsRollup returns the proportional set size and private dirty bytes of pid,
// both zero on platforms without /proc.
func smapsRollup(pid int32) (uint64, uint64) {
f, err := os.Open(fmt.Sprintf("/proc/%d/smaps_rollup", pid))
if err != nil {
return 0, 0
}
defer func() {
if err := f.Close(); err != nil {
log.Debugf("close smaps_rollup for %d: %v", pid, err)
}
}()
var pss, private uint64
scanner := bufio.NewScanner(f)
for scanner.Scan() {
fields := strings.Fields(scanner.Text())
if len(fields) < 2 {
continue
}
kb, err := strconv.ParseUint(fields[1], 10, 64)
if err != nil {
continue
}
switch fields[0] {
case "Pss:":
pss = kb * 1024
case "Private_Dirty:":
private = kb * 1024
}
}
return pss, private
}
// processMemory returns the formatted resident and virtual size of this process.
func processMemory() (string, string) {
p, err := process.NewProcess(int32(os.Getpid()))
if err != nil {
unavailable := fmt.Sprintf("unavailable (%v)", err)
return unavailable, unavailable
}
info, err := p.MemoryInfo()
if err != nil {
unavailable := fmt.Sprintf("unavailable (%v)", err)
return unavailable, unavailable
}
return formatMemBytes(info.RSS), formatMemBytes(info.VMS)
}
// formatMemBytes renders a byte count as megabytes with the raw value kept.
func formatMemBytes(n uint64) string {
return fmt.Sprintf("%8.1f MB (%d bytes)", float64(n)/(1024*1024), n)
}
// formatKB renders a byte count as megabytes for the process tree columns, and
// a dash when the platform did not report the value.
func formatKB(n uint64) string {
if n == 0 {
return "-"
}
return fmt.Sprintf("%.1f MB", float64(n)/(1024*1024))
}

View File

@@ -4,17 +4,26 @@ package main
// bindTrayClick wires the tray icon's left-click handler on Linux.
//
// Both Linux click paths converge on Wails' linuxSystemTray.Activate, which
// fires the registered clickHandler:
// - Real SNI hosts (KDE Plasma, Waybar, GNOME Shell + AppIndicator) invoke
// org.kde.StatusNotifierItem.Activate over D-Bus on left-click.
// - The in-process StatusNotifierWatcher + XEmbed host used on minimal WMs
// (Fluxbox, i3, dwm, OpenBox) maps a Button1 press to that same Activate
// call itself (xembed_host_linux.go), so it routes through the same hook.
// Registering OnClick here therefore covers both paths with one handler — no
// changes to the watcher or XEmbed C code are needed. Left-click now opens the
// main window; right-click still opens the menu via Wails' default
// SecondaryActivate→OpenMenu handler (and the XEmbed GTK popup on minimal WMs).
// Expected behaviour per tray host:
//
// Host Left click Right click
// KDE Plasma, Waybar main window (Activate) menu (host-rendered)
// GNOME Shell + AppIndicator menu only menu only
// Minimal WMs via XEmbed host main window (Activate) XEmbed GTK popup
//
// OnClick fires only on org.kde.StatusNotifierItem.Activate — a real left
// click. KDE/Waybar send it over D-Bus; the in-process XEmbed host
// (xembed_host_linux.go) maps a Button1 press to the same Activate call.
//
// GNOME Shell + AppIndicator never sends Activate: it renders the dbusmenu
// on ANY click and only reports the menu opening via dbusmenu
// Event("opened"). Upstream Wails treated that event as a click, so on GNOME
// both buttons raised the main window on top of the menu, and on KDE/Waybar
// a right click raised it over the freshly opened menu. The netbirdio/wails
// fork (go.mod replace) drops that heuristic: a menu open never fires
// OnClick. On GNOME the main window is reached via the "Open NetBird" menu
// entry; left-click-opens-window is not achievable there anyway, since the
// host always opens the menu itself.
//
// We do NOT register OnDoubleClick: Wails' Linux SNI backend never fires it
// (unlike Windows). And we deliberately skip AttachWindow — it plus Wails3's

View File

@@ -27,11 +27,10 @@ const (
finalWarningCountdownSeconds = 120
)
// handleSessionExpired notifies and brings the window forward so the frontend's /login route drives renewal.
// handleSessionExpired notifies and brings the window forward so the user can reconnect.
func (t *Tray) handleSessionExpired() {
t.notify(t.loc.T("notify.sessionExpired.title"), t.loc.T("notify.sessionExpired.body"), notifyIDSessionExpired)
if t.window != nil {
t.window.SetURL("/#/login")
t.window.Show()
t.window.Focus()
}
@@ -308,11 +307,7 @@ func (t *Tray) openSessionExtendFlow() {
}
seconds := int(time.Until(deadline).Seconds())
if seconds <= 0 {
if t.window != nil {
t.window.SetURL("/#/login")
t.window.Show()
t.window.Focus()
}
t.app.Event.Emit(services.EventTriggerLogin)
return
}
if t.svc.WindowManager == nil {

View File

@@ -160,8 +160,19 @@ func TestSettingsRoundTrip(t *testing.T) {
assert.Equal(t, before.Cluster, flipped.Cluster, "cluster must be immutable across updates")
assert.Equal(t, before.Subdomain, flipped.Subdomain, "subdomain must be immutable across updates")
// A cluster different from the pinned one must be rejected; echoing the
// pinned one back is valid.
_, err = srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
Cluster: ptr("attacker.cluster.invalid"),
EnableLogCollection: before.EnableLogCollection,
EnablePromptCollection: before.EnablePromptCollection,
RedactPii: before.RedactPii,
})
requireClientError(t, err)
// Restore the original toggles.
_, err = srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
Cluster: ptr(before.Cluster),
EnableLogCollection: before.EnableLogCollection,
EnablePromptCollection: before.EnablePromptCollection,
RedactPii: before.RedactPii,

View File

@@ -0,0 +1,114 @@
//go:build e2e
package agentnetwork
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/e2e/harness"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// harnessStartFresh boots a dedicated combined server with its own fresh
// account and registers its teardown on t.
func harnessStartFresh(ctx context.Context, t *testing.T) (*harness.Combined, error) {
t.Helper()
fresh, err := harness.StartCombined(ctx)
if err != nil {
return nil, err
}
t.Cleanup(func() { _ = fresh.Terminate(context.Background()) })
if _, err := fresh.Bootstrap(ctx); err != nil {
return nil, err
}
return fresh, nil
}
// TestSettingsBootstrapViaPut covers the settings-first bootstrap path on an
// account that has never been bootstrapped: the GET reads as the defaults
// with an empty cluster/subdomain/endpoint, a PUT without a cluster has
// nothing to pin and fails, and a PUT carrying a cluster creates the row and
// pins it immutably. The shared srv cannot provide that starting state (any
// provider-creating test bootstraps it, and test order is deliberately not
// relied on), so this boots a dedicated combined server — the image is
// already built and cached by TestMain's StartCombined, so the extra cost is
// one container start.
func TestSettingsBootstrapViaPut(t *testing.T) {
ctx := context.Background()
fresh, err := harnessStartFresh(ctx, t)
require.NoError(t, err, "start dedicated combined server")
// Before agent-network bootstrap the settings read as the defaults, not
// as an error and not as a null body.
before, err := fresh.GetSettings(ctx)
require.NoError(t, err, "get settings on a fresh account must succeed")
assert.Empty(t, before.Cluster, "cluster must be empty before bootstrap")
assert.Empty(t, before.Subdomain, "subdomain must be empty before bootstrap")
assert.Empty(t, before.Endpoint, "endpoint must be empty before bootstrap, not a bare dot")
assert.True(t, before.EnableLogCollection, "defaults must show log collection on, matching bootstrap")
assert.False(t, before.EnablePromptCollection, "defaults must show prompt collection off")
// A PUT without a cluster has nothing to pin the account to.
_, err = fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
EnableLogCollection: true,
})
requireClientError(t, err)
// A PUT carrying a cluster bootstraps the account and applies the
// mutable fields from the same request. Every toggle is set away from
// its bootstrap default so each assertion can actually fail.
const cluster = "e2e.bootstrap.netbird.selfhosted"
bootstrapped, err := fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
Cluster: ptr(cluster),
EnableLogCollection: false,
EnablePromptCollection: true,
RedactPii: true,
})
require.NoError(t, err, "bootstrap settings via PUT must succeed")
assert.Equal(t, cluster, bootstrapped.Cluster, "cluster must be pinned from the request")
require.NotEmpty(t, bootstrapped.Subdomain, "subdomain must be assigned at bootstrap")
assert.Equal(t, bootstrapped.Subdomain+"."+cluster, bootstrapped.Endpoint, "endpoint must combine subdomain and cluster")
assert.False(t, bootstrapped.EnableLogCollection, "log collection from the bootstrap request must override the default")
assert.True(t, bootstrapped.EnablePromptCollection, "prompt collection from the bootstrap request must apply")
assert.True(t, bootstrapped.RedactPii, "redact toggle from the bootstrap request must apply")
// The row is persisted: an independent read agrees on every field.
after, err := fresh.GetSettings(ctx)
require.NoError(t, err, "get settings after bootstrap must succeed")
assert.Equal(t, bootstrapped.Endpoint, after.Endpoint, "bootstrap must persist across reads")
assert.Equal(t, bootstrapped.EnableLogCollection, after.EnableLogCollection, "log collection must persist")
assert.Equal(t, bootstrapped.EnablePromptCollection, after.EnablePromptCollection, "prompt collection must persist")
assert.Equal(t, bootstrapped.RedactPii, after.RedactPii, "redact toggle must persist")
// Once bootstrapped, later updates may omit the cluster entirely.
persisted, err := fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
EnableLogCollection: true,
EnablePromptCollection: false,
RedactPii: true,
})
require.NoError(t, err, "post-bootstrap update without cluster must succeed")
assert.Equal(t, cluster, persisted.Cluster, "omitted cluster must keep the pinned value")
assert.True(t, persisted.EnableLogCollection, "post-bootstrap toggle must apply")
assert.False(t, persisted.EnablePromptCollection, "post-bootstrap toggle must apply")
// The cluster is immutable: a different value is rejected rather than
// silently ignored, and the rejected update must not disturb anything.
_, err = fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
Cluster: ptr("other.cluster.invalid"),
EnableLogCollection: false,
})
requireClientError(t, err)
final, err := fresh.GetSettings(ctx)
require.NoError(t, err, "get settings after the rejected cluster change must succeed")
assert.Equal(t, persisted.Cluster, final.Cluster, "rejected update must not change the cluster")
assert.Equal(t, persisted.Endpoint, final.Endpoint, "rejected update must not change the endpoint")
assert.Equal(t, persisted.EnableLogCollection, final.EnableLogCollection, "rejected update must not apply its toggles")
assert.Equal(t, persisted.EnablePromptCollection, final.EnablePromptCollection, "rejected update must not apply its toggles")
assert.Equal(t, persisted.RedactPii, final.RedactPii, "rejected update must not apply its toggles")
}

6
go.mod
View File

@@ -81,7 +81,7 @@ require (
github.com/miekg/dns v1.1.72
github.com/mitchellh/hashstructure/v2 v2.0.2
github.com/moby/moby/api v1.54.1
github.com/netbirdio/management-integrations/integrations v0.0.0-20260416123949-2355d972be42
github.com/netbirdio/management-integrations/integrations v0.0.0-20260803100840-78e79ba20f87
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45
github.com/oapi-codegen/runtime v1.1.2
github.com/okta/okta-sdk-golang/v2 v2.18.0
@@ -114,7 +114,7 @@ require (
github.com/ti-mo/conntrack v0.5.1
github.com/ti-mo/netfilter v0.5.2
github.com/vmihailenco/msgpack/v5 v5.4.1
github.com/wailsapp/wails/v3 v3.0.0-alpha2.117
github.com/wailsapp/wails/v3 v3.0.0-beta.3
github.com/yusufpapurcu/wmi v1.2.4
github.com/zcalusic/sysinfo v1.1.3
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.67.0
@@ -339,3 +339,5 @@ replace github.com/dexidp/dex => github.com/netbirdio/dex v0.244.1-0.20260716205
replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-20260512110716-8d70ad8647c1
replace github.com/mailru/easyjson => github.com/netbirdio/easyjson v0.9.0
replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4

8
go.sum
View File

@@ -484,12 +484,14 @@ github.com/netbirdio/easyjson v0.9.0 h1:6Nw2lghSVuy8RSkAYDhDv1thBVEmfVbKZnV7T7Z6
github.com/netbirdio/easyjson v0.9.0/go.mod h1:1+xMtQp2MRNVL/V1bOzuP3aP8VNwRW55fQUto+XFtTU=
github.com/netbirdio/ice/v4 v4.0.0-20250908184934-6202be846b51 h1:Ov4qdafATOgGMB1wbSuh+0aAHcwz9hdvB6VZjh1mVMI=
github.com/netbirdio/ice/v4 v4.0.0-20250908184934-6202be846b51/go.mod h1:ZSIbPdBn5hePO8CpF1PekH2SfpTxg1PDhEwtbqZS7R8=
github.com/netbirdio/management-integrations/integrations v0.0.0-20260416123949-2355d972be42 h1:F3zS5fT9xzD1OFLfcdAE+3FfyiwjGukF1hvj0jErgs8=
github.com/netbirdio/management-integrations/integrations v0.0.0-20260416123949-2355d972be42/go.mod h1:n47r67ZSPgwSmT/Z1o48JjZQW9YJ6m/6Bd/uAXkL3Pg=
github.com/netbirdio/management-integrations/integrations v0.0.0-20260803100840-78e79ba20f87 h1:iJeUvSMC0BTpkw7u4JyWcY4/3dl7fEL9DR/TpKf2+1w=
github.com/netbirdio/management-integrations/integrations v0.0.0-20260803100840-78e79ba20f87/go.mod h1:pmsCPx1S0nuZRxCextGpc9AV4hLgGSuTsc4NMuwGeCo=
github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502 h1:3tHlFmhTdX9axERMVN63dqyFqnvuD+EMJHzM7mNGON8=
github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502/go.mod h1:CIMRFEJVL+0DS1a3Nx06NaMn4Dz63Ng6O7dl0qH0zVM=
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45 h1:ujgviVYmx243Ksy7NdSwrdGPSRNE3pb8kEDSpH0QuAQ=
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45/go.mod h1:5/sjFmLb8O96B5737VCqhHyGRzNFIaN/Bu7ZodXc3qQ=
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4 h1:UKztc3QjWvzU5DZk+uYaOWN0x62NSe/pkxuPvzqZIy4=
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260803205919-ad21e92381f4/go.mod h1:BzATbK71VFikMMMCo434wAi0QcaI03P+xeaWgDvQvjw=
github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a h1:3CWK+yTvRKOcC0Q8VCTGy4l60TEb27CQVS7LkMxwjmw=
github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
github.com/nxadm/tail v1.4.4/go.mod h1:kenIhsEOeOJmVchQTgglprH7qJGnHDVpk1VPCcaMI8A=
@@ -660,8 +662,6 @@ github.com/vmihailenco/msgpack/v5 v5.4.1 h1:cQriyiUvjTwOHg8QZaPihLWeRAAVoCpE00IU
github.com/vmihailenco/msgpack/v5 v5.4.1/go.mod h1:GaZTsDaehaPpQVyxrf5mtQlH+pc21PIudVV/E3rRQok=
github.com/vmihailenco/tagparser/v2 v2.0.0 h1:y09buUbR+b5aycVFQs/g70pqKVZNBmxwAhO7/IwNM9g=
github.com/vmihailenco/tagparser/v2 v2.0.0/go.mod h1:Wri+At7QHww0WTrCBeu4J6bNtoV6mEfg5OIWRZA9qds=
github.com/wailsapp/wails/v3 v3.0.0-alpha2.117 h1:udyjqPG3AIgkod5QDR/WblCkpV8R86BFPSrsWxSyt5Y=
github.com/wailsapp/wails/v3 v3.0.0-alpha2.117/go.mod h1:74WH2FScMsgucZvHHvv7eOefDXCm/CjuIxqhhZgPhKg=
github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU=
github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA=
github.com/x448/float16 v0.8.4 h1:qLwI1I70+NjRFUR3zs1JPUCgaCXSh3SW62uAKT1mSBM=

View File

@@ -0,0 +1,57 @@
package networkmap_pgsql
import (
"context"
"testing"
"time"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetAccountSettings(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into accounts (id, settings_peer_login_expiration_enabled, settings_peer_login_expiration, settings_peer_inactivity_expiration_enabled,
settings_peer_inactivity_expiration, settings_dns_domain, settings_ipv6_enabled_groups, settings_routing_peer_dns_resolution_enabled,
settings_lazy_connection_enabled, settings_auto_update_version, settings_auto_update_always, settings_metrics_push_enabled)
values('account-3',null,null,null,null,null,null,null,null,null,null,null)`)
accountSettings, err := networkmap_pgsql.GetAccountSettingsViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Equal(t, accountSettings, nmdata.AccountSettingsInfo{
PeerLoginExpirationEnabled: true,
PeerLoginExpiration: 86400000000000 * time.Nanosecond,
PeerInactivityExpirationEnabled: false,
PeerInactivityExpiration: 86400000000000 * time.Nanosecond,
DNSDomain: "",
IPv6EnabledGroups: []string{"group-one-resource-id"},
RoutingPeerDNSResolutionEnabled: false,
LazyConnectionEnabled: false,
AutoUpdateVersion: "disabled",
AutoUpdateAlways: false,
MetricsPushEnabled: false,
})
accountSettings, err = networkmap_pgsql.GetAccountSettingsViaPgxConnection(ctx, conn(t, ctx), "account-2")
assert.NoError(t, err)
assert.Equal(t, accountSettings, nmdata.AccountSettingsInfo{
PeerLoginExpirationEnabled: true,
PeerLoginExpiration: 86400000000000 * time.Nanosecond,
PeerInactivityExpirationEnabled: false,
PeerInactivityExpiration: 86400000000000 * time.Nanosecond,
DNSDomain: "",
IPv6EnabledGroups: []string{"group-two-resources-id"},
RoutingPeerDNSResolutionEnabled: false,
LazyConnectionEnabled: false,
AutoUpdateVersion: "disabled",
AutoUpdateAlways: false,
MetricsPushEnabled: false,
})
accountSettings, err = networkmap_pgsql.GetAccountSettingsViaPgxConnection(ctx, conn(t, ctx), "account-3")
assert.NoError(t, err)
assert.Equal(t, accountSettings, nmdata.AccountSettingsInfo{})
}

View File

@@ -0,0 +1,52 @@
insert into accounts (id, network_identifier, network_net, network_net_v6, network_dns, network_serial,dns_settings_disabled_management_groups,
settings_peer_login_expiration_enabled, settings_peer_login_expiration, settings_peer_inactivity_expiration_enabled,
settings_peer_inactivity_expiration, settings_dns_domain, settings_ipv6_enabled_groups, settings_routing_peer_dns_resolution_enabled,
settings_lazy_connection_enabled, settings_auto_update_version, settings_auto_update_always, settings_metrics_push_enabled)
VALUES('account-1','network-1','{"IP":"100.103.0.0","Mask":"//8AAA=="}','{"IP":"fdde:e995:fd38:a465::","Mask":"//////////8AAAAAAAAAAA=="}','',1,'["disabled-group-1","disabled-group-2"]',
true, 86400000000000, false,
86400000000000, null, '["group-one-resource-id"]', false,
false, 'disabled', false, false);
insert into accounts (id, network_identifier, network_net, network_net_v6, network_dns, network_serial,dns_settings_disabled_management_groups,
settings_peer_login_expiration_enabled, settings_peer_login_expiration, settings_peer_inactivity_expiration_enabled,
settings_peer_inactivity_expiration, settings_dns_domain, settings_ipv6_enabled_groups, settings_routing_peer_dns_resolution_enabled,
settings_lazy_connection_enabled, settings_auto_update_version, settings_auto_update_always, settings_metrics_push_enabled)
VALUES('account-2','network-2','{"IP":"110.0.0.0","Mask":"//8AAA=="}','{"IP":"fddf:e995:fd38:a465::","Mask":"//////////8AAAAAAAAAAA=="}','',2,null,
true, 86400000000000, false,
86400000000000, null, '["group-two-resources-id"]', false,
false, 'disabled', false, false);
insert into groups (id, account_id, name, resources, public_id) VALUES('group-one-resource-id','account-1','group-1-name', '[{"ID":"host-id-1","Type":"host"}]','group-one-resource-id-public');
insert into groups (id, account_id, name, resources, public_id) VALUES('group-two-resources-id','account-1','group-2-name', '[{"ID":"subnet-id-1","Type":"subnet"}, {"ID":"host-id-2","Type":"host"}]','group-two-resources-id-public');
insert into groups (id, account_id, name, resources, public_id) VALUES('group-no-resources-id','account-1','group-3-name', null,'group-no-resources-id-public');
insert into group_peers (account_id, peer_id, group_id) VALUES('account-1','peer-id-1','group-one-resource-id');
insert into group_peers (account_id, peer_id, group_id) VALUES('account-1','peer-id-2','group-two-resources-id');
insert into group_peers (account_id, peer_id, group_id) VALUES('account-1','peer-id-3','group-two-resources-id');
insert into peers (id, account_id, "key", ssh_key, dns_label, extra_dns_labels, user_id, ssh_enabled, login_expiration_enabled, last_login, ip, ipv6,
peer_status_requires_approval, peer_status_connected, proxy_meta_embedded, proxy_meta_cluster,
meta_wt_version, meta_go_os, meta_os_version, meta_kernel_version, meta_network_addresses, meta_files,
meta_capabilities, meta_flags, meta_sync_message_version,
location_country_code, location_city_name, location_connection_ip)
values('peer-id-1','account-1','key-1','ssh-key-1','peer-1','["extra-peer-1"]','user-id-1',true,true,'2026-08-06 13:25:59.12999+00','"10.10.10.1"','"fdf4:ba80:6aa5:89f1:44d7:8701:8699:4940"',
false,true,true,'cluster-1.netbird.services',
'0.76.0','linux','26.4.1','6.8.0-134-generic','[{"NetIP":"fe80::8b4c:973f:a76b:3771/64","Mac":"00:15:5d:24:0c:ac"},{"NetIP":"192.168.16.1/20","Mac":"00:15:5d:24:0c:ac"}]','[{"Path":"/usr/bin/netbird","Exist":false,"ProcessIsRunning":false}]',
'[1,2]','{"RosenpassEnabled":false,"RosenpassPermissive":false,"ServerSSHAllowed":true,"DisableClientRoutes":false,"DisableServerRoutes":false,"DisableDNS":false,"DisableFirewall":false,"BlockLANAccess":false,"BlockInbound":false,"DisableIPv6":false,"LazyConnectionEnabled":false}',1,
'DE','Berlin','"46.201.148.187"');
insert into peers (id,account_id,"key", ssh_key, dns_label, extra_dns_labels, user_id, ssh_enabled, login_expiration_enabled, last_login, ip, ipv6,
peer_status_requires_approval, peer_status_connected, proxy_meta_embedded, proxy_meta_cluster,
meta_wt_version, meta_go_os, meta_os_version, meta_kernel_version, meta_network_addresses, meta_files,
meta_capabilities, meta_flags, meta_sync_message_version,
location_country_code, location_city_name, location_connection_ip)
values('peer-id-2','account-1','key-2','ssh-key-2','peer-2','["extra-peer-2"]','user-id-2',true,true,'2026-08-06 14:25:59.12999+00','"10.10.100.1"','"fdf5:ba80:6aa5:89f1:44d7:8701:8699:4940"',
false,true,true,'cluster-2.netbird.services',
'0.76.1','linux','26.4.2','6.8.0-135-generic','[{"NetIP":"fe81::8b4c:973f:a76b:3771/64","Mac":"00:15:5d:24:0c:ad"},{"NetIP":"192.168.17.1/20","Mac":"00:15:5d:24:0c:ad"}]','[{"Path":"/usr/bin/netbird","Exist":false,"ProcessIsRunning":false}]',
'[1,2]','{"RosenpassEnabled":false,"RosenpassPermissive":false,"ServerSSHAllowed":true,"DisableClientRoutes":false,"DisableServerRoutes":false,"DisableDNS":false,"DisableFirewall":false,"BlockLANAccess":false,"BlockInbound":false,"DisableIPv6":false,"LazyConnectionEnabled":false}',0,
'DE','Berlin','"46.201.149.187"');
insert into peers (id,account_id,"key", ssh_key, dns_label, extra_dns_labels, user_id, ssh_enabled, login_expiration_enabled, last_login, ip, ipv6,
peer_status_requires_approval, peer_status_connected, proxy_meta_embedded, proxy_meta_cluster,
meta_wt_version, meta_go_os, meta_os_version, meta_kernel_version, meta_network_addresses, meta_files,
meta_capabilities, meta_flags, meta_sync_message_version,
location_country_code, location_city_name, location_connection_ip)
values('peer-id-3','account-1','key-3','ssh-key-3','peer-3','["extra-peer-3"]','user-id-3',true,true,'2026-08-06 12:25:59.12999+00','"10.10.200.1"','"fdf6:ba80:6aa5:89f1:44d7:8701:8699:4940"',
false,true,true,'cluster-3.netbird.services',
'0.76.2','linux','26.4.3','6.8.0-136-generic','[{"NetIP":"fe82::8b4c:973f:a76b:3771/64","Mac":"00:15:5d:24:0c:ae"},{"NetIP":"192.168.18.1/20","Mac":"00:15:5d:24:0c:ae"}]','[{"Path":"/usr/bin/netbird","Exist":false,"ProcessIsRunning":false}]',
'[1,2]','{"RosenpassEnabled":false,"RosenpassPermissive":false,"ServerSSHAllowed":true,"DisableClientRoutes":false,"DisableServerRoutes":false,"DisableDNS":false,"DisableFirewall":false,"BlockLANAccess":false,"BlockInbound":false,"DisableIPv6":false,"LazyConnectionEnabled":false}',1,
'DE','Berlin','"46.201.150.187"');

View File

@@ -0,0 +1,23 @@
package networkmap_pgsql
import (
"context"
"testing"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetDnsSettings(t *testing.T) {
ctx := context.TODO()
settings, err := pgstore.GetDnsSettings(ctx, "account-1")
assert.NoError(t, err)
assert.Equal(t, settings, nmdata.DNSSettings{
DisabledManagementGroups: []string{"disabled-group-1", "disabled-group-2"},
})
settings, err = pgstore.GetDnsSettings(ctx, "account-2")
assert.NoError(t, err)
assert.Equal(t, settings, nmdata.DNSSettings{})
}

View File

@@ -0,0 +1,61 @@
package networkmap_pgsql
import (
"context"
"testing"
"github.com/miekg/dns"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetAppliedZoneCandidatesViaPgxConnection(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into zones (id, account_id, domain, enable_search_domain, distribution_groups)
VALUES('zone-1','account-1','test-1.com',true,'["group-one-resource-id"]')`)
execQuery(t, ctx,
`insert into zones (id, account_id, domain, enable_search_domain, distribution_groups)
VALUES('zone-2','account-1','test-2.com',false,'["group-two-resources-id"]')`)
execQuery(t, ctx,
`insert into records (id, account_id, zone_id, name, type, ttl, content)
VALUES('record-1','account-1','zone-1','test.test-1.com','A',1800,'1.1.1.1')`)
execQuery(t, ctx,
`insert into records (id, account_id, zone_id, name, type, ttl, content)
VALUES('record-2','account-1','zone-1','test2.test-1.com','A',1800,'1.1.1.2')`)
execQuery(t, ctx,
`insert into records (id, account_id, zone_id, name, type, ttl, content)
VALUES('record-3','account-1','zone-1','test3.test-1.com','CNAME',1800,'test4.test-1.com')`)
execQuery(t, ctx,
`insert into records (id, account_id, zone_id, name, type, ttl, content)
VALUES('record-4','account-1','zone-2','test2.test-2.com','CNAME',1800,'test3.test-2.com')`)
zoneCandidates, err := networkmap_pgsql.GetAppliedZoneCandidatesViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Contains(t, zoneCandidates, networkmap.AppliedZoneCandidate{
DistributionGroups: []string{"group-one-resource-id"},
Zone: nmdata.CustomZone{
Domain: "test-1.com",
SearchDomainDisabled: false,
Records: []nmdata.SimpleRecord{
{Name: "test.test-1.com", Type: int(dns.TypeA), Class: "IN", TTL: 1800, RData: "1.1.1.1"},
{Name: "test2.test-1.com", Type: int(dns.TypeA), Class: "IN", TTL: 1800, RData: "1.1.1.2"},
{Name: "test3.test-1.com", Type: int(dns.TypeCNAME), Class: "IN", TTL: 1800, RData: "test4.test-1.com."},
},
},
})
assert.Contains(t, zoneCandidates, networkmap.AppliedZoneCandidate{
DistributionGroups: []string{"group-two-resources-id"},
Zone: nmdata.CustomZone{
Domain: "test-2.com",
SearchDomainDisabled: true,
Records: []nmdata.SimpleRecord{
{Name: "test2.test-2.com", Type: int(dns.TypeCNAME), Class: "IN", TTL: 1800, RData: "test3.test-2.com."},
},
},
})
}

View File

@@ -0,0 +1,37 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/stretchr/testify/assert"
)
func TestGetDomains(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into domains (id, account_id, domain, target_cluster)
VALUES('domain-1','account-1','test-1.com','target-1.cluster.local')`)
execQuery(t, ctx,
`insert into domains (id, account_id, domain, target_cluster)
VALUES('domain-2','account-1','test-2.com','target-2.cluster.local')`)
execQuery(t, ctx,
`insert into domains (id, account_id, domain, target_cluster)
VALUES('domain-3','account-1',null,null)`)
domains, err := pgstore.GetDomains(ctx, "account-1")
assert.NoError(t, err)
assert.Len(t, domains, 2)
assert.Contains(t, domains, networkmap_pgsql.Domain{
Domain: sql.NullString{String: "test-1.com", Valid: true},
TargetCluster: sql.NullString{String: "target-1.cluster.local", Valid: true},
})
assert.Contains(t, domains, networkmap_pgsql.Domain{
Domain: sql.NullString{String: "test-2.com", Valid: true},
TargetCluster: sql.NullString{String: "target-2.cluster.local", Valid: true},
})
}

View File

@@ -0,0 +1,59 @@
package networkmap_pgsql
import (
"context"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetGroups(t *testing.T) {
ctx := context.TODO()
s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn)
assert.NoError(t, err)
groups, resourceToGroupIdx, err := s.GetGroups(ctx, "account-1")
assert.NoError(t, err)
assert.Contains(t,
groups,
nmdata.Group{ID: "group-one-resource-id", Name: "group-1-name", PublicID: "group-one-resource-id-public", Resources: []nmdata.Resource{{ID: "host-id-1", Type: "host"}}, Peers: []string{"peer-id-1"}},
)
assert.NotNil(t, resourceToGroupIdx["host-id-1"]["group-one-resource-id"])
assert.Contains(t,
groups,
nmdata.Group{ID: "group-two-resources-id", Name: "group-2-name", PublicID: "group-two-resources-id-public",
Resources: []nmdata.Resource{{ID: "subnet-id-1", Type: "subnet"}, {ID: "host-id-2", Type: "host"}},
Peers: []string{"peer-id-2", "peer-id-3"}},
)
assert.NotNil(t, resourceToGroupIdx["host-id-2"]["group-two-resources-id"])
assert.NotNil(t, resourceToGroupIdx["subnet-id-1"]["group-two-resources-id"])
assert.Contains(t,
groups,
nmdata.Group{ID: "group-no-resources-id", Name: "group-3-name", PublicID: "group-no-resources-id-public"})
}
// Verify handling of empty fields in groups table
// Verify that group's PublicID gets populated on retrieval
// TODO (dmitri) PublicID should not be populated with delta updates,
// which require stable PublicIDs
func TestGetGroupsWithoutExpectedFields(t *testing.T) {
ctx := context.TODO()
s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn)
assert.NoError(t, err)
execQuery(t, ctx,
"insert into accounts (id) VALUES('random-id')")
execQuery(t, ctx,
"insert into groups (id, account_id) VALUES('g2-test-group-id-1','random-id')")
assert.NoError(t, err)
groups, _, err := s.GetGroups(ctx, "random-id")
assert.NoError(t, err)
assert.Len(t, groups, 1)
assert.NotEmpty(t, groups[0].PublicID)
}

View File

@@ -0,0 +1,166 @@
package networkmap_pgsql
import (
"context"
_ "embed"
"fmt"
"os"
"regexp"
"strings"
"testing"
"time"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
"github.com/stretchr/testify/assert"
log "github.com/sirupsen/logrus"
"gorm.io/driver/postgres"
"gorm.io/gorm"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
gormstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/testutil"
)
//go:embed base_data.sql
var baseData string
var (
dsn string
pgstore *networkmap_pgsql.PgStore
)
func TestMain(m *testing.M) {
_, tmpdsn, err := testutil.CreatePostgresTestContainer()
if err != nil {
log.Fatalf("error starting postres container %v", err)
}
var db *gorm.DB
for i := range 5 {
db, err = gorm.Open(postgres.Open(tmpdsn), &gorm.Config{})
if err == nil {
break
}
if i < 5 {
waitTime := time.Duration(100*(i+1)) * time.Millisecond
time.Sleep(waitTime)
continue
}
log.Fatalf("error connecting to postres db %v", err)
}
var cleanup func()
dsn, cleanup, err = createRandomDB(tmpdsn, db)
sqlDB, _ := db.DB()
if sqlDB != nil {
sqlDB.Close()
}
if err != nil {
log.Fatalf("error creating postres db %v", err)
}
_, err = gormstore.NewPostgresqlStoreForTests(context.TODO(), dsn, nil, false)
if err != nil {
log.Fatalf("error running migrations %v", err)
}
ctx := context.TODO()
s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn)
if err != nil {
log.Fatal("error creating postgres store %w", err)
}
for _, query := range strings.Split(baseData, ";") {
if _, err := s.Pool.Exec(ctx, query); err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
}
pgstore, err = networkmap_pgsql.NewPostgresqlStore(ctx, dsn)
if err != nil {
log.Fatalf("error creating pg store %v", err.Error())
}
code := m.Run()
cleanup()
os.Exit(code)
}
func createRandomDB(dsn string, db *gorm.DB) (string, func(), error) {
dbName := fmt.Sprintf("test_db_%s", strings.ReplaceAll(uuid.New().String(), "-", "_"))
if err := db.Exec(fmt.Sprintf("CREATE DATABASE %s", dbName)).Error; err != nil {
return "", nil, fmt.Errorf("failed to create database: %v", err)
}
originalDSN := dsn
cleanup := func() {
var dropDB *gorm.DB
var err error
dropDB, err = gorm.Open(postgres.Open(originalDSN), &gorm.Config{
SkipDefaultTransaction: true,
PrepareStmt: false,
})
if err != nil {
log.Errorf("failed to connect for dropping database %s: %v", dbName, err)
return
}
defer func() {
if sqlDB, _ := dropDB.DB(); sqlDB != nil {
sqlDB.Close()
}
}()
if sqlDB, _ := dropDB.DB(); sqlDB != nil {
sqlDB.SetMaxOpenConns(1)
sqlDB.SetMaxIdleConns(0)
sqlDB.SetConnMaxLifetime(time.Second)
}
err = dropDB.Exec(fmt.Sprintf("DROP DATABASE IF EXISTS %s WITH (FORCE)", dbName)).Error
if err != nil {
log.Errorf("failed to drop database %s: %v", dbName, err)
}
}
return replaceDBName(dsn, dbName), cleanup, nil
}
func replaceDBName(dsn, newDBName string) string {
re := regexp.MustCompile(`(?P<pre>[:/@])(?P<dbname>[^/?]+)(?P<post>\?|$)`)
return re.ReplaceAllString(dsn, `${pre}`+newDBName+`${post}`)
}
func conn(t *testing.T, ctx context.Context) *pgx.Conn {
t.Helper()
c, err := pgstore.Pool.Acquire(ctx)
assert.NoError(t, err)
return c.Conn()
}
func execQuery(t *testing.T, ctx context.Context, q string) {
t.Helper()
_, err := pgstore.Pool.Exec(ctx, q)
assert.NoError(t, err)
}
// use to parse time in time.RFC3339Nano format
// returns the time in the local time zone, as that's what being returned from sql queries
func mustParseTime(t string) *time.Time {
tt, err := time.Parse(time.RFC3339Nano, t)
if err != nil {
panic(err)
}
local := tt.Local()
return &local
}

View File

@@ -0,0 +1,60 @@
package networkmap_pgsql
import (
"context"
"net/netip"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetNameServerGroups(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into name_server_groups (id, public_id, name, description, name_servers, groups, domains, enabled, search_domains_enabled, "primary", account_id)
VALUES('nsgroup-1','nsgroup-1-public','nsgroup-1','nsgroup-1','[{"IP":"192.168.31.2","NSType":1,"Port":53}]','["group-one-resource-id"]','["test-1.com"]',TRUE,FALSE,TRUE,'account-1')`)
execQuery(t, ctx,
`insert into name_server_groups (id, public_id, name, description, name_servers, groups, domains, enabled, search_domains_enabled,"primary",account_id)
VALUES('nsgroup-2','nsgroup-2-public','nsgroup-2','nsgroup-2','[{"IP":"192.168.32.3","NSType":1,"Port":53}]','["group-one-resource-id","group-no-resources-id"]','["test-1.com","test-2.com"]',TRUE,FALSE,TRUE,'account-1')`)
execQuery(t, ctx,
`insert into name_server_groups (id, public_id, name, description, name_servers, groups, domains, enabled, search_domains_enabled,"primary",account_id)
VALUES('nsgroup-3','nsgroup-3-public',null,null,null,null,null,TRUE,FALSE,FALSE,'account-1')`)
nsgroups, err := networkmap_pgsql.GetNameServerGroupsViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Contains(t, nsgroups, nmdata.NameServerGroup{
ID: "nsgroup-1",
PublicID: "nsgroup-1-public",
Name: "nsgroup-1",
Description: "nsgroup-1",
NameServers: []nmdata.NameServer{{IP: netip.MustParseAddr("192.168.31.2"), NSType: 1, Port: 53}},
Groups: []string{"group-one-resource-id"},
Domains: []string{"test-1.com"},
Primary: true,
SearchDomainsEnabled: false,
Enabled: true,
})
assert.Contains(t, nsgroups, nmdata.NameServerGroup{
ID: "nsgroup-2",
PublicID: "nsgroup-2-public",
Name: "nsgroup-2",
Description: "nsgroup-2",
NameServers: []nmdata.NameServer{{IP: netip.MustParseAddr("192.168.32.3"), NSType: 1, Port: 53}},
Groups: []string{"group-one-resource-id", "group-no-resources-id"},
Domains: []string{"test-1.com", "test-2.com"},
Primary: true,
SearchDomainsEnabled: false,
Enabled: true,
})
assert.Contains(t, nsgroups, nmdata.NameServerGroup{
ID: "nsgroup-3",
PublicID: "nsgroup-3-public",
Primary: false,
SearchDomainsEnabled: false,
Enabled: true,
})
}

View File

@@ -0,0 +1,67 @@
package networkmap_pgsql
import (
"context"
"net/netip"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetNetworkResources(t *testing.T) {
ctx := context.TODO()
s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn)
assert.NoError(t, err)
execQuery(t, ctx,
`insert into network_resources (id, account_id, network_id, public_id, name, description, type, domain, prefix, enabled)
VALUES('net-resource-1','account-1','network-1','net-resource-public-1','network-resource-1','network-resource-1','subnet','','"10.0.0.0/16"',TRUE)`)
execQuery(t, ctx,
`insert into network_resources (id, account_id, network_id, public_id, name, description, type, domain, prefix, enabled)
VALUES('net-resource-2','account-1','network-2','net-resource-public-2','network-resource-2','network-resource-2','domain','test.com','',TRUE)`)
execQuery(t, ctx,
`insert into network_resources (id, account_id, network_id, public_id, name, description, type, domain, prefix, enabled)
VALUES('net-resource-3','account-1','network-3','net-resource-public-3','network-resource-3','network-resource-3','host','','"10.0.0.1/32"',TRUE)`)
resources, err := s.GetNetworkResources(ctx, "account-1")
assert.NoError(t, err)
assert.Contains(t, resources, nmdata.NetworkResource{
ID: "net-resource-1",
AccountID: "account-1",
NetworkID: "network-1",
PublicID: "net-resource-public-1",
Name: "network-resource-1",
Description: "network-resource-1",
Type: "subnet",
Domain: "",
Prefix: netip.MustParsePrefix("10.0.0.0/16"),
Enabled: true,
})
assert.Contains(t, resources, nmdata.NetworkResource{
ID: "net-resource-2",
AccountID: "account-1",
NetworkID: "network-2",
PublicID: "net-resource-public-2",
Name: "network-resource-2",
Description: "network-resource-2",
Type: "domain",
Domain: "test.com",
Enabled: true,
})
assert.Contains(t, resources, nmdata.NetworkResource{
ID: "net-resource-3",
AccountID: "account-1",
NetworkID: "network-3",
PublicID: "net-resource-public-3",
Name: "network-resource-3",
Description: "network-resource-3",
Type: "host",
Domain: "",
Prefix: netip.MustParsePrefix("10.0.0.1/32"),
Enabled: true,
})
}

View File

@@ -0,0 +1,35 @@
package networkmap_pgsql
import (
"context"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetNetworkRouters(t *testing.T) {
ctx := context.TODO()
s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn)
assert.NoError(t, err)
execQuery(t, ctx,
`insert into network_routers (id, account_id, public_id, peer, network_id, masquerade, metric, enabled, peer_groups)
VALUES('test-nr-id-1','account-1','public-id-1','peer-id-1','network-id-1',TRUE,999,TRUE,'["group-one-resource-id"]')`)
execQuery(t, ctx,
`insert into network_routers (id, account_id, public_id, peer, network_id, masquerade, metric, enabled, peer_groups)
VALUES('test-nr-id-2','account-1','public-id-2','','network-id-2',TRUE,333,TRUE,'["group-two-resources-id","group-no-resources-id"]')`)
routers, err := s.GetNetworkRouters(ctx, "account-1")
assert.NoError(t, err)
assert.NotEmpty(t, routers)
assert.Equal(t, routers["network-id-1"],
map[string]*nmdata.NetworkRouter{"peer-id-1": {PublicID: "public-id-1", Masquerade: true, Metric: 999, Enabled: true, PeerGroups: []string{"group-one-resource-id"}}})
assert.Equal(t, routers["network-id-2"],
map[string]*nmdata.NetworkRouter{
"peer-id-2": {PublicID: "public-id-2", Masquerade: true, Metric: 333, Enabled: true, PeerGroups: []string{"group-two-resources-id", "group-no-resources-id"}},
"peer-id-3": {PublicID: "public-id-2", Masquerade: true, Metric: 333, Enabled: true, PeerGroups: []string{"group-two-resources-id", "group-no-resources-id"}}})
}

View File

@@ -0,0 +1,58 @@
package networkmap_pgsql
import (
"context"
"encoding/json"
"net"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetNetwork(t *testing.T) {
ctx := context.TODO()
s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn)
assert.NoError(t, err)
network, err := s.GetNetwork(ctx, "account-1")
assert.NoError(t, err)
assert.Equal(t, network, nmdata.Network{
Identifier: "network-1",
Net: mustParseCIDR("100.103.0.0/16"),
NetV6: mustParseCIDR("fdde:e995:fd38:a465::/64"),
Serial: 1,
})
network, err = s.GetNetwork(ctx, "account-2")
assert.NoError(t, err)
assert.Equal(t, network, nmdata.Network{
Identifier: "network-2",
Net: mustParseCIDR("110.0.0.0/16"),
NetV6: mustParseCIDR("fddf:e995:fd38:a465::/64"),
Serial: 2,
})
}
func mustParseCIDR(s string) net.IPNet {
var toret net.IPNet
_, net, err := net.ParseCIDR(s)
if err != nil {
panic(err)
}
jn, err := json.Marshal(net)
if err != nil {
panic(err)
}
err = json.Unmarshal(jn, &toret)
if err != nil {
panic(err)
}
return toret
}

View File

@@ -0,0 +1,25 @@
package networkmap_pgsql
import (
"context"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/stretchr/testify/assert"
)
func TestGetNetworks(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into networks (id, account_id, public_id) VALUES('network-1','account-1','network-1-public')`)
execQuery(t, ctx,
`insert into networks (id, account_id, public_id) VALUES('network-2','account-1','network-2-public')`)
networksIdx, err := networkmap_pgsql.GetNetworkXIDToPublicIdMapViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Equal(t, networksIdx, map[string]string{
"network-1": "network-1-public",
"network-2": "network-2-public",
})
}

View File

@@ -0,0 +1,162 @@
package networkmap_pgsql
import (
"context"
"net"
"net/netip"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetPeers(t *testing.T) {
ctx := context.TODO()
peers, clusterToPeersIdx, err := networkmap_pgsql.GetPeersViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
// shouldn't be returned in the index, as it's not connected
execQuery(t, ctx,
`insert into peers (id,account_id,"key",ssh_key,proxy_meta_embedded,peer_status_connected)
values('peer-4','account-1','key-4','ssh-key-4',true,false)`)
// shouldn't be returned in the index as it doesn't have cluster set
execQuery(t, ctx,
`insert into peers (id,account_id,"key",ssh_key,proxy_meta_embedded,peer_status_connected)
values('peer-5','account-1','key-5','ssh-key-5',false,true)`)
peer1 := nmdata.Peer{
ID: "peer-id-1",
Key: "key-1",
SSHKey: "ssh-key-1",
DNSLabel: "peer-1",
ExtraDNSLabels: []string{"extra-peer-1"},
UserID: "user-id-1",
SSHEnabled: true,
LoginExpirationEnabled: true,
LastLogin: mustParseTime("2026-08-06T13:25:59.12999+00:00"),
IP: netip.MustParseAddr("10.10.10.1"),
IPv6: netip.MustParseAddr("fdf4:ba80:6aa5:89f1:44d7:8701:8699:4940"),
RequiresApproval: false,
Meta: nmdata.PeerSystemMeta{
WtVersion: "0.76.0",
GoOS: "linux",
OSVersion: "26.4.1",
KernelVersion: "6.8.0-134-generic",
NetworkAddresses: []nmdata.NetworkAddress{
{NetIP: netip.MustParsePrefix("fe80::8b4c:973f:a76b:3771/64")},
{NetIP: netip.MustParsePrefix("192.168.16.1/20")},
},
Files: []nmdata.File{
{Path: "/usr/bin/netbird", ProcessIsRunning: false},
},
Capabilities: []int32{1, 2},
Flags: nmdata.Flags{
ServerSSHAllowed: true,
DisableIPv6: false,
},
SyncMessageVersion: 1,
},
ProxyMeta: nmdata.ProxyMeta{
Embedded: true,
},
Location: nmdata.PeerLocation{
CountryCode: "DE",
CityName: "Berlin",
ConnectionIP: net.ParseIP("46.201.148.187"),
},
}
peer2 := nmdata.Peer{
ID: "peer-id-2",
Key: "key-2",
SSHKey: "ssh-key-2",
DNSLabel: "peer-2",
ExtraDNSLabels: []string{"extra-peer-2"},
UserID: "user-id-2",
SSHEnabled: true,
LoginExpirationEnabled: true,
LastLogin: mustParseTime("2026-08-06T14:25:59.12999+00:00"),
IP: netip.MustParseAddr("10.10.100.1"),
IPv6: netip.MustParseAddr("fdf5:ba80:6aa5:89f1:44d7:8701:8699:4940"),
RequiresApproval: false,
Meta: nmdata.PeerSystemMeta{
WtVersion: "0.76.1",
GoOS: "linux",
OSVersion: "26.4.2",
KernelVersion: "6.8.0-135-generic",
NetworkAddresses: []nmdata.NetworkAddress{
{NetIP: netip.MustParsePrefix("fe81::8b4c:973f:a76b:3771/64")},
{NetIP: netip.MustParsePrefix("192.168.17.1/20")},
},
Files: []nmdata.File{
{Path: "/usr/bin/netbird", ProcessIsRunning: false},
},
Capabilities: []int32{1, 2},
Flags: nmdata.Flags{
ServerSSHAllowed: true,
DisableIPv6: false,
},
SyncMessageVersion: 0,
},
ProxyMeta: nmdata.ProxyMeta{
Embedded: true,
},
Location: nmdata.PeerLocation{
CountryCode: "DE",
CityName: "Berlin",
ConnectionIP: net.ParseIP("46.201.149.187"),
},
}
peer3 := nmdata.Peer{
ID: "peer-id-3",
Key: "key-3",
SSHKey: "ssh-key-3",
DNSLabel: "peer-3",
ExtraDNSLabels: []string{"extra-peer-3"},
UserID: "user-id-3",
SSHEnabled: true,
LoginExpirationEnabled: true,
LastLogin: mustParseTime("2026-08-06T12:25:59.12999+00:00"),
IP: netip.MustParseAddr("10.10.200.1"),
IPv6: netip.MustParseAddr("fdf6:ba80:6aa5:89f1:44d7:8701:8699:4940"),
RequiresApproval: false,
Meta: nmdata.PeerSystemMeta{
WtVersion: "0.76.2",
GoOS: "linux",
OSVersion: "26.4.3",
KernelVersion: "6.8.0-136-generic",
NetworkAddresses: []nmdata.NetworkAddress{
{NetIP: netip.MustParsePrefix("fe82::8b4c:973f:a76b:3771/64")},
{NetIP: netip.MustParsePrefix("192.168.18.1/20")},
},
Files: []nmdata.File{
{Path: "/usr/bin/netbird", ProcessIsRunning: false},
},
Capabilities: []int32{1, 2},
Flags: nmdata.Flags{
ServerSSHAllowed: true,
DisableIPv6: false,
},
SyncMessageVersion: 1,
},
ProxyMeta: nmdata.ProxyMeta{
Embedded: true,
},
Location: nmdata.PeerLocation{
CountryCode: "DE",
CityName: "Berlin",
ConnectionIP: net.ParseIP("46.201.150.187"),
},
}
assert.Contains(t, peers, peer1)
assert.Contains(t, peers, peer2)
assert.Contains(t, peers, peer3)
assert.Equal(t, clusterToPeersIdx, map[string][]*nmdata.Peer{
"cluster-1.netbird.services": {&peer1},
"cluster-2.netbird.services": {&peer2},
"cluster-3.netbird.services": {&peer3},
})
}

View File

@@ -0,0 +1,145 @@
package networkmap_pgsql
import (
"context"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetPolicies(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into policies (id, public_id, account_id, enabled, source_posture_checks)
values('policy-1','policy-1-public','account-1',true,'["posture-checks-1","posture-checks-2"]')`)
execQuery(t, ctx,
`insert into policy_rules (id, policy_id, enabled, action, protocol, bidirectional, sources, destinations,
source_resource, destination_resource, ports, port_ranges,
authorized_groups, authorized_user)
values('policy-1-rule-1','policy-1',true,'accept','tcp',true,'["group-one-resource-id","group-two-resources-id"]','["group-one-resource-id","group-two-resources-id"]',
'{"ID":"host-id-1","Type":"host"}','{"ID":"domain-1","Type":"domain"}','["8080","8443"]', '[{"Start":8080,"End":8090}]',
'{"group-one-resource-id":["user-1", "user-2"]}','user-3')`)
execQuery(t, ctx,
`insert into policies (id, public_id, account_id, enabled, source_posture_checks)
values('policy-2','policy-2-public','account-1',true,'["posture-checks-3","posture-checks-4"]')`)
execQuery(t, ctx,
`insert into policy_rules (id, policy_id, enabled, action, protocol, bidirectional, sources, destinations,
source_resource, destination_resource, ports, port_ranges,
authorized_groups, authorized_user)
values('policy-2-rule-1','policy-2',true,'accept','tcp',true,'["group-one-resource-id"]','["group-two-resources-id"]',
'{"ID":"host-id-3","Type":"host"}','{"ID":"domain-3","Type":"domain"}','["8080","8443"]', '[{"Start":8080,"End":8090}]',
'{"group-one-resource-id":["user-6", "user-7"]}','user-8')`)
// policy with a rule with null fields
execQuery(t, ctx,
`insert into policies (id, public_id, account_id, enabled, source_posture_checks)
values('policy-3','policy-3-public','account-1',true,null)`)
execQuery(t, ctx,
`insert into policy_rules (id, policy_id, enabled, action, protocol, bidirectional, sources, destinations,
source_resource, destination_resource, ports, port_ranges,
authorized_groups, authorized_user)
values('policy-3-rule-1','policy-3',true,null,null,null,null,null,null,null,null,null,null,null)`)
// policy with a disabled rule, destination resource and groups should not be in indexes
execQuery(t, ctx,
`insert into policies (id, public_id, account_id, enabled, source_posture_checks)
values('policy-4','policy-4-public','account-1',true,null)`)
execQuery(t, ctx,
`insert into policy_rules (id, policy_id, enabled, action, protocol, bidirectional, sources, destinations,
source_resource, destination_resource, ports, port_ranges,
authorized_groups, authorized_user)
values('policy-4-rule-1','policy-4',false,null,null,null,null,'["group-two-resources-id"]',
null,'{"ID":"domain-3","Type":"domain"}',null,null,null,null)`)
policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := networkmap_pgsql.GetPoliciesViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Contains(t, policies, nmdata.Policy{
ID: "policy-1",
PublicID: "policy-1-public",
Enabled: true,
SourcePostureChecks: []string{"posture-checks-1", "posture-checks-2"},
Rules: []*nmdata.PolicyRule{
{
ID: "policy-1",
PolicyID: "policy-1",
Enabled: true,
Action: "accept",
Protocol: "tcp",
Bidirectional: true,
Sources: []string{"group-one-resource-id", "group-two-resources-id"},
Destinations: []string{"group-one-resource-id", "group-two-resources-id"},
SourceResource: nmdata.Resource{ID: "host-id-1", Type: "host"},
DestinationResource: nmdata.Resource{ID: "domain-1", Type: "domain"},
Ports: []string{"8080", "8443"},
PortRanges: []nmdata.RulePortRange{{Start: 8080, End: 8090}},
AuthorizedGroups: map[string][]string{"group-one-resource-id": {"user-1", "user-2"}},
AuthorizedUser: "user-3",
},
},
})
assert.Contains(t, policies, nmdata.Policy{
ID: "policy-2",
PublicID: "policy-2-public",
Enabled: true,
SourcePostureChecks: []string{"posture-checks-3", "posture-checks-4"},
Rules: []*nmdata.PolicyRule{
{
ID: "policy-2",
PolicyID: "policy-2",
Enabled: true,
Action: "accept",
Protocol: "tcp",
Bidirectional: true,
Sources: []string{"group-one-resource-id"},
Destinations: []string{"group-two-resources-id"},
SourceResource: nmdata.Resource{ID: "host-id-3", Type: "host"},
DestinationResource: nmdata.Resource{ID: "domain-3", Type: "domain"},
Ports: []string{"8080", "8443"},
PortRanges: []nmdata.RulePortRange{{Start: 8080, End: 8090}},
AuthorizedGroups: map[string][]string{"group-one-resource-id": {"user-6", "user-7"}},
AuthorizedUser: "user-8",
},
},
})
assert.Contains(t, policies, nmdata.Policy{
ID: "policy-3",
PublicID: "policy-3-public",
Enabled: true,
SourcePostureChecks: []string{},
Rules: []*nmdata.PolicyRule{
{
ID: "policy-3",
PolicyID: "policy-3",
Enabled: true,
},
},
})
assert.Contains(t, policies, nmdata.Policy{
ID: "policy-4",
PublicID: "policy-4-public",
Enabled: true,
SourcePostureChecks: []string{},
Rules: []*nmdata.PolicyRule{
{
ID: "policy-4",
PolicyID: "policy-4",
Enabled: false,
Destinations: []string{"group-two-resources-id"},
DestinationResource: nmdata.Resource{ID: "domain-3", Type: "domain"},
},
},
})
assert.Equal(t, policyToDestinationGroupIdx, map[string]map[string]any{
"policy-1": {"group-one-resource-id": struct{}{}, "group-two-resources-id": struct{}{}},
"policy-2": {"group-two-resources-id": struct{}{}},
})
assert.Equal(t, policyToDestinationResourceIdx, map[string]map[string]any{
"policy-1": {"domain-1": struct{}{}},
"policy-2": {"domain-3": struct{}{}},
})
}

View File

@@ -0,0 +1,60 @@
package networkmap_pgsql
import (
"context"
"net/netip"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetPostureChecks(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into posture_checks (id, account_id, public_id, checks)
VALUES('posturecheck-1','account-1','posturecheck-1-public',
'{"NBVersionCheck":{"MinVersion":"0.25.0"},
"OSVersionCheck":{"Darwin":{"MinVersion":"12.0"}},
"GeoLocationCheck":{"Locations":[{"CountryCode":"FI","CityName":""}],"Action":"allow"},
"PeerNetworkRangeCheck":{"Action":"deny","Ranges":["192.168.0.1/24"]}}')`)
execQuery(t, ctx,
`insert into posture_checks (id, account_id, public_id, checks)
VALUES('posturecheck-2','account-1','posturecheck-2-public',
'{"NBVersionCheck":{"MinVersion":"0.25.0"},
"OSVersionCheck":{"Android":{"MinVersion":"0"}},
"GeoLocationCheck":{"Locations":[{"CountryCode":"US","CityName":"Harker Heights"}],"Action":"allow"},
"PeerNetworkRangeCheck":{"Action":"allow","Ranges":["0.0.0.0/0"]}}')`)
execQuery(t, ctx,
`insert into posture_checks (id, account_id, public_id, checks)
VALUES('posturecheck-3','account-1','posturecheck-3-public', null)`)
postureChecks, idToPublicIDIdx, err := networkmap_pgsql.GetPostureChecksViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Equal(t, idToPublicIDIdx, map[string]string{
"posturecheck-1": "posturecheck-1-public",
"posturecheck-2": "posturecheck-2-public",
"posturecheck-3": "posturecheck-3-public",
})
assert.Contains(t, postureChecks, nmdata.PostureChecks{
ID: "posturecheck-1",
Checks: nmdata.ChecksDefinition{
NBVersionCheck: &nmdata.NBVersionCheck{MinVersion: "0.25.0"},
OSVersionCheck: &nmdata.OSVersionCheck{Darwin: &nmdata.MinVersionCheck{MinVersion: "12.0"}},
GeoLocationCheck: &nmdata.GeoLocationCheck{Locations: []nmdata.GeoLocation{{CountryCode: "FI"}}, Action: "allow"},
PeerNetworkRangeCheck: &nmdata.PeerNetworkRangeCheck{Action: "deny", Ranges: []netip.Prefix{netip.MustParsePrefix("192.168.0.1/24")}},
}})
assert.Contains(t, postureChecks, nmdata.PostureChecks{
ID: "posturecheck-2",
Checks: nmdata.ChecksDefinition{
NBVersionCheck: &nmdata.NBVersionCheck{MinVersion: "0.25.0"},
OSVersionCheck: &nmdata.OSVersionCheck{Android: &nmdata.MinVersionCheck{MinVersion: "0"}},
GeoLocationCheck: &nmdata.GeoLocationCheck{Locations: []nmdata.GeoLocation{{CountryCode: "US", CityName: "Harker Heights"}}, Action: "allow"},
PeerNetworkRangeCheck: &nmdata.PeerNetworkRangeCheck{Action: "allow", Ranges: []netip.Prefix{netip.MustParsePrefix("0.0.0.0/0")}},
}})
assert.Contains(t, postureChecks, nmdata.PostureChecks{
ID: "posturecheck-3"})
}

View File

@@ -0,0 +1,86 @@
package networkmap_pgsql
import (
"context"
"net/netip"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetRoutes(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into routes (id, account_id, public_id, network, domains, keep_route, net_id, description,
peer, peer_groups, network_type, masquerade, metric, enabled,
groups, access_control_groups, skip_auto_apply)
VALUES('route-1','account-1','route-1-public','"172.0.0.0/16"','["test-1.com"]',true,'route-1-net-id','route-1',
'peer-id-1','["group-one-resource-id"]',1,true,9999,true,
'["group-one-resource-id"]','["group-one-resource-id"]',false)`)
execQuery(t, ctx,
`insert into routes (id, account_id, public_id, network, domains, keep_route, net_id, description,
peer, peer_groups, network_type, masquerade, metric, enabled,
groups, access_control_groups, skip_auto_apply)
VALUES('route-2','account-1','route-2-public','"172.10.0.0/16"','["test-1.com","test-2.com"]',true,'route-2-net-id','route-2',
'peer-id-2','["group-two-resources-id"]',1,true,9999,true,
'["group-two-resources-id"]','["group-two-resources-id"]',false)`)
execQuery(t, ctx,
`insert into routes (id, account_id, public_id, network, domains, keep_route, net_id, description,
peer, peer_groups, network_type, masquerade, metric, enabled,
groups, access_control_groups, skip_auto_apply)
VALUES('route-3','account-1','route-3-public',null,null,null,null,'route-3',
null,null,null,null,null,null,null,null,null)`)
routes, err := networkmap_pgsql.GetRoutesViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Contains(t, routes, nmdata.Route{
ID: "route-1",
AccountID: "account-1",
PublicID: "route-1-public",
Network: netip.MustParsePrefix("172.0.0.0/16"),
Domains: domain.List{"test-1.com"},
KeepRoute: true,
NetID: "route-1-net-id",
Description: "route-1",
Peer: "peer-id-1",
PeerID: "peer-id-1",
PeerGroups: []string{"group-one-resource-id"},
NetworkType: 1,
Masquerade: true,
Metric: 9999,
Enabled: true,
Groups: []string{"group-one-resource-id"},
AccessControlGroups: []string{"group-one-resource-id"},
SkipAutoApply: false,
})
assert.Contains(t, routes, nmdata.Route{
ID: "route-2",
AccountID: "account-1",
PublicID: "route-2-public",
Network: netip.MustParsePrefix("172.10.0.0/16"),
Domains: domain.List{"test-1.com", "test-2.com"},
KeepRoute: true,
NetID: "route-2-net-id",
Description: "route-2",
Peer: "peer-id-2",
PeerID: "peer-id-2",
PeerGroups: []string{"group-two-resources-id"},
NetworkType: 1,
Masquerade: true,
Metric: 9999,
Enabled: true,
Groups: []string{"group-two-resources-id"},
AccessControlGroups: []string{"group-two-resources-id"},
SkipAutoApply: false,
})
assert.Contains(t, routes, nmdata.Route{
ID: "route-3",
AccountID: "account-1",
PublicID: "route-3-public",
Description: "route-3",
})
}

View File

@@ -0,0 +1,110 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"testing"
"github.com/stretchr/testify/assert"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
)
func TestGetPrivateServicesViaPgxConnection(t *testing.T) {
ctx := context.TODO()
_, err := pgstore.Pool.Exec(ctx,
`insert into services (id, account_id, enabled, private, access_groups, proxy_cluster, domain)
values('service-1','account-1',true,true,'["group-one-resource-id"]','test-1.com','test-2.com')`)
assert.NoError(t, err)
execQuery(t, ctx,
`insert into services (id, account_id, enabled, private, access_groups, proxy_cluster, domain)
values('service-2','account-1',true,true,'["group-one-resource-id","group-two-resources-id"]','test-3.com','test-4.com')`)
assert.NoError(t, err)
execQuery(t, ctx,
`insert into services (id, account_id, enabled, private, access_groups, proxy_cluster, domain)
values('service-3','account-1',null,null,null,null,null)`)
assert.NoError(t, err)
services, err := networkmap_pgsql.GetPrivateServicesViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Contains(t, services, networkmap_pgsql.Service{
Enabled: sql.NullBool{Bool: true, Valid: true},
Private: sql.NullBool{Bool: true, Valid: true},
AccessGroups: []string{"group-one-resource-id"},
ProxyCluster: sql.NullString{String: "test-1.com", Valid: true},
Domain: sql.NullString{String: "test-2.com", Valid: true},
})
assert.Contains(t, services, networkmap_pgsql.Service{
Enabled: sql.NullBool{Bool: true, Valid: true},
Private: sql.NullBool{Bool: true, Valid: true},
AccessGroups: []string{"group-one-resource-id", "group-two-resources-id"},
ProxyCluster: sql.NullString{String: "test-3.com", Valid: true},
Domain: sql.NullString{String: "test-4.com", Valid: true},
})
assert.Contains(t, services, networkmap_pgsql.Service{
Enabled: sql.NullBool{Bool: false, Valid: false},
Private: sql.NullBool{Bool: false, Valid: false},
AccessGroups: []string{},
ProxyCluster: sql.NullString{String: "", Valid: false},
Domain: sql.NullString{String: "", Valid: false},
})
}
func TestGetProxyTargetedDomainResourceIDsViaPgxConnection(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into services (id, account_id, enabled, terminated)
values('service-4','account-1',true,false)`)
execQuery(t, ctx,
`insert into targets (target_id, account_id, service_id, enabled, target_type)
values('target-1','account-1','service-4',true,'domain')`)
// id shouldn't be returned as the taget_type is not "domain"
execQuery(t, ctx,
`insert into targets (target_id, account_id, service_id, enabled, target_type)
values('target-2','account-1','service-4',true,'cluster')`)
// id shouldn't be included as the target is disabled
execQuery(t, ctx,
`insert into targets (target_id, account_id, service_id, enabled, target_type)
values('target-3','account-1','service-4',false,'domain')`)
// id shouldn't be included as the service is disabled
execQuery(t, ctx,
`insert into services (id, account_id, enabled, terminated)
values('service-5','account-1',false,false)`)
execQuery(t, ctx,
`insert into targets (target_id, account_id, service_id, enabled, target_type)
values('target-4','account-1','service-5',false,'domain')`)
// id shouldn't be included as the service is terminated (explicitly)
execQuery(t, ctx,
`insert into services (id, account_id, enabled, terminated)
values('service-6','account-1',true,true)`)
execQuery(t, ctx,
`insert into targets (target_id, account_id, service_id, enabled, target_type)
values('target-5','account-1','service-6',true,'domain')`)
// id shouldn't be included as the service is terminated (implicitly)
execQuery(t, ctx,
`insert into services (id, account_id, enabled, terminated)
values('service-7','account-1',true,null)`)
execQuery(t, ctx,
`insert into targets (target_id, account_id, service_id, enabled, target_type)
values('target-6','account-1','service-7',true,'domain')`)
execQuery(t, ctx,
`insert into services (id, account_id, enabled, terminated)
values('service-8','account-1',true,false)`)
execQuery(t, ctx,
`insert into targets (target_id, account_id, service_id, enabled, target_type)
values('target-7','account-1','service-8',true,'domain')`)
// id shouldn't be returned as the taget_id is null
execQuery(t, ctx,
`insert into targets (target_id, account_id, service_id, enabled, target_type)
values(null,'account-1','service-4',true,'cluster')`)
servtargetedDomains, err := networkmap_pgsql.GetProxyTargetedDomainResourceIDsViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Equal(t, servtargetedDomains, map[string]struct{}{
"target-1": {},
"target-6": {},
"target-7": {},
})
}

View File

@@ -0,0 +1,56 @@
package networkmap_pgsql
import (
"context"
"testing"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
"github.com/stretchr/testify/assert"
)
func TestGetAllowedUsers(t *testing.T) {
ctx := context.TODO()
execQuery(t, ctx,
`insert into users (id, name, account_id, auto_groups, blocked, is_service_user)
VALUES('user-1','user-1','account-1','["group-one-resource-id"]',false,false)`)
execQuery(t, ctx,
`insert into users (id, name, account_id, auto_groups, blocked, is_service_user)
VALUES('user-2','user-2','account-1','["group-one-resource-id","group-two-resources-id"]',false,false)`)
execQuery(t, ctx,
`insert into users (id, name, account_id, auto_groups, blocked, is_service_user)
VALUES('user-3','user-3','account-1','["group-two-resources-id"]',false,false)`)
// shouldn't be included as it's blocked
execQuery(t, ctx,
`insert into users (id, name, account_id, auto_groups, blocked, is_service_user)
VALUES('user-4','user-4','account-1','["group-two-resources-id"]',true,false)`)
// shouldn't be included as it's a service_user
execQuery(t, ctx,
`insert into users (id, name, account_id, auto_groups, blocked, is_service_user)
VALUES('user-5','user-5','account-1','["group-two-resources-id"]',false,true)`)
execQuery(t, ctx,
`insert into groups (id, name, account_id)
VALUES('all-group-1','All','account-1')`)
execQuery(t, ctx,
`insert into groups (id, name, account_id)
VALUES('all-group-2','All','account-1')`)
execQuery(t, ctx,
`insert into groups (id, name, account_id)
VALUES('all-group-3','All','account-1')`)
userIdx, groupIdToUserIds, err := networkmap_pgsql.GetAllowedUsersViaPgxConnection(ctx, conn(t, ctx), "account-1")
assert.NoError(t, err)
assert.Equal(t, userIdx, map[string]struct{}{
"user-1": {},
"user-2": {},
"user-3": {},
})
assert.Equal(t, groupIdToUserIds, map[string][]string{
"group-one-resource-id": {"user-1", "user-2"},
"group-two-resources-id": {"user-2", "user-3"},
"all-group-1": {"user-1", "user-2", "user-3"},
"all-group-2": {"user-1", "user-2", "user-3"},
"all-group-3": {"user-1", "user-2", "user-3"},
})
}

View File

@@ -18,6 +18,7 @@ import (
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache"
"github.com/netbirdio/netbird/management/internals/modules/peers/ephemeral"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/management/internals/server/config"
"github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/account"
@@ -30,6 +31,8 @@ import (
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/management/server/types"
sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/management/status"
"github.com/netbirdio/netbird/util"
@@ -61,6 +64,8 @@ type Controller struct {
serverSupportedSyncMessageVersion sharedgrpc.SyncMessageVersion
perAccountServerSupportedSyncMessageVersions map[string]sharedgrpc.SyncMessageVersion
nmdataStore *networkmapdb.NetworkMapDBStoreImpl
}
type bufferUpdate struct {
@@ -78,7 +83,7 @@ type bufferAffectedUpdate struct {
var _ network_map.Controller = (*Controller)(nil)
func NewController(ctx context.Context, store store.Store, metrics telemetry.AppMetrics, peersUpdateManager network_map.PeersUpdateManager, requestBuffer account.RequestBuffer, integratedPeerValidator integrated_validator.IntegratedValidator, settingsManager settings.Manager, dnsDomain string, proxyController port_forwarding.Controller, ephemeralPeersManager ephemeral.Manager, config *config.Config) *Controller {
func NewController(ctx context.Context, store store.Store, metrics telemetry.AppMetrics, peersUpdateManager network_map.PeersUpdateManager, requestBuffer account.RequestBuffer, integratedPeerValidator integrated_validator.IntegratedValidator, settingsManager settings.Manager, dnsDomain string, proxyController port_forwarding.Controller, ephemeralPeersManager ephemeral.Manager, config *config.Config, nmdataStore *networkmapdb.NetworkMapDBStoreImpl) *Controller {
nMetrics, err := newMetrics(metrics.UpdateChannelMetrics())
if err != nil {
log.Fatal(fmt.Errorf("error creating metrics: %w", err))
@@ -99,6 +104,7 @@ func NewController(ctx context.Context, store store.Store, metrics telemetry.App
EphemeralPeersManager: ephemeralPeersManager,
serverSupportedSyncMessageVersion: sharedgrpc.SyncMessageVersionFromConfig(config.HighestSupportedSyncMessageVersion),
perAccountServerSupportedSyncMessageVersions: sharedgrpc.SyncMessageVersionsFromMap(config.PerAccountHighestSupportedSyncMessageVersion),
nmdataStore: nmdataStore,
}
}
@@ -147,6 +153,11 @@ func (c *Controller) CountStreams() int {
func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID string, reason types.UpdateReason) error {
log.WithContext(ctx).Tracef("updating peers for account %s from %s", accountID, util.GetCallerName())
if nmData := c.getNetworkMapData(ctx, accountID); nmData != nil {
return c.sendUpdateAccountPeersFromData(ctx, accountID, reason, nmData)
}
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to get account: %v", err)
@@ -167,7 +178,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
return nil
}
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validate peers: %v", err)
}
@@ -254,7 +265,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
// proxyNetworkMap rides the envelope as a ProxyPatch sidecar;
// the client merges it into Calculate()'s output the same
// way the legacy server did via NetworkMap.Merge.
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort)
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(p), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -275,7 +286,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
}
start = time.Now()
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, nmap, dnsDomain, postureChecks, dnsCache, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort)
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(p), nil, nil, nmap, dnsDomain, postureChecks, dnsCache, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
c.metrics.CountToSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -293,6 +304,259 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
return nil
}
// sendUpdateAccountPeersFromData is the account-free variant of
// sendUpdateAccountPeers: everything is computed from the network-map DB
// store's twin data; only extra settings and validated peers are resolved at
// runtime. Proxy network maps and policy injection, private-service zones,
// group-to-user SSH mappings and forced routing-peer DNS resolution have no
// DB-backed source yet and are omitted.
func (c *Controller) sendUpdateAccountPeersFromData(ctx context.Context, accountID string, reason types.UpdateReason, nmData *networkmap.NetworkMapData) error {
peersToUpdate := c.connectedPeersFromData(nmData, nil)
if len(peersToUpdate) == 0 {
return nil
}
return c.sendUpdatesFromData(ctx, accountID, nmData, peersToUpdate, &reason)
}
// sendUpdateForAffectedPeersFromData is the account-free variant of
// sendUpdateForAffectedPeers.
func (c *Controller) sendUpdateForAffectedPeersFromData(ctx context.Context, accountID string, peerIDs []string, nmData *networkmap.NetworkMapData) error {
if len(peerIDs) == 0 {
log.WithContext(ctx).Tracef("sendUpdateForAffectedPeersFromData: no affected peers")
return nil
}
peersToUpdate := c.connectedPeersFromData(nmData, peerIDs)
if len(peersToUpdate) == 0 {
log.WithContext(ctx).Tracef("sendUpdateForAffectedPeersFromData: no peers to update (affected peers not found in data or no channels)")
return nil
}
log.WithContext(ctx).Tracef("sendUpdateForAffectedPeersFromData: sending network map to %d connected peers", len(peersToUpdate))
return c.sendUpdatesFromData(ctx, accountID, nmData, peersToUpdate, nil)
}
// connectedPeersFromData returns the peers with an open update channel. An
// empty affected list means all peers; a non-empty list restricts the result
// to those peer IDs.
func (c *Controller) connectedPeersFromData(nmData *networkmap.NetworkMapData, affected []string) []*nmdata.Peer {
if len(affected) == 0 {
result := make([]*nmdata.Peer, 0, len(nmData.Peers))
for _, peer := range nmData.Peers {
if c.peersUpdateManager.HasChannel(peer.ID) {
result = append(result, peer)
}
}
return result
}
result := make([]*nmdata.Peer, 0, len(affected))
for _, peerID := range affected {
peer := nmData.Peers[peerID]
if peer == nil {
continue
}
if c.peersUpdateManager.HasChannel(peerID) {
result = append(result, peer)
}
}
return result
}
func (c *Controller) sendUpdatesFromData(ctx context.Context, accountID string, nmData *networkmap.NetworkMapData, peersToUpdate []*nmdata.Peer, reason *types.UpdateReason) error {
globalStart := time.Now()
extraSettings, err := c.settingsManager.GetExtraSettings(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to get flow enabled status: %v", err)
}
dnsCache := &cache.DNSConfigCache{}
dnsDomain := c.getDNSDomainFromData(nmData.AccountSettings)
peersCustomZone := networkmap.PeersCustomZone(ctx, accountID, dnsDomain, nmData.Peers, ipv6AllowedPeersFromData(nmData))
dnsFwdPort := computeForwarderPortFromData(nmData.Peers, network_map.DnsForwarderPortMinVersion)
var wg sync.WaitGroup
semaphore := make(chan struct{}, 10)
for _, peer := range peersToUpdate {
if reason != nil && c.accountManagerMetrics != nil {
c.accountManagerMetrics.CountNmapTriggered(string(reason.Resource), string(reason.Operation))
}
wg.Add(1)
semaphore <- struct{}{}
go func(p *nmdata.Peer) {
defer wg.Done()
defer func() { <-semaphore }()
start := time.Now()
postureChecks := peerPostureChecksFromData(nmData, p.ID)
c.metrics.CountCalcPostureChecksDuration(time.Since(start))
start = time.Now()
peerGroups := maps.Keys(nmData.GetPeerGroups(p.ID))
var update *proto.SyncResponse
commonSyncMessageVersion := sharedgrpc.HighestCommonSyncMessageVersion(
c.perAccountOrGlobalSupportedSyncMessageVersions(accountID),
sharedgrpc.SyncMessageVersionFromConfig(&p.Meta.SyncMessageVersion))
log.WithContext(ctx).
WithFields(log.Fields{
"sync_message_version": commonSyncMessageVersion,
"server_sync_message_version": c.perAccountOrGlobalSupportedSyncMessageVersions(accountID),
"peer_sync_message_version": sharedgrpc.SyncMessageVersionFromConfig(&p.Meta.SyncMessageVersion),
}).Debug("common highest sync message version")
if commonSyncMessageVersion == sharedgrpc.ComponentNetworkMap {
components := nmData.GetPeerNetworkMapComponents(p.ID, peersCustomZone)
c.metrics.CountCalcPeerNetworkMapDuration(time.Since(start))
start = time.Now()
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, components, nil, dnsDomain, postureChecks, nmData.AccountSettings, extraSettings, peerGroups, dnsFwdPort)
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
Update: update,
MessageType: network_map.MessageTypeNetworkMap,
})
return
}
nmap := networkMapFromData(ctx, nmData, p.ID, peersCustomZone)
c.metrics.CountCalcPeerNetworkMapDuration(time.Since(start))
start = time.Now()
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, nmap, dnsDomain, postureChecks, dnsCache, nmData.AccountSettings, extraSettings, peerGroups, dnsFwdPort)
c.metrics.CountToSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
Update: update,
MessageType: network_map.MessageTypeNetworkMap,
})
}(peer)
}
wg.Wait()
if c.accountManagerMetrics != nil {
c.accountManagerMetrics.CountUpdateAccountPeersDuration(time.Since(globalStart))
}
return nil
}
func (c *Controller) getNetworkMapData(ctx context.Context, accountID string) *networkmap.NetworkMapData {
if c.nmdataStore == nil {
return nil
}
nmData, err := c.nmdataStore.GetNetworkMapData(ctx, accountID)
if err != nil {
log.WithContext(ctx).Errorf("failed to get network map data for account %s, falling back to account-based computation: %v", accountID, err)
return nil
}
return nmData
}
func (c *Controller) getDNSDomainFromData(settings *nmdata.AccountSettingsInfo) string {
if settings == nil || settings.DNSDomain == "" {
return c.dnsDomain
}
return settings.DNSDomain
}
func ipv6AllowedPeersFromData(nmData *networkmap.NetworkMapData) map[string]struct{} {
result := make(map[string]struct{})
if nmData.AccountSettings != nil {
for _, groupID := range nmData.AccountSettings.IPv6EnabledGroups {
group := nmData.Groups[groupID]
if group == nil {
continue
}
for _, peerID := range group.Peers {
result[peerID] = struct{}{}
}
}
}
for id, p := range nmData.Peers {
if p != nil && p.ProxyMeta.Embedded {
result[id] = struct{}{}
}
}
return result
}
func networkMapFromData(ctx context.Context, nmData *networkmap.NetworkMapData, peerID string, peersCustomZone nmdata.CustomZone) *types.NetworkMap {
components := nmData.GetPeerNetworkMapComponents(peerID, peersCustomZone)
if components.IsEmpty() {
return &types.NetworkMap{Network: components.Network}
}
return types.CalculateNetworkMapFromComponents(ctx, components)
}
// peerPostureChecksFromData mirrors getPeerPostureChecks on the twin store. The
// sync response only encodes process-check file paths, so only ProcessCheck is
// converted back to the server posture type.
func peerPostureChecksFromData(nmData *networkmap.NetworkMapData, peerID string) []*posture.Checks {
if len(nmData.PostureChecks) == 0 {
return nil
}
peerPostureChecks := make(map[string]*posture.Checks)
for _, policy := range nmData.Policies {
if policy == nil || !policy.Enabled || len(policy.SourcePostureChecks) == 0 {
continue
}
if !isPeerInPolicySourceGroupsFromData(nmData, peerID, policy) {
continue
}
for _, checkID := range policy.SourcePostureChecks {
twin := nmData.PostureChecks[checkID]
if twin == nil {
continue
}
peerPostureChecks[checkID] = postureChecksFromTwin(twin)
}
}
return maps.Values(peerPostureChecks)
}
func isPeerInPolicySourceGroupsFromData(nmData *networkmap.NetworkMapData, peerID string, policy *nmdata.Policy) bool {
for _, rule := range policy.Rules {
if rule == nil || !rule.Enabled {
continue
}
for _, groupID := range rule.Sources {
if group := nmData.Groups[groupID]; group != nil && slices.Contains(group.Peers, peerID) {
return true
}
}
}
return false
}
func postureChecksFromTwin(twin *nmdata.PostureChecks) *posture.Checks {
checks := &posture.Checks{ID: twin.ID}
if twin.Checks.ProcessCheck != nil {
processes := make([]posture.Process, 0, len(twin.Checks.ProcessCheck.Processes))
for _, p := range twin.Checks.ProcessCheck.Processes {
processes = append(processes, posture.Process{LinuxPath: p.LinuxPath, MacPath: p.MacPath, WindowsPath: p.WindowsPath})
}
checks.Checks.ProcessCheck = &posture.ProcessCheck{Processes: processes}
}
return checks
}
func (c *Controller) perAccountOrGlobalSupportedSyncMessageVersions(accountId string) sharedgrpc.SyncMessageVersion {
if perAccount, ok := c.perAccountServerSupportedSyncMessageVersions[accountId]; ok {
return perAccount
@@ -325,6 +589,10 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
return nil
}
if nmData := c.getNetworkMapData(ctx, accountID); nmData != nil {
return c.sendUpdateForAffectedPeersFromData(ctx, accountID, peerIDs, nmData)
}
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
if err != nil {
return fmt.Errorf("failed to get account: %v", err)
@@ -340,7 +608,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
log.WithContext(ctx).Tracef("sendUpdateForAffectedPeers: sending network map to %d connected peers", len(peersToUpdate))
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validate peers: %v", err)
}
@@ -426,7 +694,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
// proxyNetworkMap rides the envelope as a ProxyPatch sidecar;
// the client merges it into Calculate()'s output the same
// way the legacy server did via NetworkMap.Merge.
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort)
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(p), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -447,7 +715,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
}
start = time.Now()
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, p, nil, nil, nmap, dnsDomain, postureChecks, dnsCache, account.Settings, extraSetting, maps.Keys(peerGroups), dnsFwdPort)
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(p), nil, nil, nmap, dnsDomain, postureChecks, dnsCache, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
c.metrics.CountToSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -504,7 +772,7 @@ func (c *Controller) UpdateAccountPeer(ctx context.Context, accountId string, pe
return fmt.Errorf("peer %s doesn't exists in account %s", peerId, accountId)
}
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validated peers: %v", err)
}
@@ -564,7 +832,7 @@ func (c *Controller) UpdateAccountPeer(ctx context.Context, accountId string, pe
// proxyNetworkMap rides the envelope as a ProxyPatch sidecar;
// the client merges it into Calculate()'s output the same
// way the legacy server did via NetworkMap.Merge.
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, peer, nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, account.Settings, extraSettings, maps.Keys(peerGroups), dnsFwdPort)
update = grpc.ToComponentSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(peer), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSettings, maps.Keys(peerGroups), dnsFwdPort)
c.peersUpdateManager.SendUpdate(ctx, peer.ID, &network_map.UpdateMessage{
Update: update,
@@ -581,7 +849,7 @@ func (c *Controller) UpdateAccountPeer(ctx context.Context, accountId string, pe
nmap.Merge(proxyNetworkMap)
}
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, peer, nil, nil, nmap, dnsDomain, postureChecks, dnsCache, account.Settings, extraSettings, maps.Keys(peerGroups), dnsFwdPort)
update = grpc.ToSyncResponse(ctx, nil, c.config.HttpConfig, c.config.DeviceAuthorizationFlow, types.TwinPeer(peer), nil, nil, nmap, dnsDomain, postureChecks, dnsCache, types.TwinAccountSettings(account.Settings), extraSettings, maps.Keys(peerGroups), dnsFwdPort)
c.peersUpdateManager.SendUpdate(ctx, peer.ID, &network_map.UpdateMessage{
Update: update,
@@ -641,7 +909,11 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
if err != nil {
return nil, nil, nil, nil, 0, err
}
return peer, &types.NetworkMapComponents{Network: network.Copy()}, nil, nil, 0, nil
return peer, &types.NetworkMapComponents{Network: types.TwinNetwork(network)}, nil, nil, 0, nil
}
if nmData := c.getNetworkMapData(ctx, accountID); nmData != nil {
return c.getValidatedPeerWithComponentsFromData(ctx, accountID, peer, nmData)
}
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
@@ -651,7 +923,7 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
c.injectAllProxyPolicies(ctx, account)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
if err != nil {
return nil, nil, nil, nil, 0, err
}
@@ -688,6 +960,21 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
return peer, components, proxyNetworkMaps[peer.ID], postureChecks, dnsFwdPort, nil
}
// getValidatedPeerWithComponentsFromData is the account-free variant of
// GetValidatedPeerWithComponents. The proxy network map fragment is omitted
// like on the other nmdata paths.
func (c *Controller) getValidatedPeerWithComponentsFromData(ctx context.Context, accountID string, peer *nbpeer.Peer, nmData *networkmap.NetworkMapData) (*nbpeer.Peer, *types.NetworkMapComponents, *types.NetworkMap, []*posture.Checks, int64, error) {
postureChecks := peerPostureChecksFromData(nmData, peer.ID)
dnsDomain := c.getDNSDomainFromData(nmData.AccountSettings)
peersCustomZone := networkmap.PeersCustomZone(ctx, accountID, dnsDomain, nmData.Peers, ipv6AllowedPeersFromData(nmData))
components := nmData.GetPeerNetworkMapComponents(peer.ID, peersCustomZone)
dnsFwdPort := computeForwarderPortFromData(nmData.Peers, network_map.DnsForwarderPortMinVersion)
return peer, components, nil, postureChecks, dnsFwdPort, nil
}
// BufferUpdateAffectedPeers accumulates peer IDs and flushes them after the buffer interval.
func (c *Controller) BufferUpdateAffectedPeers(ctx context.Context, accountID string, peerIDs []string, reason types.UpdateReason) error {
if len(peerIDs) == 0 {
@@ -794,11 +1081,15 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
}
emptyMap := &types.NetworkMap{
Network: network.Copy(),
Network: types.TwinNetwork(network),
}
return emptyMap, nil, 0, nil
}
if nmData := c.getNetworkMapData(ctx, accountID); nmData != nil {
return c.getValidatedPeerWithMapFromData(ctx, accountID, peerID, nmData)
}
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
if err != nil {
return nil, nil, 0, err
@@ -806,7 +1097,7 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
c.injectAllProxyPolicies(ctx, account)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
if err != nil {
return nil, nil, 0, err
}
@@ -846,6 +1137,21 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
return networkMap, postureChecks, dnsFwdPort, nil
}
// getValidatedPeerWithMapFromData is the account-free variant of
// GetValidatedPeerWithMap. The proxy network map fragment is omitted like on
// the other nmdata paths.
func (c *Controller) getValidatedPeerWithMapFromData(ctx context.Context, accountID string, peerID string, nmData *networkmap.NetworkMapData) (*types.NetworkMap, []*posture.Checks, int64, error) {
postureChecks := peerPostureChecksFromData(nmData, peerID)
dnsDomain := c.getDNSDomainFromData(nmData.AccountSettings)
peersCustomZone := networkmap.PeersCustomZone(ctx, accountID, dnsDomain, nmData.Peers, ipv6AllowedPeersFromData(nmData))
networkMap := networkMapFromData(ctx, nmData, peerID, peersCustomZone)
dnsFwdPort := computeForwarderPortFromData(nmData.Peers, network_map.DnsForwarderPortMinVersion)
return networkMap, postureChecks, dnsFwdPort, nil
}
// GetDNSDomain returns the configured dnsDomain
func (c *Controller) GetDNSDomain(settings *types.Settings) string {
if settings == nil {
@@ -908,20 +1214,36 @@ func (c *Controller) StartWarmup(ctx context.Context) {
// computeForwarderPort checks if all peers in the account have updated to a specific version or newer.
// If all peers have the required version, it returns the new well-known port (22054), otherwise returns 0.
func computeForwarderPort(peers []*nbpeer.Peer, requiredVersion string) int64 {
if len(peers) == 0 {
versions := make([]string, 0, len(peers))
for _, peer := range peers {
versions = append(versions, peer.Meta.WtVersion)
}
return computeForwarderPortFromVersions(versions, requiredVersion)
}
func computeForwarderPortFromData(peers map[string]*nmdata.Peer, requiredVersion string) int64 {
versions := make([]string, 0, len(peers))
for _, peer := range peers {
versions = append(versions, peer.Meta.WtVersion)
}
return computeForwarderPortFromVersions(versions, requiredVersion)
}
func computeForwarderPortFromVersions(wtVersions []string, requiredVersion string) int64 {
if len(wtVersions) == 0 {
return int64(network_map.OldForwarderPort)
}
reqVer := semver.Canonical(requiredVersion)
// Check if all peers have the required version or newer
for _, peer := range peers {
for _, wtVersion := range wtVersions {
// Development version is always supported
if version.IsDevelopmentVersion(peer.Meta.WtVersion) {
if version.IsDevelopmentVersion(wtVersion) {
continue
}
peerVersion := semver.Canonical("v" + peer.Meta.WtVersion)
peerVersion := semver.Canonical("v" + wtVersion)
if peerVersion == "" {
// If any peer doesn't have version info, return 0
return int64(network_map.OldForwarderPort)
@@ -1055,7 +1377,7 @@ func (c *Controller) GetNetworkMap(ctx context.Context, peerID string) (*types.N
groups[groupID] = group.Peers
}
validatedPeers, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
validatedPeers, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
if err != nil {
return nil, err
}

View File

@@ -17,6 +17,7 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/account"
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/store"
@@ -61,10 +62,23 @@ func newAgentNetworkHandlerFixture(t *testing.T) *agentNetworkHandlerFixture {
Return(true, context.Background(), nil).
AnyTimes()
manager := agentnetwork.NewManager(st, perms, nil, nil)
// Swallow activity events so the mutation paths (create/update/delete)
// are exercisable through the HTTP layer.
accounts := account.NewMockManager(ctrl)
accounts.EXPECT().
StoreEvent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).
AnyTimes()
accounts.EXPECT().
UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).
AnyTimes()
manager := agentnetwork.NewManager(st, perms, accounts, nil)
h := &handler{manager: manager}
router := mux.NewRouter()
router.HandleFunc("/agent-network/providers", h.createProvider).Methods("POST")
router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET")
router.HandleFunc("/agent-network/providers/{providerId}", h.updateProvider).Methods("PUT")
h.addPolicyEndpoints(router)
h.addConsumptionEndpoints(router)
h.addBudgetRuleEndpoints(router)

View File

@@ -1,7 +1,9 @@
package handlers
import (
"encoding/json"
"math"
nethttp "net/http"
"testing"
"github.com/stretchr/testify/assert"
@@ -51,3 +53,50 @@ func TestValidate_ModelRates(t *testing.T) {
assert.Error(t, validate(base(m), true), "case %q must be rejected", name)
}
}
// TestProviderHandler_UpdateReplacesFullState pins the update contract shared
// with the other PUT endpoints: the request replaces the provider's mutable
// state, so optional fields absent from the JSON land as their zero values.
// The two exceptions are server-side: the api_key (a secret — omitted means
// "not rotated") and the session keypair, both preserved by the manager. The
// identity headers stay on the wire as explicit empty strings so a cleared
// value round-trips.
func TestProviderHandler_UpdateReplacesFullState(t *testing.T) {
f := newAgentNetworkHandlerFixture(t)
create := `{
"provider_id": "openai_api",
"name": "openai",
"upstream_url": "https://api.openai.com",
"api_key": "sk-test",
"enabled": true,
"metadata_disabled": true,
"skip_tls_verification": true,
"extra_values": {"x-portkey-config": "pc-prod-3f2a"},
"identity_header_user_id": "x-bf-dim-netbird_user_id",
"models": [{"id": "gpt-4o", "input_per_1k": 0.0025, "output_per_1k": 0.01}]
}`
rec := f.do(t, nethttp.MethodPost, "/agent-network/providers", create)
require.Equal(t, nethttp.StatusOK, rec.Code, "create must succeed: %s", rec.Body.String())
var created api.AgentNetworkProvider
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &created))
// Minimal update: only the required fields, no api_key. Everything
// optional must land as its zero value.
update := `{"provider_id": "openai_api", "name": "openai-renamed", "upstream_url": "https://api.openai.com", "enabled": true}`
rec = f.do(t, nethttp.MethodPut, "/agent-network/providers/"+created.Id, update)
require.Equal(t, nethttp.StatusOK, rec.Code, "update without api_key must succeed (key is preserved): %s", rec.Body.String())
var updated api.AgentNetworkProvider
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &updated))
assert.Equal(t, "openai-renamed", updated.Name, "sent field must apply")
assert.True(t, updated.Enabled, "sent field must apply")
assert.False(t, updated.MetadataDisabled, "omitted metadata_disabled must land as false — PUT replaces the full state")
assert.False(t, updated.SkipTlsVerification, "omitted skip_tls_verification must land as false")
assert.Nil(t, updated.ExtraValues, "omitted extra_values must be cleared")
assert.Equal(t, "", updated.IdentityHeaderUserId, "omitted identity header must be cleared yet stay on the wire")
assert.Empty(t, updated.Models, "omitted models must be cleared")
assert.Contains(t, rec.Body.String(), `"identity_header_user_id":""`,
"cleared identity header must round-trip as an explicit empty string")
}

View File

@@ -2,7 +2,6 @@ package handlers
import (
"encoding/json"
"errors"
"net/http"
"github.com/gorilla/mux"
@@ -11,19 +10,20 @@ import (
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/shared/management/http/api"
"github.com/netbirdio/netbird/shared/management/http/util"
"github.com/netbirdio/netbird/shared/management/status"
)
// addSettingsEndpoints registers the Agent Network settings routes. The
// settings row is bootstrapped server-side on first provider create; GET reads
// it and PUT updates the mutable collection toggles (cluster/subdomain stay
// immutable).
// settings row is bootstrapped server-side on first provider create or on the
// first PUT carrying a cluster; GET reads it and PUT applies a partial update
// of the mutable collection toggles (cluster/subdomain stay immutable).
func (h *handler) addSettingsEndpoints(router *mux.Router) {
router.HandleFunc("/agent-network/settings", h.getSettings).Methods("GET", "OPTIONS")
router.HandleFunc("/agent-network/settings", h.updateSettings).Methods("PUT", "OPTIONS")
}
// updateSettings applies the collection toggles to the account's settings row.
// updateSettings replaces the mutable settings fields on the account's row.
// A request carrying a cluster bootstraps the row when the account doesn't
// have one yet.
func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) {
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
if err != nil {
@@ -48,11 +48,9 @@ func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) {
util.WriteJSONObject(r.Context(), w, updated.ToAPIResponse())
}
// getSettings returns the account's agent-network settings. The settings
// row is bootstrapped on first provider create, so freshly-onboarded
// accounts have nothing to read. Rather than 404-ing in that case (which
// the dashboard would have to special-case), return a JSON null with 200
// so consumers can branch on the body alone.
// getSettings returns the account's agent-network settings. Accounts that
// haven't been bootstrapped yet read as the defaults with an empty cluster,
// subdomain and endpoint; the manager synthesises that view.
func (h *handler) getSettings(w http.ResponseWriter, r *http.Request) {
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
if err != nil {
@@ -62,11 +60,6 @@ func (h *handler) getSettings(w http.ResponseWriter, r *http.Request) {
settings, err := h.manager.GetSettings(r.Context(), userAuth.AccountId, userAuth.UserId)
if err != nil {
var sErr *status.Error
if errors.As(err, &sErr) && sErr.Type() == status.NotFound {
util.WriteJSONObject(r.Context(), w, nil)
return
}
util.WriteError(r.Context(), err, w)
return
}

View File

@@ -0,0 +1,137 @@
package handlers
import (
"encoding/json"
"net/http"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// TestSettingsHandler_GetUnbootstrappedReturnsDefaults pins the settings-read
// convention shared with the account and DNS settings endpoints: settings
// always read as a JSON object. Before bootstrap that object carries the
// defaults with an empty cluster/subdomain/endpoint (the "not bootstrapped"
// signal) and no timestamps — never a 404 and never the legacy null body.
func TestSettingsHandler_GetUnbootstrappedReturnsDefaults(t *testing.T) {
f := newAgentNetworkHandlerFixture(t)
rec := f.do(t, http.MethodGet, "/agent-network/settings", "")
require.Equal(t, http.StatusOK, rec.Code,
"unbootstrapped account must read as 200 with defaults: got %d body=%s", rec.Code, rec.Body.String())
require.NotEqual(t, "null", trimSpace(rec.Body.String()),
"the legacy 200+null shape must not come back")
var got api.AgentNetworkSettings
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
assert.Empty(t, got.Cluster, "cluster must be empty until bootstrapped")
assert.Empty(t, got.Subdomain, "subdomain must be empty until bootstrapped")
assert.Empty(t, got.Endpoint, "endpoint must be empty until bootstrapped, not a bare dot")
assert.True(t, got.EnableLogCollection, "defaults must show log collection on, matching bootstrap")
assert.False(t, got.EnablePromptCollection, "defaults must show prompt collection off")
assert.False(t, got.RedactPii, "defaults must show redaction off")
require.NotNil(t, got.AccessLogRetentionDays)
assert.Equal(t, 30, *got.AccessLogRetentionDays, "defaults must show the bootstrap retention")
assert.Nil(t, got.CreatedAt, "no timestamps before a row exists")
assert.Nil(t, got.UpdatedAt, "no timestamps before a row exists")
}
// TestSettingsHandler_PutBootstrapsWithCluster covers the settings-first
// bootstrap path: a PUT carrying a cluster on an unbootstrapped account
// creates the row (cluster pinned, subdomain assigned) and applies the
// mutable fields from the same request.
func TestSettingsHandler_PutBootstrapsWithCluster(t *testing.T) {
f := newAgentNetworkHandlerFixture(t)
rec := f.do(t, http.MethodPut, "/agent-network/settings",
`{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": false, "access_log_retention_days": 30}`)
require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String())
var got api.AgentNetworkSettings
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
assert.Equal(t, "eu.proxy.netbird.io", got.Cluster, "cluster must be pinned from the request")
assert.NotEmpty(t, got.Subdomain, "subdomain must be assigned at bootstrap")
assert.Equal(t, got.Subdomain+".eu.proxy.netbird.io", got.Endpoint, "endpoint must combine subdomain and cluster")
assert.True(t, got.EnableLogCollection, "toggle from the bootstrap request must apply")
assert.True(t, got.EnablePromptCollection, "toggle from the bootstrap request must apply")
require.NotNil(t, got.AccessLogRetentionDays)
assert.Equal(t, 30, *got.AccessLogRetentionDays, "retention from the bootstrap request must apply")
// The row is now readable via GET.
rec = f.do(t, http.MethodGet, "/agent-network/settings", "")
require.Equal(t, http.StatusOK, rec.Code, "GET after bootstrap must succeed")
}
// TestSettingsHandler_PutWithoutClusterOnUnbootstrapped pins that a PUT
// without a cluster cannot conjure a settings row out of nothing — there is
// no cluster to pin — and surfaces as 404 like the GET.
func TestSettingsHandler_PutWithoutClusterOnUnbootstrapped(t *testing.T) {
f := newAgentNetworkHandlerFixture(t)
rec := f.do(t, http.MethodPut, "/agent-network/settings",
`{"enable_log_collection": false, "enable_prompt_collection": false, "redact_pii": false}`)
assert.Equal(t, http.StatusNotFound, rec.Code,
"cluster-less PUT on an unbootstrapped account must 404: got %d body=%s", rec.Code, rec.Body.String())
assert.Contains(t, rec.Body.String(), "cluster",
"the error must point the caller at the bootstrap paths: %s", rec.Body.String())
}
// TestSettingsHandler_PutReplacesMutableFields pins the update contract shared
// with the other PUT endpoints: the request replaces every mutable field, so a
// toggle absent from the JSON lands as its zero value rather than being
// preserved. Cluster and subdomain survive untouched.
func TestSettingsHandler_PutReplacesMutableFields(t *testing.T) {
f := newAgentNetworkHandlerFixture(t)
rec := f.do(t, http.MethodPut, "/agent-network/settings",
`{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": true, "access_log_retention_days": 14}`)
require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String())
var before api.AgentNetworkSettings
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &before))
rec = f.do(t, http.MethodPut, "/agent-network/settings",
`{"enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`)
require.Equal(t, http.StatusOK, rec.Code, "update PUT must succeed: %s", rec.Body.String())
var got api.AgentNetworkSettings
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
assert.True(t, got.EnableLogCollection, "sent toggle must apply")
assert.False(t, got.EnablePromptCollection, "sent toggle must apply")
assert.False(t, got.RedactPii, "sent toggle must apply")
require.NotNil(t, got.AccessLogRetentionDays)
assert.Equal(t, 0, *got.AccessLogRetentionDays,
"retention absent from the request must land as the zero value — PUT replaces all mutable fields")
assert.Equal(t, before.Cluster, got.Cluster, "cluster must survive updates untouched")
assert.Equal(t, before.Subdomain, got.Subdomain, "subdomain must survive updates untouched")
}
// TestSettingsHandler_PutRejectsClusterChange pins cluster immutability: once
// assigned, a differing cluster is rejected as a validation error instead of
// being silently ignored, so callers never observe a value other than the one
// they sent. Echoing the assigned cluster back stays valid, which lets
// declarative clients send their full desired state idempotently.
func TestSettingsHandler_PutRejectsClusterChange(t *testing.T) {
f := newAgentNetworkHandlerFixture(t)
rec := f.do(t, http.MethodPut, "/agent-network/settings",
`{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`)
require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String())
rec = f.do(t, http.MethodPut, "/agent-network/settings",
`{"cluster": "us.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`)
assert.Equal(t, http.StatusUnprocessableEntity, rec.Code,
"cluster change must be rejected as a validation error: got %d body=%s", rec.Code, rec.Body.String())
rec = f.do(t, http.MethodPut, "/agent-network/settings",
`{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": true}`)
require.Equal(t, http.StatusOK, rec.Code, "echoing the assigned cluster must stay valid: %s", rec.Body.String())
var got api.AgentNetworkSettings
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
assert.Equal(t, "eu.proxy.netbird.io", got.Cluster, "cluster must be unchanged")
assert.True(t, got.RedactPii, "toggle sent alongside the echoed cluster must apply")
}

View File

@@ -207,7 +207,7 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide
}
if strings.TrimSpace(bootstrapCluster) != "" {
if _, err := m.bootstrapSettingsIfNeeded(ctx, provider.AccountID, bootstrapCluster); err != nil {
if _, err := m.bootstrapSettingsIfNeeded(ctx, m.store, provider.AccountID, bootstrapCluster); err != nil {
// The provider create has already succeeded; logging the
// bootstrap miss matches the plan's PoC behaviour. The synth
// path treats a missing settings row as a no-op, and the next
@@ -559,40 +559,83 @@ func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, r
return nil
}
// UpdateSettings applies the mutable account-level settings — the collection
// toggles — onto the existing row. Cluster and Subdomain are immutable and are
// preserved from the persisted row regardless of the input. Because the
// collection toggles change the synthesised service config (prompt-capture
// gating, access-log emission), a reconcile is triggered so the proxy and peer
// network maps converge on the new state.
// UpdateSettings replaces the mutable account-level settings — the collection
// toggles and retention — on the account's row. When the account has no
// settings row yet, a non-empty settings.Cluster bootstraps one (same path as
// first provider create); without it the update fails with NotFound. On an
// existing row the cluster and subdomain are immutable: a differing
// settings.Cluster is rejected rather than silently ignored so callers never
// observe a value other than what they sent. Because the collection toggles
// change the synthesised service config (prompt-capture gating, access-log
// emission), a reconcile is triggered so the proxy and peer network maps
// converge on the new state.
func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, settings *types.Settings) (*types.Settings, error) {
if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Update); err != nil {
return nil, err
}
existing, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, settings.AccountID)
requestedCluster := strings.TrimSpace(settings.Cluster)
// The row lock from LockingStrengthUpdate only holds for the duration of
// the surrounding transaction, so the read, the cluster-immutability
// check, and the save must share one — otherwise concurrent PUTs could
// interleave between them.
var updated *types.Settings
err := m.store.ExecuteInTransaction(ctx, func(tx store.Store) error {
existing, err := tx.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, settings.AccountID)
switch {
case err == nil:
if requestedCluster != "" && requestedCluster != existing.Cluster {
return status.Errorf(status.InvalidArgument, "cluster is immutable once assigned (current: %s)", existing.Cluster)
}
case isNotFound(err):
if requestedCluster == "" {
return status.Errorf(status.NotFound, "agent network settings have not been bootstrapped yet; pass cluster to bootstrap them, or create a provider with bootstrap_cluster set")
}
// Bootstrapping pins the cluster and subdomain — a settings
// create on top of the update the caller already passed, matching
// the gate on the provider-create bootstrap path.
if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Create); err != nil {
return err
}
existing, err = m.bootstrapSettingsIfNeeded(ctx, tx, settings.AccountID, requestedCluster)
if err != nil {
return err
}
default:
return fmt.Errorf("get agent network settings: %w", err)
}
existing.EnableLogCollection = settings.EnableLogCollection
existing.EnablePromptCollection = settings.EnablePromptCollection
existing.RedactPii = settings.RedactPii
existing.AccessLogRetentionDays = settings.AccessLogRetentionDays
existing.UpdatedAt = time.Now().UTC()
if err := tx.SaveAgentNetworkSettings(ctx, existing); err != nil {
return fmt.Errorf("save agent network settings: %w", err)
}
updated = existing
return nil
})
if err != nil {
return nil, fmt.Errorf("get agent network settings: %w", err)
}
existing.EnableLogCollection = settings.EnableLogCollection
existing.EnablePromptCollection = settings.EnablePromptCollection
existing.RedactPii = settings.RedactPii
existing.AccessLogRetentionDays = settings.AccessLogRetentionDays
existing.UpdatedAt = time.Now().UTC()
if err := m.store.SaveAgentNetworkSettings(ctx, existing); err != nil {
return nil, fmt.Errorf("save agent network settings: %w", err)
return nil, err
}
m.accountManager.StoreEvent(ctx, userID, settings.AccountID, settings.AccountID, activity.AgentNetworkSettingsUpdated, map[string]any{
"log_collection": existing.EnableLogCollection,
"prompt_collection": existing.EnablePromptCollection,
"redact_pii": existing.RedactPii,
"log_collection": updated.EnableLogCollection,
"prompt_collection": updated.EnablePromptCollection,
"redact_pii": updated.RedactPii,
})
m.reconcile(ctx, settings.AccountID)
return existing, nil
return updated, nil
}
// isNotFound reports whether err is a status.NotFound error.
func isNotFound(err error) bool {
var sErr *status.Error
return errors.As(err, &sErr) && sErr.Type() == status.NotFound
}
// validateProviderRefs ensures every destination provider id refers to a
@@ -616,22 +659,25 @@ func (m *managerImpl) validateProviderRefs(ctx context.Context, accountID string
return nil
}
// GetSettings returns the agent-network settings row for the account.
// Returns the underlying status.NotFound when no row has been
// bootstrapped yet (i.e. the account has no providers).
// GetSettings returns the agent-network settings row for the account. When no
// row has been bootstrapped yet, the defaults are returned (without
// persisting) with cluster and subdomain empty — settings always read as an
// object, like the account and DNS settings endpoints.
func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) (*types.Settings, error) {
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Read); err != nil {
return nil, err
}
return m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
settings, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
switch {
case err == nil:
return settings, nil
case isNotFound(err):
return types.DefaultSettings(accountID), nil
default:
return nil, err
}
}
// bootstrapSettingsIfNeeded creates the per-account agent-network
// settings row when missing. The cluster comes from the create-time
// hint the dashboard sends (auto-picked from the active cluster list);
// the subdomain is picked from the curated wordlist avoiding
// collisions on the same cluster. Idempotent: if a row already exists
// it is returned untouched and the hint is ignored.
// requireSettingsBootstrapPermission gates the one-time settings bootstrap a
// first provider create performs. Pinning the account's cluster and subdomain
// is a settings write, so it needs the settings permission on top of the
@@ -641,14 +687,20 @@ func (m *managerImpl) requireSettingsBootstrapPermission(ctx context.Context, ac
if err == nil {
return nil
}
var sErr *status.Error
if !errors.As(err, &sErr) || sErr.Type() != status.NotFound {
if !isNotFound(err) {
return fmt.Errorf("get agent network settings: %w", err)
}
return m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Create)
}
func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, providerCluster string) (*types.Settings, error) {
// bootstrapSettingsIfNeeded creates the per-account agent-network
// settings row when missing. The cluster comes from the create-time
// hint the dashboard sends (auto-picked from the active cluster list);
// the subdomain is picked from the curated wordlist avoiding
// collisions on the same cluster. Idempotent: if a row already exists
// it is returned untouched and the hint is ignored. st is the store to
// operate on — pass the transaction store when calling from within one.
func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, st store.Store, accountID, providerCluster string) (*types.Settings, error) {
if accountID == "" {
return nil, fmt.Errorf("bootstrap settings: account id is required")
}
@@ -656,16 +708,15 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID,
return nil, fmt.Errorf("bootstrap settings: provider cluster is required")
}
existing, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
existing, err := st.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
if err == nil {
return existing, nil
}
var sErr *status.Error
if !errors.As(err, &sErr) || sErr.Type() != status.NotFound {
if !isNotFound(err) {
return nil, fmt.Errorf("get agent network settings: %w", err)
}
siblings, err := m.store.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster)
siblings, err := st.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster)
if err != nil {
return nil, fmt.Errorf("list agent network settings on cluster: %w", err)
}
@@ -684,18 +735,12 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID,
m.labelRngMu.Unlock()
now := time.Now().UTC()
settings := &types.Settings{
AccountID: accountID,
Cluster: providerCluster,
Subdomain: subdomain,
// Logs on by default; usage is collected regardless. Retention bounds
// how long full log rows are kept.
EnableLogCollection: true,
AccessLogRetentionDays: types.DefaultAccessLogRetentionDays,
CreatedAt: now,
UpdatedAt: now,
}
if err := m.store.SaveAgentNetworkSettings(ctx, settings); err != nil {
settings := types.DefaultSettings(accountID)
settings.Cluster = providerCluster
settings.Subdomain = subdomain
settings.CreatedAt = now
settings.UpdatedAt = now
if err := st.SaveAgentNetworkSettings(ctx, settings); err != nil {
return nil, fmt.Errorf("save agent network settings: %w", err)
}
return settings, nil
@@ -898,8 +943,8 @@ func (*mockManager) UpdateBudgetRule(_ context.Context, _ string, r *types.Accou
func (*mockManager) DeleteBudgetRule(_ context.Context, _, _, _ string) error { return nil }
func (*mockManager) GetSettings(_ context.Context, _, _ string) (*types.Settings, error) {
return nil, status.Errorf(status.NotFound, "agent network settings not found")
func (*mockManager) GetSettings(_ context.Context, accountID, _ string) (*types.Settings, error) {
return types.DefaultSettings(accountID), nil
}
func (*mockManager) UpdateSettings(_ context.Context, _ string, s *types.Settings) (*types.Settings, error) {

View File

@@ -164,9 +164,7 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) {
p.MetadataDisabled = *req.MetadataDisabled
}
// Identity-header overrides for catalogs flagged Customizable.
// nil pointer = "field omitted on the wire" → leave the stored
// value untouched (per the openapi description). Empty string is
// an explicit clear that disables stamping for this dimension.
// Empty or omitted disables stamping for this dimension.
if req.IdentityHeaderUserId != nil {
p.IdentityHeaderUserID = strings.TrimSpace(*req.IdentityHeaderUserId)
}
@@ -192,16 +190,20 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
created := p.CreatedAt
updated := p.UpdatedAt
resp := &api.AgentNetworkProvider{
Id: p.ID,
ProviderId: p.ProviderID,
Name: p.Name,
UpstreamUrl: p.UpstreamURL,
Models: models,
Enabled: p.Enabled,
SkipTlsVerification: p.SkipTLSVerification,
MetadataDisabled: p.MetadataDisabled,
CreatedAt: &created,
UpdatedAt: &updated,
Id: p.ID,
ProviderId: p.ProviderID,
Name: p.Name,
UpstreamUrl: p.UpstreamURL,
Models: models,
// Always present on the wire so an explicitly cleared header
// round-trips as "" instead of vanishing from the response.
IdentityHeaderUserId: p.IdentityHeaderUserID,
IdentityHeaderGroups: p.IdentityHeaderGroups,
Enabled: p.Enabled,
SkipTlsVerification: p.SkipTLSVerification,
MetadataDisabled: p.MetadataDisabled,
CreatedAt: &created,
UpdatedAt: &updated,
}
if len(p.ExtraValues) > 0 {
out := make(map[string]string, len(p.ExtraValues))
@@ -210,14 +212,6 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
}
resp.ExtraValues = &out
}
if p.IdentityHeaderUserID != "" {
v := p.IdentityHeaderUserID
resp.IdentityHeaderUserId = &v
}
if p.IdentityHeaderGroups != "" {
v := p.IdentityHeaderGroups
resp.IdentityHeaderGroups = &v
}
return resp
}

View File

@@ -77,3 +77,41 @@ func TestProvider_MetadataDisabled_RoundTrip(t *testing.T) {
assert.False(t, p.MetadataDisabled, "explicit false must clear metadata_disabled")
assert.False(t, p.ToAPIResponse().MetadataDisabled, "response must reflect the cleared value")
}
// TestProvider_IdentityHeaders_AlwaysOnWire pins that the identity header
// fields are always present in the API response — an explicitly cleared
// ("") header must round-trip as "" rather than vanish, so API consumers
// (e.g. the Terraform provider) never observe a value other than the one
// they wrote.
func TestProvider_IdentityHeaders_AlwaysOnWire(t *testing.T) {
set := "x-bf-dim-netbird_user_id"
empty := ""
base := func() *api.AgentNetworkProviderRequest {
return &api.AgentNetworkProviderRequest{
ProviderId: "custom",
Name: "bifrost",
UpstreamUrl: "https://bifrost.internal",
}
}
p := NewProvider("acc-1")
resp := p.ToAPIResponse()
assert.Equal(t, "", resp.IdentityHeaderUserId, "unset header must surface as empty string, not be omitted")
assert.Equal(t, "", resp.IdentityHeaderGroups, "unset header must surface as empty string, not be omitted")
req := base()
req.IdentityHeaderUserId = &set
p.FromAPIRequest(req)
assert.Equal(t, set, p.ToAPIResponse().IdentityHeaderUserId, "configured header must round-trip")
// Omitting the field preserves it.
p.FromAPIRequest(base())
assert.Equal(t, set, p.ToAPIResponse().IdentityHeaderUserId, "omitted header must preserve the stored value")
// An explicit "" clears it AND stays visible on the wire.
req = base()
req.IdentityHeaderUserId = &empty
p.FromAPIRequest(req)
assert.Equal(t, "", p.ToAPIResponse().IdentityHeaderUserId, "cleared header must round-trip as empty string")
}

View File

@@ -1,6 +1,7 @@
package types
import (
"strings"
"time"
"github.com/netbirdio/netbird/shared/management/http/api"
@@ -42,18 +43,34 @@ type Settings struct {
// schema cohesive.
func (Settings) TableName() string { return "agent_network_settings" }
// DefaultSettings returns the settings an account observes before its row is
// bootstrapped: log collection on with the default retention, everything else
// off, and no cluster/subdomain assigned yet. Bootstrap persists exactly these
// values plus the assigned cluster and subdomain, so the pre-bootstrap read
// and the freshly bootstrapped row agree.
func DefaultSettings(accountID string) *Settings {
return &Settings{
AccountID: accountID,
EnableLogCollection: true,
AccessLogRetentionDays: DefaultAccessLogRetentionDays,
}
}
// Endpoint returns the bare hostname agents reach this account at:
// `<subdomain>.<cluster>`.
// `<subdomain>.<cluster>`. Empty until both halves are assigned at bootstrap.
func (s *Settings) Endpoint() string {
if s.Cluster == "" || s.Subdomain == "" {
return ""
}
return s.Subdomain + "." + s.Cluster
}
// ToAPIResponse renders the settings as the API representation.
// ToAPIResponse renders the settings as the API representation. The
// timestamps are omitted while zero — a default (not yet bootstrapped) view
// has no persisted row to date.
func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings {
created := s.CreatedAt
updated := s.UpdatedAt
retention := s.AccessLogRetentionDays
return &api.AgentNetworkSettings{
resp := &api.AgentNetworkSettings{
Cluster: s.Cluster,
Subdomain: s.Subdomain,
Endpoint: s.Endpoint(),
@@ -61,14 +78,27 @@ func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings {
EnablePromptCollection: s.EnablePromptCollection,
RedactPii: s.RedactPii,
AccessLogRetentionDays: &retention,
CreatedAt: &created,
UpdatedAt: &updated,
}
if !s.CreatedAt.IsZero() {
created := s.CreatedAt
resp.CreatedAt = &created
}
if !s.UpdatedAt.IsZero() {
updated := s.UpdatedAt
resp.UpdatedAt = &updated
}
return resp
}
// FromAPIRequest applies the mutable settings fields from the request. Cluster
// and Subdomain are immutable and intentionally not touched here.
// FromAPIRequest applies the request onto the receiver. The mutable
// collection fields are always replaced with the request values. Cluster
// participates only in bootstrap and the immutability check (see
// Manager.UpdateSettings); Subdomain is server-assigned and never taken
// from a request.
func (s *Settings) FromAPIRequest(req *api.AgentNetworkSettingsRequest) {
if req.Cluster != nil {
s.Cluster = strings.TrimSpace(*req.Cluster)
}
s.EnableLogCollection = req.EnableLogCollection
s.EnablePromptCollection = req.EnablePromptCollection
s.RedactPii = req.RedactPii

View File

@@ -0,0 +1,171 @@
package networkmapdb
import (
"context"
"database/sql"
"encoding/json"
"errors"
"reflect"
"strings"
"github.com/rs/xid"
"golang.org/x/exp/maps"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
NMAP_STRUCT_TAG = "nmap"
NMAP_SKIP = "skip"
NMAP_MAP_TO = "map_to"
)
type NetworkMapDBStore interface { //nolint:revive // established name across the codebase
GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string]map[string]any, error)
GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error)
GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error)
GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error)
GetNameServerGroups(ctx context.Context, accountId string) ([]nmdata.NameServerGroup, error)
GetNetworkResources(ctx context.Context, accountId string) ([]nmdata.NetworkResource, error)
GetNetworkRouters(ctx context.Context, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error)
GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error)
GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error)
GetAccountSettings(ctx context.Context, accountId string) (nmdata.AccountSettingsInfo, error)
GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, map[string]string, error)
GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error)
GetAllowedUsers(ctx context.Context, accountId string) (map[string]struct{}, map[string][]string, error)
GetDnsSettings(ctx context.Context, accountId string) (nmdata.DNSSettings, error)
}
type NetworkMapDBStoreImpl struct { //nolint:revive // established name across the codebase
store NetworkMapDBStore
integratedPeerValidator integrated_validator.IntegratedValidator
extraSettingsManager settings.Manager
}
func NewNetworkMapDBStoreImpl(store NetworkMapDBStore, integratedPeerValidator integrated_validator.IntegratedValidator, extraSettingsManager settings.Manager) *NetworkMapDBStoreImpl {
return &NetworkMapDBStoreImpl{
store: store,
integratedPeerValidator: integratedPeerValidator,
extraSettingsManager: extraSettingsManager,
}
}
func (s *NetworkMapDBStoreImpl) GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error) {
nmdata, err := s.store.GetNetworkMapData(ctx, accountId)
if err != nil {
return nil, err
}
extraSettings, err := s.extraSettingsManager.GetExtraSettings(ctx, accountId)
if err != nil {
return nil, err
}
nmdata.ValidatedPeers, err = s.integratedPeerValidator.GetValidatedPeers(ctx, accountId, maps.Values(nmdata.Groups), maps.Values(nmdata.Peers), extraSettings)
if err != nil {
return nil, err
}
return nmdata, nil
}
func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error {
typ := src.Elem().Type()
for i := 0; i < typ.NumField(); i++ {
f := typ.Field(i)
fieldTags := make(map[string]string)
if v := f.Tag.Get(NMAP_STRUCT_TAG); v != "" {
for _, t := range strings.Split(v, ",") {
kv := tagFromString(t)
fieldTags[kv.Key] = kv.Value
}
}
if _, ok := fieldTags[NMAP_SKIP]; ok {
continue
}
if f.PkgPath != "" { // skip unexported fields
continue
}
dstFieldName := f.Name
if override, ok := fieldTags[NMAP_MAP_TO]; ok {
dstFieldName = override
}
dstField := dst.Elem().FieldByName(dstFieldName)
if !dstField.IsValid() {
return errors.New("unsupported type in destination field: " + dstFieldName)
}
srcField := src.Elem().Field(i)
srcFieldType := srcField.Type().String()
switch srcFieldType {
case "string":
s := srcField.Interface().(string)
dstField.SetString(s)
case "sql.NullString":
s := srcField.Interface().(sql.NullString)
if s.Valid {
dstField.SetString(s.String)
}
if (dstFieldName == "PublicId" || dstFieldName == "PublicID") && s.String == "" {
dstField.SetString(xid.New().String()) // TODO (dmitri) this needs to be removed to support delta updates
}
case "sql.NullTime":
s := srcField.Interface().(sql.NullTime)
if s.Valid {
if dstField.Kind() == reflect.Ptr {
t := reflect.ValueOf(&s.Time).Elem()
dstField.Set(t.Addr())
} else {
dstField.Set(reflect.ValueOf(s.Time))
}
}
case "sql.NullBool":
s := srcField.Interface().(sql.NullBool)
if s.Valid {
dstField.SetBool(s.Bool)
}
case "sql.NullInt64":
s := srcField.Interface().(sql.NullInt64)
if s.Valid {
dstField.SetInt(s.Int64)
}
case "json.RawMessage":
s := srcField.Interface().(json.RawMessage)
if len(s) == 0 {
continue
}
if err := json.Unmarshal(s, dstField.Addr().Interface()); err != nil {
return err
}
case "[]string":
if srcField.IsNil() {
continue
}
dstv := reflect.MakeSlice(dstField.Type(), srcField.Len(), srcField.Cap())
reflect.Copy(dstv, srcField)
dstField.Set(dstv)
}
}
return nil
}
type fieldTag struct {
Key string
Value string
}
func tagFromString(t string) fieldTag {
kv := strings.Split(t, ":")
if len(kv) == 1 {
return fieldTag{Key: strings.TrimSpace(kv[0])}
}
return fieldTag{Key: strings.TrimSpace(kv[0]), Value: strings.TrimSpace(kv[1])}
}

View File

@@ -0,0 +1,84 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"time"
"github.com/jackc/pgx/v5"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetAccountSettingsQuery = `
select settings_peer_login_expiration_enabled as peer_login_expiration_enabled,
settings_peer_login_expiration as peer_login_expiration,
settings_peer_inactivity_expiration_enabled as peer_inactivity_expiration_enabled,
settings_peer_inactivity_expiration as peer_inactivity_expiration,
settings_dns_domain as dns_domain,
settings_ipv6_enabled_groups as ipv6_enabled_groups,
settings_routing_peer_dns_resolution_enabled as routing_peer_dns_resolution_enabled,
settings_lazy_connection_enabled as lazy_connection_enabled,
settings_auto_update_version as auto_update_version,
settings_auto_update_always as auto_update_always,
settings_metrics_push_enabled as metrics_push_enabled
from accounts
where id=$1
`
)
func (pg *PgStore) GetAccountSettings(ctx context.Context, accountId string) (nmdata.AccountSettingsInfo, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nmdata.AccountSettingsInfo{}, err
}
return GetAccountSettingsViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetAccountSettingsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.AccountSettingsInfo, error) {
rows, err := con.Query(ctx, GetAccountSettingsQuery, accountId)
if err != nil {
return nmdata.AccountSettingsInfo{}, err
}
settings, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[account])
if err != nil {
return nmdata.AccountSettingsInfo{}, err
}
settingsInfo := nmdata.AccountSettingsInfo{
PeerLoginExpirationEnabled: settings.PeerLoginExpirationEnabled.Bool,
PeerLoginExpiration: time.Duration(settings.PeerLoginExpiration.Int64),
PeerInactivityExpirationEnabled: settings.PeerInactivityExpirationEnabled.Bool,
PeerInactivityExpiration: time.Duration(settings.PeerInactivityExpiration.Int64),
DNSDomain: settings.DNSDomain.String,
RoutingPeerDNSResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled.Bool,
LazyConnectionEnabled: settings.LazyConnectionEnabled.Bool,
AutoUpdateVersion: settings.AutoUpdateVersion.String,
AutoUpdateAlways: settings.AutoUpdateAlways.Bool,
MetricsPushEnabled: settings.MetricsPushEnabled.Bool,
}
if settings.IPv6EnabledGroups != nil {
if err := json.Unmarshal(settings.IPv6EnabledGroups, &settingsInfo.IPv6EnabledGroups); err != nil {
return nmdata.AccountSettingsInfo{}, err
}
}
return settingsInfo, nil
}
type account struct {
PeerLoginExpirationEnabled sql.NullBool
PeerLoginExpiration sql.NullInt64
PeerInactivityExpirationEnabled sql.NullBool
PeerInactivityExpiration sql.NullInt64
DNSDomain sql.NullString
IPv6EnabledGroups json.RawMessage
RoutingPeerDNSResolutionEnabled sql.NullBool
LazyConnectionEnabled sql.NullBool
AutoUpdateVersion sql.NullString
AutoUpdateAlways sql.NullBool
MetricsPushEnabled sql.NullBool
}

View File

@@ -0,0 +1,124 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"reflect"
"github.com/jackc/pgx/v5"
"github.com/miekg/dns"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
var ErrDnsUnsupportedRecordType = errors.New("unsupported record type")
const (
GetAccountZonesQuery = `
select zones.id as id, domain, not enable_search_domain as search_domain_disabled, distribution_groups,
r.name as record_name, r.type as record_type, 'IN' record_class, r.ttl as record_ttl, r.content as record_rdata
from zones
left join records as r on r.zone_id = zones.id
where zones.account_id=$1
`
)
func (pg *PgStore) GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, err
}
return GetAppliedZoneCandidatesViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetAppliedZoneCandidatesViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) ([]networkmap.AppliedZoneCandidate, error) {
rows, err := conn.Query(ctx, GetAccountZonesQuery, accountId)
if err != nil {
return nil, err
}
zones, err := pgx.CollectRows(rows, pgx.RowToStructByName[zone])
if err != nil {
return nil, err
}
toret := make([]networkmap.AppliedZoneCandidate, 0, len(zones))
currentZoneId := ""
for _, z := range zones {
if !z.RecordType.Valid {
continue
}
zone := nmdata.CustomZone{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&z), reflect.ValueOf(&zone))
if err != nil {
return nil, err
}
var distributionGroups []string
if err := json.Unmarshal(z.DistributionGroups, &distributionGroups); err != nil {
return nil, err
}
if z.Id != currentZoneId {
zone.Records = []nmdata.SimpleRecord{}
toret = append(toret, appliedZoneCandidateFromZone(zone, distributionGroups))
currentZoneId = z.Id
}
rtype, rdata, err := recordTypeAndRdata(z.RecordType.String, z.RecordRData.String)
if err != nil {
if errors.Is(err, ErrDnsUnsupportedRecordType) {
continue
}
return nil, err
}
lastZone := &toret[len(toret)-1]
lastZone.Zone.Records = append(lastZone.Zone.Records, nmdata.SimpleRecord{
Name: z.RecordName.String,
Class: z.RecordClass.String,
TTL: int(z.RecordTTL.Int64),
RData: rdata,
Type: rtype,
})
}
return toret, nil
}
type zone struct {
Id string `nmap:"skip"`
DistributionGroups json.RawMessage `nmap:"skip"`
Domain sql.NullString
SearchDomainDisabled sql.NullBool
RecordName sql.NullString `nmap:"skip"`
RecordType sql.NullString `nmap:"skip"`
RecordClass sql.NullString `nmap:"skip"`
RecordTTL sql.NullInt64 `nmap:"skip"`
RecordRData sql.NullString `nmap:"skip"`
}
func recordTypeAndRdata(t, rdata string) (int, string, error) {
switch t {
case "A":
return int(dns.TypeA), rdata, nil
case "AAAA":
return int(dns.TypeAAAA), rdata, nil
case "CNAME":
return int(dns.TypeCNAME), dns.Fqdn(rdata), nil
default:
return 0, "", fmt.Errorf("record type: %s %w", t, ErrDnsUnsupportedRecordType)
}
}
func appliedZoneCandidateFromZone(z nmdata.CustomZone, distributionGroups []string) networkmap.AppliedZoneCandidate {
return networkmap.AppliedZoneCandidate{
DistributionGroups: distributionGroups,
Zone: z,
}
}

View File

@@ -0,0 +1,53 @@
package networkmap_pgsql
import (
"context"
"encoding/json"
"github.com/jackc/pgx/v5"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetDnsSettingsQuery = `
select dns_settings_disabled_management_groups
from accounts
where id=$1
`
)
func (pg *PgStore) GetDnsSettings(ctx context.Context, accountId string) (nmdata.DNSSettings, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nmdata.DNSSettings{}, err
}
return GetDnsSettingsViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetDnsSettingsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.DNSSettings, error) {
rows, err := con.Query(ctx, GetDnsSettingsQuery, accountId)
if err != nil {
return nmdata.DNSSettings{}, err
}
return pgx.CollectOneRow(rows, rowToDnsSettings)
}
func rowToDnsSettings(row pgx.CollectableRow) (nmdata.DNSSettings, error) {
var value nmdata.DNSSettings
var settings json.RawMessage
if err := row.Scan(&settings); err != nil {
return value, err
}
if settings == nil {
return nmdata.DNSSettings{}, nil
}
if err := json.Unmarshal(settings, &value.DisabledManagementGroups); err != nil {
return value, err
}
return value, nil
}

View File

@@ -0,0 +1,38 @@
package networkmap_pgsql
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestRecordTypeAndRdata(t *testing.T) {
var tests = []struct {
recordType string
expectedRecordType int
rdata string
expectedRdata string
expectedErr error
}{
{recordType: "A", expectedRecordType: 1, rdata: "test.com", expectedRdata: "test.com", expectedErr: nil},
{recordType: "AAAA", expectedRecordType: 28, rdata: "test.com", expectedRdata: "test.com", expectedErr: nil},
{recordType: "CNAME", expectedRecordType: 5, rdata: "test.com", expectedRdata: "test.com.", expectedErr: nil},
{recordType: "CNAME", expectedRecordType: 5, rdata: "test.com.", expectedRdata: "test.com.", expectedErr: nil},
{recordType: "TypeMX", expectedErr: ErrDnsUnsupportedRecordType},
}
for _, tt := range tests {
t.Run(tt.recordType, func(t *testing.T) {
recordType, rdata, err := recordTypeAndRdata(tt.recordType, tt.rdata)
if tt.expectedErr != nil {
assert.ErrorIs(t, err, ErrDnsUnsupportedRecordType)
return
}
assert.NoError(t, err)
assert.Equal(t, recordType, tt.expectedRecordType)
assert.Equal(t, rdata, tt.expectedRdata)
})
}
}

View File

@@ -0,0 +1,38 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"github.com/jackc/pgx/v5"
)
const (
GetDomainsQuery = `
select domain, target_cluster
from domains
where account_id=$1 and domain<>'' and target_cluster<>''
`
)
func (pg *PgStore) GetDomains(ctx context.Context, accountId string) ([]Domain, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, err
}
return GetDomainsViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetDomainsViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) ([]Domain, error) {
rows, err := conn.Query(ctx, GetDomainsQuery, accountId)
if err != nil {
return nil, err
}
return pgx.CollectRows(rows, pgx.RowToStructByName[Domain])
}
type Domain struct {
Domain sql.NullString
TargetCluster sql.NullString
}

View File

@@ -0,0 +1,72 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetGroupsQuery = `
select id, name, public_id, resources,
(
select array_agg(group_peers.peer_id)
from group_peers
where group_peers.group_id = groups.id and group_peers.account_id=$1
) as peers
from groups where account_id=$1
`
)
// we also return a resource-to-group index.
// an alternative is to add json indexes, query this directly. Not sure how expensive
// json indexes are. TODO (dmitri) verify and maybe change the implementation here.
func (pg *PgStore) GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string]map[string]any, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, nil, err
}
return GetGroupsViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetGroupsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Group, map[string]map[string]any, error) {
rows, err := con.Query(ctx, GetGroupsQuery, accountId)
if err != nil {
return nil, nil, err
}
groups, err := pgx.CollectRows(rows, pgx.RowToStructByName[group])
toret := make([]nmdata.Group, 0, len(groups))
resourceToGroupIdx := make(map[string]map[string]any)
for _, g := range groups {
dg := nmdata.Group{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&g), reflect.ValueOf(&dg))
if err != nil {
return nil, nil, err
}
toret = append(toret, dg)
for _, resource := range dg.Resources {
if _, ok := resourceToGroupIdx[resource.ID]; !ok {
resourceToGroupIdx[resource.ID] = make(map[string]any)
}
resourceToGroupIdx[resource.ID][g.ID] = struct{}{}
}
}
return toret, resourceToGroupIdx, err
}
type group struct {
ID string
Name sql.NullString
PublicID sql.NullString
Resources json.RawMessage
Peers []string
}

View File

@@ -0,0 +1,65 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetNameserversQuery = `
select id, public_id, name, description, name_servers, groups, "primary", domains, enabled, search_domains_enabled
from name_server_groups
where account_id=$1
`
)
func (pg *PgStore) GetNameServerGroups(ctx context.Context, accountId string) ([]nmdata.NameServerGroup, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, err
}
return GetNameServerGroupsViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetNameServerGroupsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.NameServerGroup, error) {
rows, err := con.Query(ctx, GetNameserversQuery, accountId)
if err != nil {
return nil, err
}
nsgroups, err := pgx.CollectRows(rows, pgx.RowToStructByName[nameserverGroup])
if err != nil {
return nil, err
}
toret := make([]nmdata.NameServerGroup, 0, len(nsgroups))
for _, nsg := range nsgroups {
group := nmdata.NameServerGroup{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&nsg), reflect.ValueOf(&group))
if err != nil {
return nil, err
}
toret = append(toret, group)
}
return toret, nil
}
type nameserverGroup struct {
ID string
PublicID sql.NullString
Name sql.NullString
Description sql.NullString
NameServers json.RawMessage
Groups json.RawMessage
Primary sql.NullBool
Domains json.RawMessage
Enabled sql.NullBool
SearchDomainsEnabled sql.NullBool
}

View File

@@ -0,0 +1,57 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetNetworkQuery = `
select network_identifier as identifier, network_net as net, network_net_v6 as net_v6, network_dns as dns, network_serial as serial
from accounts
where id=$1
`
)
func (pg *PgStore) GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nmdata.Network{}, err
}
return GetNetworkViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetNetworkViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.Network, error) {
rows, err := con.Query(ctx, GetNetworkQuery, accountId)
if err != nil {
return nmdata.Network{}, err
}
n, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[accountnetwork])
if err != nil {
return nmdata.Network{}, err
}
toret := nmdata.Network{}
err = networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&n), reflect.ValueOf(&toret))
if err != nil {
return nmdata.Network{}, err
}
return toret, nil
}
type accountnetwork struct {
Identifier sql.NullString
Net json.RawMessage
NetV6 json.RawMessage
Dns sql.NullString
Serial sql.NullInt64
}

View File

@@ -0,0 +1,242 @@
package networkmap_pgsql
import (
"context"
"fmt"
"strings"
"github.com/jackc/pgx/v5"
"github.com/miekg/dns"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
func (pg *PgStore) GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error) {
tx, err := pg.Pool.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly})
if err != nil {
return nil, err
}
acctSettings, err := GetAccountSettingsViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get account settings: %w", err))
}
dnsZones, err := GetAppliedZoneCandidatesViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get applied zone candidates: %w", err))
}
groups, resourceToGroupIdx, err := GetGroupsViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get groups: %w", err))
}
nsGroups, err := GetNameServerGroupsViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get nameserver groups: %w", err))
}
networkResources, err := GetNetworkResourcesViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network resources: %w", err))
}
routers, err := GetNetworkRoutersViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network routers: %w", err))
}
network, err := GetNetworkViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network: %w", err))
}
peers, proxyPeers, err := GetPeersViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get peers: %w", err))
}
policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := GetPoliciesViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get policies: %w", err))
}
postureChecks, postureCheckXIDToPublicID, err := GetPostureChecksViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get posture checks: %w", err))
}
routes, err := GetRoutesViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get routes: %w", err))
}
networkXIDToPublicID, err := GetNetworkXIDToPublicIdMapViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network xid to public id map: %w", err))
}
allowedUserIds, groupsToUserIds, err := GetAllowedUsersViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get allowed users: %w", err))
}
dnsSettings, err := GetDnsSettingsViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get dns settings: %w", err))
}
domains, err := GetDomainsViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, err)
}
services, err := GetPrivateServicesViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, err)
}
proxyTargetedDomainResourceIDs, err := GetProxyTargetedDomainResourceIDsViaPgxConnection(ctx, tx.Conn(), accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get proxy targeted domain resources: %w", err))
}
resourcePolicies := make(map[string][]*nmdata.Policy)
for _, resource := range networkResources {
if !resource.Enabled {
continue
}
networkResourceGroups := resourceToGroupIdx[resource.ID]
for _, policy := range policies {
if !policy.Enabled {
continue
}
if _, ok := policyToDestinationResourceIdx[policy.ID][resource.ID]; ok {
resourcePolicies[resource.ID] = append(resourcePolicies[resource.ID], &policy) // TODO (dmitri) maybe use public id?
continue
}
if groupIds, ok := policyToDestinationGroupIdx[policy.ID]; ok {
for networkResourceGroup := range networkResourceGroups {
if _, ok := groupIds[networkResourceGroup]; ok {
resourcePolicies[resource.ID] = append(resourcePolicies[resource.ID], &policy)
break
}
}
}
}
}
if err = tx.Commit(ctx); err != nil {
log.WithContext(ctx).Warnf("failed to commit network map read transaction: %v", err)
}
toret := networkmap.NetworkMapData{
AccountSettings: &acctSettings,
DNSSettings: &dnsSettings,
Network: &network,
Peers: toMap(peers, func(p nmdata.Peer) string { return p.ID }),
Groups: toMap(groups, func(g nmdata.Group) string { return g.ID }),
Policies: toSliceOfPtrs(policies),
ResourcePolicies: resourcePolicies,
Routes: toSliceOfPtrs(routes),
Routers: routers,
NameServerGroups: toSliceOfPtrs(nsGroups),
NetworkResources: toSliceOfPtrs(networkResources),
PostureChecks: toMap(postureChecks, func(pc nmdata.PostureChecks) string { return pc.ID }),
AllowedUserIDs: allowedUserIds,
GroupIDToUserIDs: groupsToUserIds,
NetworkXIDToPublicID: networkXIDToPublicID, // TODO (dmitri) maybe we can switch to public ids everywhere?
AppliedZoneCandidates: dnsZones,
PrivateServiceCandidates: buildPrivateServiceCandidates(services, domains, proxyPeers),
PostureCheckXIDToPublicID: postureCheckXIDToPublicID,
ProxyTargetedDomainResourceIDs: proxyTargetedDomainResourceIDs,
}
return &toret, nil
}
func rollbackAndReturnError(ctx context.Context, tx pgx.Tx, err error) (*networkmap.NetworkMapData, error) {
if errr := tx.Rollback(ctx); errr != nil {
log.WithContext(ctx).Warnf("failed to rollback network map read transaction: %v", errr)
}
return nil, err
}
func toMap[T any](all []T, id func(t T) string) map[string]*T {
toret := make(map[string]*T, len(all))
for _, t := range all {
toret[id(t)] = &t
}
return toret
}
func toSliceOfPtrs[T any](all []T) []*T {
toret := make([]*T, 0, len(all))
for _, t := range all {
toret = append(toret, &t)
}
return toret
}
func serviceDomainZone(svc Service, ds []Domain) string {
if domainFromSuffix(svc.Domain.String, svc.ProxyCluster.String) {
return svc.ProxyCluster.String
}
var zoneName string
for _, domain := range ds {
if domain.TargetCluster.String != svc.ProxyCluster.String {
continue
}
if domainFromSuffix(svc.Domain.String, domain.Domain.String) && len(domain.Domain.String) > len(zoneName) {
zoneName = domain.Domain.String
}
}
return zoneName
}
func domainFromSuffix(domain, suffix string) bool {
if suffix == "" {
return false
}
return domain == suffix || strings.HasSuffix(domain, "."+suffix)
}
func buildPrivateServiceCandidates(svcs []Service, domains []Domain, proxyPeersByCluster map[string][]*nmdata.Peer) []networkmap.PrivateServiceCandidate {
var out []networkmap.PrivateServiceCandidate
if len(proxyPeersByCluster) == 0 {
return out
}
for _, svc := range svcs {
if !svc.Enabled.Bool || !svc.Private.Bool {
continue
}
if len(svc.AccessGroups) == 0 {
continue
}
domainZone := serviceDomainZone(svc, domains)
if domainZone == "" {
continue
}
var records []nmdata.SimpleRecord
for _, proxyPeer := range proxyPeersByCluster[svc.ProxyCluster.String] {
if !proxyPeer.IP.IsValid() {
continue
}
records = append(records, nmdata.SimpleRecord{
Name: dns.Fqdn(svc.Domain.String),
Type: int(dns.TypeA),
Class: "IN",
TTL: 5,
RData: proxyPeer.IP.String(),
})
}
if len(records) == 0 {
continue
}
out = append(out, networkmap.PrivateServiceCandidate{
AccessGroups: svc.AccessGroups,
Zone: nmdata.CustomZone{
Domain: dns.Fqdn(domainZone),
Records: records,
NonAuthoritative: true,
SearchDomainDisabled: true,
},
})
}
return out
}

View File

@@ -0,0 +1,65 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetNetworkResourcesQuery = `
select id, network_id, account_id, public_id, name, description, type, domain, prefix, enabled
from network_resources
where account_id=$1
`
)
func (pg *PgStore) GetNetworkResources(ctx context.Context, accountId string) ([]nmdata.NetworkResource, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, err
}
return GetNetworkResourcesViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetNetworkResourcesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.NetworkResource, error) {
rows, err := con.Query(ctx, GetNetworkResourcesQuery, accountId)
if err != nil {
return nil, err
}
netresorces, err := pgx.CollectRows(rows, pgx.RowToStructByName[networkresource])
if err != nil {
return nil, err
}
toret := make([]nmdata.NetworkResource, 0, len(netresorces))
for _, nres := range netresorces {
resource := nmdata.NetworkResource{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&nres), reflect.ValueOf(&resource))
if err != nil {
return nil, err
}
toret = append(toret, resource)
}
return toret, nil
}
type networkresource struct {
ID string
NetworkID sql.NullString
AccountID sql.NullString
PublicID sql.NullString
Name sql.NullString
Description sql.NullString
Type sql.NullString
Domain sql.NullString
Prefix json.RawMessage
Enabled sql.NullBool
}

View File

@@ -0,0 +1,88 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetNetworkRouterQuery = `
select public_id, peer, network_id, masquerade, metric, enabled, peer_groups,
(
select array_agg(group_peers.peer_id)
from group_peers
where group_peers.account_id=$1 and group_peers.group_id in (select json_array_elements_text(peer_groups::json))
) as peers_via_groups
from network_routers
where account_id=$1
`
)
func (pg *PgStore) GetNetworkRouters(ctx context.Context, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, err
}
return GetNetworkRoutersViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetNetworkRoutersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error) {
rows, err := con.Query(ctx, GetNetworkRouterQuery, accountId)
if err != nil {
return nil, err
}
routers, err := pgx.CollectRows(rows, pgx.RowToStructByName[networkrouter])
if err != nil {
return nil, err
}
toret := make(map[string]map[string]*nmdata.NetworkRouter)
for _, router := range routers {
if !router.Enabled.Bool {
continue
}
networkId := router.NetworkID.String
if networkId == "" {
return nil, fmt.Errorf("router with public_id %s doesn't have network_id set", router.PublicID.String)
}
nmdatarouter := nmdata.NetworkRouter{}
err := networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&router), reflect.ValueOf(&nmdatarouter))
if err != nil {
return nil, err
}
if toret[networkId] == nil {
toret[networkId] = make(map[string]*nmdata.NetworkRouter)
}
if router.Peer.String != "" {
toret[networkId][router.Peer.String] = &nmdatarouter
continue
}
for _, peerId := range router.PeersViaGroups {
toret[networkId][peerId] = &nmdatarouter
}
}
return toret, nil
}
type networkrouter struct {
PublicID sql.NullString
NetworkID sql.NullString `nmap:"skip"`
Peer sql.NullString `nmap:"skip"`
PeerGroups json.RawMessage
PeersViaGroups []string `nmap:"skip"`
Masquerade sql.NullBool
Metric sql.NullInt64
Enabled sql.NullBool
}

View File

@@ -0,0 +1,52 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"github.com/jackc/pgx/v5"
)
const (
GetNetworksQuery = `
select id, public_id
from networks where account_id=$1
`
)
func (pg *PgStore) GetNetworks(ctx context.Context, accountId string) ([]network, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, err
}
return GetNetworksViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetNetworksViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]network, error) {
rows, err := con.Query(ctx, GetNetworksQuery, accountId)
if err != nil {
return nil, err
}
return pgx.CollectRows(rows, pgx.RowToStructByName[network])
}
func GetNetworkXIDToPublicIdMapViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]string, error) {
networks, err := GetNetworksViaPgxConnection(ctx, con, accountId)
if err != nil {
return nil, err
}
toret := make(map[string]string)
for _, n := range networks {
if n.PublicID.Valid {
toret[n.ID] = n.PublicID.String
}
}
return toret, nil
}
type network struct {
ID string
PublicID sql.NullString
}

View File

@@ -0,0 +1,148 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetPeersQuery = `
select id, key, ssh_key, dns_label, extra_dns_labels, user_id, ssh_enabled, login_expiration_enabled, last_login, ip, ipv6,
peer_status_requires_approval, peer_status_connected, proxy_meta_embedded, proxy_meta_cluster,
meta_wt_version, meta_go_os, meta_os_version, meta_kernel_version, meta_network_addresses, meta_files, meta_capabilities, meta_flags, meta_sync_message_version,
location_country_code, location_city_name, location_connection_ip
from peers
where account_id = $1
`
)
func (pg *PgStore) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, nil, err
}
return GetPeersViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetPeersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error) {
rows, err := con.Query(ctx, GetPeersQuery, accountId)
if err != nil {
return nil, nil, err
}
peers, err := pgx.CollectRows(rows, pgx.RowToStructByName[peer])
if err != nil {
return nil, nil, err
}
toret := make([]nmdata.Peer, 0, len(peers))
clusterToPeerIdx := make(map[string][]*nmdata.Peer)
for _, p := range peers {
dp := nmdata.Peer{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&p), reflect.ValueOf(&dp))
if err != nil {
return nil, nil, err
}
if p.ProxyMetaEmbedded.Valid {
dp.ProxyMeta.Embedded = p.ProxyMetaEmbedded.Bool
}
// This is only used to build private service candidates, not connected peers are skipped
if dp.ProxyMeta.Embedded && p.PeerStatusConnected.Bool {
clusterToPeerIdx[p.ProxyMetaCluster.String] = append(clusterToPeerIdx[p.ProxyMetaCluster.String], &dp)
}
if p.MetaWtVersion.Valid {
dp.Meta.WtVersion = p.MetaWtVersion.String
}
if p.MetaSyncMessageVersion.Valid {
dp.Meta.SyncMessageVersion = int(p.MetaSyncMessageVersion.Int64)
}
if p.MetaGoOS.Valid {
dp.Meta.GoOS = p.MetaGoOS.String
}
if p.MetaOSVersion.Valid {
dp.Meta.OSVersion = p.MetaOSVersion.String
}
if p.MetaKernelVersion.Valid {
dp.Meta.KernelVersion = p.MetaKernelVersion.String
}
if p.LocationCountryCode.Valid {
dp.Location.CountryCode = p.LocationCountryCode.String
}
if p.LocationCityName.Valid {
dp.Location.CityName = p.LocationCityName.String
}
if p.LocationConnectionIp != nil {
err := json.Unmarshal(p.LocationConnectionIp, &dp.Location.ConnectionIP)
if err != nil {
return toret, nil, err
}
}
if p.MetaFiles != nil {
err := json.Unmarshal(p.MetaFiles, &dp.Meta.Files)
if err != nil {
return toret, nil, err
}
}
if p.MetaCapabilities != nil {
err := json.Unmarshal(p.MetaCapabilities, &dp.Meta.Capabilities)
if err != nil {
return toret, nil, err
}
}
if p.MetaFlags != nil {
err := json.Unmarshal(p.MetaFlags, &dp.Meta.Flags)
if err != nil {
return toret, nil, err
}
}
if p.MetaNetworkAddresses != nil {
err := json.Unmarshal(p.MetaNetworkAddresses, &dp.Meta.NetworkAddresses)
if err != nil {
return toret, nil, err
}
}
toret = append(toret, dp)
}
return toret, clusterToPeerIdx, nil
}
// TODO add support for creating struct fields from denormalized fields
type peer struct {
ID string
Key sql.NullString
SSHKey sql.NullString
DNSLabel sql.NullString
ExtraDNSLabels json.RawMessage
UserID sql.NullString
LastLogin sql.NullTime
SSHEnabled sql.NullBool
LoginExpirationEnabled sql.NullBool
PeerStatusConnected sql.NullBool `nmap:"skip"`
PeerStatusRequiresApproval sql.NullBool `nmap:"map_to:RequiresApproval"`
ProxyMetaEmbedded sql.NullBool `nmap:"skip"`
ProxyMetaCluster sql.NullString `nmap:"skip"`
IP json.RawMessage
IPv6 json.RawMessage
LocationConnectionIp json.RawMessage `nmap:"skip"`
MetaFiles json.RawMessage `nmap:"skip"`
MetaCapabilities json.RawMessage `nmap:"skip"`
MetaFlags json.RawMessage `nmap:"skip"`
MetaNetworkAddresses json.RawMessage `nmap:"skip"`
MetaWtVersion sql.NullString `nmap:"skip"`
MetaGoOS sql.NullString `nmap:"skip"`
MetaOSVersion sql.NullString `nmap:"skip"`
MetaKernelVersion sql.NullString `nmap:"skip"`
MetaSyncMessageVersion sql.NullInt64 `nmap:"skip"`
LocationCountryCode sql.NullString `nmap:"skip"`
LocationCityName sql.NullString `nmap:"skip"`
}

View File

@@ -0,0 +1,56 @@
package networkmap_pgsql
import (
"context"
"fmt"
"time"
"github.com/jackc/pgx/v5/pgxpool"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
)
const (
pgMaxConnections = 30
pgMinConnections = 1
pgMaxConnLifetime = 60 * time.Minute
pgHealthCheckPeriod = 1 * time.Minute
)
var _ networkmapdb.NetworkMapDBStore = &PgStore{}
type PgStore struct {
Pool *pgxpool.Pool
}
func NewPostgresqlStore(ctx context.Context, dsn string) (*PgStore, error) {
pool, err := connectToPgDb(ctx, dsn)
if err != nil {
return nil, err
}
return &PgStore{Pool: pool}, nil
}
func connectToPgDb(ctx context.Context, dsn string) (*pgxpool.Pool, error) {
config, err := pgxpool.ParseConfig(dsn)
if err != nil {
return nil, fmt.Errorf("unable to parse database config: %w", err)
}
config.MaxConns = pgMaxConnections
config.MinConns = pgMinConnections
config.MaxConnLifetime = pgMaxConnLifetime
config.HealthCheckPeriod = pgHealthCheckPeriod
pool, err := pgxpool.NewWithConfig(ctx, config)
if err != nil {
return nil, fmt.Errorf("unable to create connection pool: %w", err)
}
if err := pool.Ping(ctx); err != nil {
pool.Close()
return nil, fmt.Errorf("unable to ping database: %w", err)
}
return pool, nil
}

View File

@@ -0,0 +1,168 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetPoliciesQuery = `
select p.id, p.public_id, p.enabled, array (select json_array_elements_text(p.source_posture_checks::json)) as source_posture_checks, pr.enabled as rule_enabled, pr.action, pr.protocol, pr.bidirectional,
pr.sources, pr.destinations, pr.source_resource, pr.destination_resource, pr.ports, pr.port_ranges,
pr.authorized_groups, pr.authorized_user
from policies as p
left join policy_rules as pr on p.id = pr.policy_id
where account_id=$1
`
)
func (pg *PgStore) GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, nil, nil, err
}
return GetPoliciesViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetPoliciesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) {
rows, err := con.Query(ctx, GetPoliciesQuery, accountId)
if err != nil {
return nil, nil, nil, err
}
policies, err := pgx.CollectRows(rows, pgx.RowToStructByName[policy])
if err != nil {
return nil, nil, nil, err
}
toret := make([]nmdata.Policy, 0, len(policies))
policyToDestinationResourceIdx := make(map[string]map[string]any) // policy id to destination resource id
policyToDestinationGroupIdx := make(map[string]map[string]any) // policy id to destination group id
for _, p := range policies {
policy := nmdata.Policy{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&p), reflect.ValueOf(&policy))
if err != nil {
return nil, nil, nil, err
}
var policyRule *nmdata.PolicyRule
pr := func() *nmdata.PolicyRule {
if policyRule != nil {
return policyRule
}
policyRule = &nmdata.PolicyRule{}
return policyRule
}
if p.RuleEnabled.Valid {
pr().Enabled = p.RuleEnabled.Bool
}
if p.Action.Valid {
pr().Action = p.Action.String
}
if p.Protocol.Valid {
pr().Protocol = p.Protocol.String
}
if p.Bidirectional.Valid {
pr().Bidirectional = p.Bidirectional.Bool
}
if len(p.Sources) > 0 {
err := json.Unmarshal([]byte(p.Sources), &pr().Sources)
if err != nil {
return toret, nil, nil, err
}
}
if len(p.Destinations) > 0 {
err := json.Unmarshal([]byte(p.Destinations), &pr().Destinations)
if err != nil {
return toret, nil, nil, err
}
if p.RuleEnabled.Valid && p.RuleEnabled.Bool {
for _, dst := range pr().Destinations {
if _, ok := policyToDestinationGroupIdx[p.ID]; !ok {
policyToDestinationGroupIdx[p.ID] = make(map[string]any)
}
policyToDestinationGroupIdx[p.ID][dst] = struct{}{}
}
}
}
if len(p.SourceResource) > 0 {
err := json.Unmarshal([]byte(p.SourceResource), &pr().SourceResource)
if err != nil {
return toret, nil, nil, err
}
}
if len(p.DestinationResource) > 0 {
err := json.Unmarshal([]byte(p.DestinationResource), &pr().DestinationResource)
if err != nil {
return toret, nil, nil, err
}
if p.RuleEnabled.Valid && p.RuleEnabled.Bool {
if _, ok := policyToDestinationResourceIdx[p.ID]; !ok {
policyToDestinationResourceIdx[p.ID] = make(map[string]any)
}
policyToDestinationResourceIdx[p.ID][pr().DestinationResource.ID] = struct{}{}
}
}
if len(p.Ports) > 0 {
err := json.Unmarshal([]byte(p.Ports), &pr().Ports)
if err != nil {
return toret, nil, nil, err
}
}
if len(p.PortRanges) > 0 {
err := json.Unmarshal([]byte(p.PortRanges), &pr().PortRanges)
if err != nil {
return toret, nil, nil, err
}
}
if len(p.AuthorizedGroups) > 0 {
err := json.Unmarshal([]byte(p.AuthorizedGroups), &pr().AuthorizedGroups)
if err != nil {
return toret, nil, nil, err
}
}
if p.AuthorizedUser.Valid {
pr().AuthorizedUser = p.AuthorizedUser.String
}
if policyRule != nil {
policyRule.ID = p.ID
policyRule.PolicyID = p.ID
policy.Rules = []*nmdata.PolicyRule{policyRule}
}
toret = append(toret, policy)
}
return toret, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err
}
type policy struct {
ID string
PublicID sql.NullString
SourcePostureChecks []string
Enabled sql.NullBool
RuleEnabled sql.NullBool `nmap:"skip"`
Bidirectional sql.NullBool `nmap:"skip"`
Action sql.NullString `nmap:"skip"`
Protocol sql.NullString `nmap:"skip"`
Sources json.RawMessage `nmap:"skip"`
Destinations json.RawMessage `nmap:"skip"`
SourceResource json.RawMessage `nmap:"skip"`
DestinationResource json.RawMessage `nmap:"skip"`
Ports json.RawMessage `nmap:"skip"`
PortRanges json.RawMessage `nmap:"skip"`
AuthorizedGroups json.RawMessage `nmap:"skip"`
AuthorizedUser sql.NullString `nmap:"skip"`
}

View File

@@ -0,0 +1,60 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetPostureChecksQuery = `
select id, public_id, checks
from posture_checks
where account_id=$1
`
)
func (pg *PgStore) GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, map[string]string, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, nil, err
}
return GetPostureChecksViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetPostureChecksViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.PostureChecks, map[string]string, error) {
rows, err := con.Query(ctx, GetPostureChecksQuery, accountId)
if err != nil {
return nil, nil, err
}
checks, err := pgx.CollectRows(rows, pgx.RowToStructByName[posturechecks])
if err != nil {
return nil, nil, err
}
toret := make([]nmdata.PostureChecks, 0, len(checks))
idToPublicIDIdx := make(map[string]string)
for _, c := range checks {
checks := nmdata.PostureChecks{}
err := networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&c), reflect.ValueOf(&checks))
if err != nil {
return nil, nil, err
}
toret = append(toret, checks)
idToPublicIDIdx[checks.ID] = c.PublicID.String
}
return toret, idToPublicIDIdx, nil
}
type posturechecks struct {
ID string
PublicID sql.NullString `nmap:"skip"`
Checks json.RawMessage
}

View File

@@ -0,0 +1,75 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetRoutesQuery = `
select id, account_id, public_id, network, domains, keep_route, net_id, description,
peer, peer as peer_id, peer_groups, network_type, masquerade, metric, enabled,
groups, access_control_groups, skip_auto_apply
from routes
where account_id=$1
`
)
func (pg *PgStore) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, err
}
return GetRoutesViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetRoutesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Route, error) {
rows, err := con.Query(ctx, GetRoutesQuery, accountId)
if err != nil {
return nil, err
}
routes, err := pgx.CollectRows(rows, pgx.RowToStructByName[route])
if err != nil {
return nil, err
}
toret := make([]nmdata.Route, 0, len(routes))
for _, r := range routes {
route := nmdata.Route{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&r), reflect.ValueOf(&route))
if err != nil {
return nil, err
}
toret = append(toret, route)
}
return toret, nil
}
type route struct {
ID string
AccountID sql.NullString
PublicID sql.NullString
Network json.RawMessage
Domains json.RawMessage
KeepRoute sql.NullBool
NetID sql.NullString
Description sql.NullString
Peer sql.NullString
PeerID sql.NullString
PeerGroups json.RawMessage
NetworkType sql.NullInt64
Masquerade sql.NullBool
Metric sql.NullInt64
Enabled sql.NullBool
Groups json.RawMessage
AccessControlGroups json.RawMessage
SkipAutoApply sql.NullBool
}

View File

@@ -0,0 +1,67 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"github.com/jackc/pgx/v5"
)
const (
GetServicesQuery = `
select enabled, private, array (select json_array_elements_text(access_groups::json)) as access_groups, proxy_cluster, domain
from services
where account_id=$1
`
GetProxyTargetedDomainResourcesQuery = `
select t.target_id
from targets as t
join services as s on s.id = t.service_id
where s.account_id=$1 and s.enabled and not coalesce(s.terminated, false)
and t.enabled and t.target_type='domain' and t.target_id is not null
`
)
func (pg *PgStore) GetPrivateServices(ctx context.Context, accountId string) ([]Service, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, err
}
return GetPrivateServicesViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetPrivateServicesViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) ([]Service, error) {
rows, err := conn.Query(ctx, GetServicesQuery, accountId)
if err != nil {
return nil, err
}
return pgx.CollectRows(rows, pgx.RowToStructByName[Service])
}
func GetProxyTargetedDomainResourceIDsViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) (map[string]struct{}, error) {
rows, err := conn.Query(ctx, GetProxyTargetedDomainResourcesQuery, accountId)
if err != nil {
return nil, err
}
ids, err := pgx.CollectRows(rows, pgx.RowTo[string])
if err != nil {
return nil, err
}
toret := make(map[string]struct{}, len(ids))
for _, id := range ids {
toret[id] = struct{}{}
}
return toret, nil
}
type Service struct {
Enabled sql.NullBool
Private sql.NullBool
AccessGroups []string
ProxyCluster sql.NullString
Domain sql.NullString
}

View File

@@ -0,0 +1,219 @@
package networkmap_pgsql
import (
"database/sql"
"encoding/json"
"reflect"
"testing"
"time"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/stretchr/testify/assert"
)
func TestNullStringSupport(t *testing.T) {
src := withNullString{Name: sql.NullString{String: "string", Valid: true}}
dst := withString{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, withString{Name: "string"}, dst)
src = withNullString{Name: sql.NullString{Valid: false}}
dst = withString{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, withString{Name: ""}, dst)
}
func TestNullBoolSupport(t *testing.T) {
src := withNullBool{TrueOrFalse: sql.NullBool{Bool: true, Valid: true}}
dst := withBool{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, withBool{TrueOrFalse: true}, dst)
}
func TestRawJsonSupport(t *testing.T) {
jb, _ := json.Marshal(embeddedS{Name: "blob-name", SomeField: 1})
src := withRawJson{Blob: json.RawMessage(jb)}
dst := fromJson{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, fromJson{Blob: embeddedS{Name: "blob-name", SomeField: 1}}, dst)
src1 := withRawJson{}
dst1 := fromJson{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src1), reflect.ValueOf(&dst1)))
assert.Equal(t, fromJson{}, dst1)
}
func TestShouldSkipTag(t *testing.T) {
src5 := withSkipTag{Field: "shouldskip"}
dst5 := emptySkipTagTarget{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src5), reflect.ValueOf(&dst5)))
assert.Equal(t, emptySkipTagTarget{}, dst5)
}
func TestMapToTag(t *testing.T) {
src6 := withMapToTag{Field: "fieldvalue"}
dst6 := mapToTagTarget{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src6), reflect.ValueOf(&dst6)))
assert.Equal(t, mapToTagTarget{AnotherField: "fieldvalue"}, dst6)
}
func TestNullableInt64Support(t *testing.T) {
src := withInt64{Field: sql.NullInt64{Int64: int64(1), Valid: true}}
dst := int64Target{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, int64Target{Field: 1}, dst)
}
func TestNullableTimeSupport(t *testing.T) {
now := time.Now()
src := withNullableTime{Field: sql.NullTime{Time: now, Valid: true}}
dst := nullableTimeTarget{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, nullableTimeTarget{Field: now}, dst)
}
func TestNullableTimePointerSupport(t *testing.T) {
now := time.Now()
src := withNullableTime{Field: sql.NullTime{Time: now, Valid: true}}
dst := nullableTimePointerTarget{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, nullableTimePointerTarget{Field: &now}, dst)
}
func TestStringSLiceSupport(t *testing.T) {
src := withStringSlice{Field: []string{"one"}}
dst := withStringSlice{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, withStringSlice{Field: []string{"one"}}, dst)
}
func TestNullStringSLiceSupport(t *testing.T) {
src := withStringSlice{}
dst := withStringSlice{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, withStringSlice{}, dst)
}
func TestWithMultipleFields(t *testing.T) {
now := time.Now()
src := withMultipleFields{
Field1: sql.NullString{String: "aaa", Valid: true},
Field2: sql.NullBool{Bool: true, Valid: true},
Field3: sql.NullTime{Time: now, Valid: true},
Field4: sql.NullInt64{Int64: 1, Valid: true},
Field5: "another",
}
dst := multipleFieldsTarget{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, multipleFieldsTarget{
Field1: "aaa",
Field2: true,
Field3: now,
Field4: 1,
Field5: "another",
}, dst)
}
func TestEmptyPublicIdsFilled(t *testing.T) {
src := withEmptyPublicIds{}
dst := emptyPublicIdTarget{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.NotEmpty(t, dst.PublicID)
assert.NotEmpty(t, dst.PublicId)
}
type withNullString struct {
Name sql.NullString
}
type withString struct {
Name string
}
type withMultipleFields struct {
Field1 sql.NullString
Field2 sql.NullBool
Field3 sql.NullTime
Field4 sql.NullInt64
Field5 string
}
type multipleFieldsTarget struct {
Field1 string
Field2 bool
Field3 time.Time
Field4 int64
Field5 string
}
type withNullBool struct {
TrueOrFalse sql.NullBool
}
type withBool struct {
TrueOrFalse bool
}
type withRawJson struct {
Blob json.RawMessage
}
type embeddedS struct {
Name string
SomeField int
}
type fromJson struct {
Blob embeddedS
}
type withSkipTag struct {
Field string `nmap:"skip"`
}
type emptySkipTagTarget struct {
Field string
}
type withMapToTag struct {
Field string `nmap:"map_to:AnotherField"`
}
type mapToTagTarget struct {
AnotherField string
}
type withInt64 struct {
Field sql.NullInt64
}
type int64Target struct {
Field int
}
type withNullableTime struct {
Field sql.NullTime
}
type nullableTimeTarget struct {
Field time.Time
}
type nullableTimePointerTarget struct {
Field *time.Time
}
type withStringSlice struct {
Field []string
}
type withEmptyPublicIds struct {
PublicID sql.NullString
PublicId sql.NullString
}
type emptyPublicIdTarget struct {
PublicID string
PublicId string
}

View File

@@ -0,0 +1,68 @@
package networkmap_pgsql
import (
"context"
"github.com/jackc/pgx/v5"
)
const (
GetAllowedUserIdsQuery = `
select id, array (select json_array_elements_text(auto_groups::json)) as auto_groups
from users
where account_id=$1 and not blocked and not is_service_user
`
GetAllGroupIdQuery = `
select array_agg(id) from groups
where account_id=$1 and name='All'
`
)
func (pg *PgStore) GetAllowedUsers(ctx context.Context, accountId string) (map[string]struct{}, map[string][]string, error) {
c, err := pg.Pool.Acquire(ctx)
if err != nil {
return nil, nil, err
}
return GetAllowedUsersViaPgxConnection(ctx, c.Conn(), accountId)
}
func GetAllowedUsersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]struct{}, map[string][]string, error) {
rows, err := con.Query(ctx, GetAllowedUserIdsQuery, accountId)
if err != nil {
return nil, nil, err
}
users, err := pgx.CollectRows(rows, pgx.RowToStructByName[user])
if err != nil {
return nil, nil, err
}
rows, err = con.Query(ctx, GetAllGroupIdQuery, accountId)
if err != nil {
return nil, nil, err
}
allGroupIds, err := pgx.CollectOneRow(rows, pgx.RowTo[[]string])
if err != nil {
return nil, nil, err
}
userIdIdx := make(map[string]struct{})
groupIdToUserIds := make(map[string][]string)
for _, user := range users {
userIdIdx[user.ID] = struct{}{}
for _, groupId := range user.AutoGroups {
groupIdToUserIds[groupId] = append(groupIdToUserIds[groupId], user.ID)
}
for _, allgid := range allGroupIds {
groupIdToUserIds[allgid] = append(groupIdToUserIds[allgid], user.ID)
}
}
return userIdIdx, groupIdToUserIds, nil
}
type user struct {
ID string
AutoGroups []string
}

View File

@@ -7,6 +7,7 @@ import (
"crypto/tls"
"net/http"
"net/netip"
"os"
"slices"
"time"
@@ -28,6 +29,8 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/activity"
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
@@ -99,6 +102,22 @@ func (s *BaseServer) Store() store.Store {
})
}
func (s *BaseServer) NetworkMapStore() *networkmapdb.NetworkMapDBStoreImpl {
return Create(s, func() *networkmapdb.NetworkMapDBStoreImpl {
dsn := os.Getenv("NETBIRD_NMAP_STORE_DSN") // Todo: this needs to be hoocked up properly
if dsn == "" {
return nil
}
store, err := networkmap_pgsql.NewPostgresqlStore(context.Background(), dsn)
if err != nil {
log.Fatalf("failed to create network map store: %v", err)
}
return networkmapdb.NewNetworkMapDBStoreImpl(store, s.IntegratedValidator(), s.SettingsManager())
})
}
func (s *BaseServer) EventStore() activity.Store {
return Create(s, func() activity.Store {
var err error

View File

@@ -123,7 +123,7 @@ func (s *BaseServer) EphemeralManager() ephemeral.Manager {
func (s *BaseServer) NetworkMapController() network_map.Controller {
return Create(s, func() network_map.Controller {
return nmapcontroller.NewController(context.Background(), s.Store(), s.Metrics(), s.PeersUpdateManager(), s.AccountRequestBuffer(), s.IntegratedValidator(), s.SettingsManager(), s.DNSDomain(), s.ProxyController(), s.EphemeralManager(), s.Config)
return nmapcontroller.NewController(context.Background(), s.Store(), s.Metrics(), s.PeersUpdateManager(), s.AccountRequestBuffer(), s.IntegratedValidator(), s.SettingsManager(), s.DNSDomain(), s.ProxyController(), s.EphemeralManager(), s.Config, s.NetworkMapStore())
})
}

View File

@@ -4,10 +4,9 @@ import (
"encoding/base64"
"strconv"
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/management/server/types"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
)
@@ -84,6 +83,7 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel
enc := newComponentEncoder(c)
enc.indexAllPeers()
routerIdxs := enc.indexRouterPeers(c.RouterPeers)
enc.indexAllNetworkResources()
// Phase 2: gather every policy that any consumer references (peer-pair
// policies + resource-only policies) so encodeResourcePoliciesMap can
@@ -105,7 +105,6 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel
DnsSettings: enc.encodeDNSSettings(c.DNSSettings),
DnsDomain: in.DNSDomain,
CustomZoneDomain: c.CustomZoneDomain,
AgentVersions: enc.agentVersions,
Peers: enc.peers,
RouterPeerIndexes: routerIdxs,
Policies: policies,
@@ -130,7 +129,7 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel
// networkSerial returns c.Network.CurrentSerial() with a nil guard. The
// production path always populates c.Network, but the encoder is exported
// and a hand-built components struct may omit it.
func networkSerial(n *types.Network) uint64 {
func networkSerial(n *nmdata.Network) uint64 {
if n == nil {
return 0
}
@@ -143,16 +142,15 @@ type componentEncoder struct {
peerOrder map[string]uint32
peers []*proto.PeerCompact
agentVersionOrder map[string]uint32
agentVersions []string
networkIdToPublicId map[string]string
}
func newComponentEncoder(c *types.NetworkMapComponents) *componentEncoder {
return &componentEncoder{
components: c,
peerOrder: make(map[string]uint32, len(c.Peers)),
peers: make([]*proto.PeerCompact, 0, len(c.Peers)),
agentVersionOrder: make(map[string]uint32),
components: c,
peerOrder: make(map[string]uint32, len(c.Peers)),
peers: make([]*proto.PeerCompact, 0, len(c.Peers)),
networkIdToPublicId: make(map[string]string),
}
}
@@ -165,7 +163,7 @@ func (e *componentEncoder) indexAllPeers() {
}
}
func (e *componentEncoder) appendPeer(p *types.ComponentPeer) uint32 {
func (e *componentEncoder) appendPeer(p *nmdata.Peer) uint32 {
if idx, ok := e.peerOrder[p.ID]; ok {
return idx
}
@@ -179,7 +177,7 @@ func (e *componentEncoder) appendPeer(p *types.ComponentPeer) uint32 {
// (c.RouterPeers may contain peers not in c.Peers when validation rules drop
// them) and returns their wire indexes for the RouterPeerIndexes field. Must
// run before any encoder that resolves peer ids via e.peerOrder.
func (e *componentEncoder) indexRouterPeers(routers map[string]*types.ComponentPeer) []uint32 {
func (e *componentEncoder) indexRouterPeers(routers map[string]*nmdata.Peer) []uint32 {
if len(routers) == 0 {
return nil
}
@@ -193,6 +191,15 @@ func (e *componentEncoder) indexRouterPeers(routers map[string]*types.ComponentP
return out
}
func (e *componentEncoder) indexAllNetworkResources() {
for _, r := range e.components.NetworkResources {
if !r.Enabled {
continue
}
e.networkIdToPublicId[r.ID] = r.PublicID
}
}
func (e *componentEncoder) encodeGroups() []*proto.GroupCompact {
if len(e.components.Groups) == 0 {
return nil
@@ -206,10 +213,22 @@ func (e *componentEncoder) encodeGroups() []*proto.GroupCompact {
peerIdxs = append(peerIdxs, idx)
}
}
groupCompactResources := func() []*proto.ResourceCompact {
var toret []*proto.ResourceCompact
for _, r := range g.Resources {
if pr := e.resourceToProto(r); pr != nil {
toret = append(toret, pr)
}
}
return toret
}
out = append(out, &proto.GroupCompact{
Id: g.PublicID,
PeerIndexes: peerIdxs,
IsAll: g.IsGroupAll(),
Resources: groupCompactResources(),
})
}
return out
@@ -219,7 +238,7 @@ func (e *componentEncoder) encodeGroups() []*proto.GroupCompact {
// list and a map from policy pointer to the indexes of its emitted rules in
// that list — used by encodeResourcePoliciesMap to translate
// ResourcePoliciesMap[resourceID][]*Policy into wire-side indexes.
func (e *componentEncoder) encodePolicies(policies []*types.Policy) []*proto.PolicyCompact {
func (e *componentEncoder) encodePolicies(policies []*nmdata.Policy) []*proto.PolicyCompact {
if len(policies) == 0 {
return nil
}
@@ -241,7 +260,7 @@ func (e *componentEncoder) encodePolicies(policies []*types.Policy) []*proto.Pol
}
// encodePolicyRule maps a single PolicyRule under pol to a PolicyCompact entry.
func (e *componentEncoder) encodePolicyRule(pol *types.Policy, r *types.PolicyRule) *proto.PolicyCompact {
func (e *componentEncoder) encodePolicyRule(pol *nmdata.Policy, r *nmdata.PolicyRule) *proto.PolicyCompact {
return &proto.PolicyCompact{
Id: pol.PublicID,
Action: networkmap.GetProtoAction(string(r.Action)),
@@ -280,14 +299,14 @@ func (e *componentEncoder) groupPublicXids(src []string) []string {
// only live in ResourcePoliciesMap; without this union step they'd be lost
// from the wire and the client's resource-policy lookup would come back
// empty.
func unionPolicies(policies []*types.Policy, resourcePolicies map[string][]*types.Policy) []*types.Policy {
func unionPolicies(policies []*nmdata.Policy, resourcePolicies map[string][]*nmdata.Policy) []*nmdata.Policy {
// Fast path: non-router peers have no resource-only policies, so the
// "union" is identical to `policies`. Skip the dedup map allocation.
if len(resourcePolicies) == 0 {
return policies
}
seen := make(map[string]struct{}, len(policies))
out := make([]*types.Policy, 0, len(policies))
out := make([]*nmdata.Policy, 0, len(policies))
for _, p := range policies {
if p == nil {
continue
@@ -345,18 +364,31 @@ func (e *componentEncoder) groupPublicXid(groupID string) (string, bool) {
// peers array. For other resource types only the type string is shipped
// today (Calculate's resource-typed rule path consults SourceResource only
// for "peer" — other types fall through to group-based lookup).
func (e *componentEncoder) resourceToProto(r types.Resource) *proto.ResourceCompact {
if r.ID == "" && r.Type == "" {
func (e *componentEncoder) resourceToProto(r nmdata.Resource) *proto.ResourceCompact {
t, ok := proto.ResourceCompactType_value[string(r.Type)]
if !ok || t == 0 || r.ID == "" {
return nil
}
out := &proto.ResourceCompact{Type: string(r.Type)}
if r.Type == types.ResourceTypePeer && r.ID != "" {
if idx, ok := e.peerOrder[r.ID]; ok {
out.PeerIndexSet = true
out.PeerIndex = idx
if t == int32(proto.ResourceCompactType_peer) {
idx, ok := e.peerOrder[r.ID]
if !ok {
return nil
}
return &proto.ResourceCompact{
Type: proto.ResourceCompactType_peer,
ResourceId: &proto.ResourceCompact_PeerIndex{PeerIndex: idx},
}
}
return out
publicID, ok := e.networkIdToPublicId[r.ID]
if !ok {
return nil
}
return &proto.ResourceCompact{
Type: proto.ResourceCompactType(t),
ResourceId: &proto.ResourceCompact_Id{Id: publicID},
}
}
// postureCheckSeqs translates a slice of posture-check xids to their
@@ -389,7 +421,7 @@ func (e *componentEncoder) networkPublicId(xid string) (string, bool) {
return id, true
}
func (e *componentEncoder) encodeDNSSettings(s *types.DNSSettings) *proto.DNSSettingsCompact {
func (e *componentEncoder) encodeDNSSettings(s *nmdata.DNSSettings) *proto.DNSSettingsCompact {
if s == nil || len(s.DisabledManagementGroups) == 0 {
return nil
}
@@ -404,7 +436,7 @@ func (e *componentEncoder) encodeDNSSettings(s *types.DNSSettings) *proto.DNSSet
return out
}
func (e *componentEncoder) encodeRoutes(routes []*nbroute.Route) []*proto.RouteRaw {
func (e *componentEncoder) encodeRoutes(routes []*nmdata.Route) []*proto.RouteRaw {
if len(routes) == 0 {
return nil
}
@@ -442,7 +474,7 @@ func (e *componentEncoder) encodeRoutes(routes []*nbroute.Route) []*proto.RouteR
return out
}
func (e *componentEncoder) encodeNameServerGroups(nsgs []*nbdns.NameServerGroup) []*proto.NameServerGroupRaw {
func (e *componentEncoder) encodeNameServerGroups(nsgs []*nmdata.NameServerGroup) []*proto.NameServerGroupRaw {
if len(nsgs) == 0 {
return nil
}
@@ -465,7 +497,7 @@ func (e *componentEncoder) encodeNameServerGroups(nsgs []*nbdns.NameServerGroup)
return out
}
func encodeNameServers(servers []nbdns.NameServer) []*proto.NameServer {
func encodeNameServers(servers []nmdata.NameServer) []*proto.NameServer {
if len(servers) == 0 {
return nil
}
@@ -480,7 +512,7 @@ func encodeNameServers(servers []nbdns.NameServer) []*proto.NameServer {
return out
}
func encodeSimpleRecords(records []nbdns.SimpleRecord) []*proto.SimpleRecord {
func encodeSimpleRecords(records []nmdata.SimpleRecord) []*proto.SimpleRecord {
if len(records) == 0 {
return nil
}
@@ -497,7 +529,7 @@ func encodeSimpleRecords(records []nbdns.SimpleRecord) []*proto.SimpleRecord {
return out
}
func encodeCustomZones(zones []nbdns.CustomZone) []*proto.CustomZone {
func encodeCustomZones(zones []nmdata.CustomZone) []*proto.CustomZone {
if len(zones) == 0 {
return nil
}
@@ -513,7 +545,7 @@ func encodeCustomZones(zones []nbdns.CustomZone) []*proto.CustomZone {
return out
}
func (e *componentEncoder) encodeNetworkResources(resources []*types.ComponentResource) []*proto.NetworkResourceRaw {
func (e *componentEncoder) encodeNetworkResources(resources []*nmdata.NetworkResource) []*proto.NetworkResourceRaw {
if len(resources) == 0 {
return nil
}
@@ -542,7 +574,7 @@ func (e *componentEncoder) encodeNetworkResources(resources []*types.ComponentRe
return out
}
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*types.ComponentRouter) map[string]*proto.NetworkRouterList {
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*nmdata.NetworkRouter) map[string]*proto.NetworkRouterList {
if len(routersMap) == 0 {
return nil
}
@@ -578,7 +610,7 @@ func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*ty
return out
}
func (e *componentEncoder) encodeResourcePoliciesMap(rpm map[string][]*types.Policy) map[string]*proto.PolicyIds {
func (e *componentEncoder) encodeResourcePoliciesMap(rpm map[string][]*nmdata.Policy) map[string]*proto.PolicyIds {
if len(rpm) == 0 {
return nil
}
@@ -599,6 +631,9 @@ func (e *componentEncoder) encodeResourcePoliciesMap(rpm map[string][]*types.Pol
}
ids := make([]string, 0, len(policies))
for _, pol := range policies {
if pol == nil {
continue
}
ids = append(ids, pol.PublicID)
}
if len(ids) == 0 {
@@ -665,7 +700,7 @@ func (e *componentEncoder) encodePostureFailedPeers(m map[string]map[string]stru
// (which shouldn't happen in production but the encoder is exported)
// degrades to login_expiration_enabled = false, which makes
// LoginExpired() return false for every peer.
func toAccountSettingsCompact(s *types.AccountSettingsInfo) *proto.AccountSettingsCompact {
func toAccountSettingsCompact(s *nmdata.AccountSettingsInfo) *proto.AccountSettingsCompact {
if s == nil {
return &proto.AccountSettingsCompact{}
}
@@ -675,7 +710,7 @@ func toAccountSettingsCompact(s *types.AccountSettingsInfo) *proto.AccountSettin
}
}
func toAccountNetwork(n *types.Network) *proto.AccountNetwork {
func toAccountNetwork(n *nmdata.Network) *proto.AccountNetwork {
if n == nil {
return nil
}
@@ -691,20 +726,20 @@ func toAccountNetwork(n *types.Network) *proto.AccountNetwork {
return out
}
func toPeerCompact(p *types.ComponentPeer) *proto.PeerCompact {
func toPeerCompact(p *nmdata.Peer) *proto.PeerCompact {
pc := &proto.PeerCompact{
WgPubKey: decodeWgKey(p.Key),
SshPubKey: []byte(p.SSHKey),
DnsLabel: p.DNSLabel,
AgentVersion: p.AgentVersion,
AddedWithSsoLogin: p.AddedWithSSOLogin,
AgentVersion: p.Meta.WtVersion,
AddedWithSsoLogin: p.UserID != "",
LoginExpirationEnabled: p.LoginExpirationEnabled,
SshEnabled: p.SSHEnabled,
SupportsIpv6: p.SupportsIPv6,
SupportsSourcePrefixes: p.SupportsSourcePrefixes,
ServerSshAllowed: p.ServerSSHAllowed,
SupportsIpv6: p.SupportsIPv6(),
SupportsSourcePrefixes: p.SupportsSourcePrefixes(),
ServerSshAllowed: p.Meta.Flags.ServerSSHAllowed,
}
if !p.LastLogin.IsZero() {
if p.LastLogin != nil {
pc.LastLoginUnixNano = p.LastLogin.UnixNano()
}
switch {
@@ -753,7 +788,7 @@ func portsToUint32(ports []string) []uint32 {
return out
}
func portRangesToProto(ranges []types.RulePortRange) []*proto.PortInfo_Range {
func portRangesToProto(ranges []nmdata.RulePortRange) []*proto.PortInfo_Range {
if len(ranges) == 0 {
return nil
}

View File

@@ -16,7 +16,7 @@ import (
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/management/server/types"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
)
@@ -152,66 +152,66 @@ func envelopesEquivalent(a, b *proto.NetworkMapEnvelope) bool {
}
func newTestComponents() *types.NetworkMapComponents {
peerA := &types.ComponentPeer{
ID: "peer-a",
Key: testWgKeyA,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
DNSLabel: "peera",
SSHKey: "ssh-a",
AgentVersion: "0.40.0",
peerA := &nmdata.Peer{
ID: "peer-a",
Key: testWgKeyA,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
DNSLabel: "peera",
SSHKey: "ssh-a",
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
}
peerB := &types.ComponentPeer{
ID: "peer-b",
Key: testWgKeyB,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}),
DNSLabel: "peerb",
AgentVersion: "0.25.0",
peerB := &nmdata.Peer{
ID: "peer-b",
Key: testWgKeyB,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}),
DNSLabel: "peerb",
Meta: nmdata.PeerSystemMeta{WtVersion: "0.25.0"},
}
peerC := &types.ComponentPeer{
ID: "peer-c",
Key: testWgKeyC,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc",
AgentVersion: "0.40.0",
peerC := &nmdata.Peer{
ID: "peer-c",
Key: testWgKeyC,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc",
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
}
return &types.NetworkMapComponents{
PeerID: "peer-a",
Network: &types.Network{
Network: &nmdata.Network{
Identifier: "net-test",
Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)},
Serial: 7,
},
AccountSettings: &types.AccountSettingsInfo{
AccountSettings: &nmdata.AccountSettingsInfo{
PeerLoginExpirationEnabled: true,
PeerLoginExpiration: 2 * time.Hour,
},
Peers: map[string]*types.ComponentPeer{
Peers: map[string]*nmdata.Peer{
"peer-a": peerA,
"peer-b": peerB,
"peer-c": peerC,
},
Groups: map[string]*types.ComponentGroup{
"group-src": {ID: "group-src", PublicID: "1", Name: "Src", Peers: []string{"peer-a"}},
"group-dst": {ID: "group-dst", PublicID: "2", Name: "Dst", Peers: []string{"peer-b", "peer-c"}},
Groups: map[string]*nmdata.Group{
"group-src": {PublicID: "1", Name: "Src", Peers: []string{"peer-a"}},
"group-dst": {PublicID: "2", Name: "Dst", Peers: []string{"peer-b", "peer-c"}},
},
Policies: []*types.Policy{
Policies: []*nmdata.Policy{
{
ID: "pol-1",
PublicID: "10",
Enabled: true,
Rules: []*types.PolicyRule{{
ID: "rule-1", Enabled: true, Action: types.PolicyTrafficActionAccept,
Protocol: types.PolicyRuleProtocolTCP, Bidirectional: true,
Rules: []*nmdata.PolicyRule{{
ID: "rule-1", Enabled: true, Action: string(types.PolicyTrafficActionAccept),
Protocol: string(types.PolicyRuleProtocolTCP), Bidirectional: true,
Ports: []string{"22", "80"},
PortRanges: []types.RulePortRange{{Start: 8000, End: 8100}},
PortRanges: []nmdata.RulePortRange{{Start: 8000, End: 8100}},
Sources: []string{"group-src"},
Destinations: []string{"group-dst"},
}},
},
},
RouterPeers: map[string]*types.ComponentPeer{"peer-c": peerC},
RouterPeers: map[string]*nmdata.Peer{"peer-c": peerC},
}
}
@@ -304,6 +304,31 @@ func TestEncodeNetworkMapEnvelope_GroupsByAccountPublicId(t *testing.T) {
assert.Len(t, groupByID["2"].PeerIndexes, 2)
}
func TestEncodePolicy(t *testing.T) {
encoder := componentEncoder{peerOrder: map[string]uint32{"peerId": uint32(1234)}, networkIdToPublicId: map[string]string{"domain": "publicDomain", "host": "publicHost", "subnet": "publicSubnet"}}
assert.Equal(t,
encoder.resourceToProto(nmdata.Resource{Type: "peer", ID: "peerId"}),
&proto.ResourceCompact{Type: proto.ResourceCompactType_peer, ResourceId: &proto.ResourceCompact_PeerIndex{PeerIndex: uint32(1234)}})
// verify invalid peer id results in nil
assert.Nil(t,
encoder.resourceToProto(nmdata.Resource{Type: "peer", ID: "boom"}))
assert.Equal(t,
encoder.resourceToProto(nmdata.Resource{Type: "domain", ID: "domain"}),
&proto.ResourceCompact{Type: proto.ResourceCompactType_domain, ResourceId: &proto.ResourceCompact_Id{Id: "publicDomain"}})
assert.Equal(t,
encoder.resourceToProto(nmdata.Resource{Type: "host", ID: "host"}),
&proto.ResourceCompact{Type: proto.ResourceCompactType_host, ResourceId: &proto.ResourceCompact_Id{Id: "publicHost"}})
assert.Equal(t,
encoder.resourceToProto(nmdata.Resource{Type: "subnet", ID: "subnet"}),
&proto.ResourceCompact{Type: proto.ResourceCompactType_subnet, ResourceId: &proto.ResourceCompact_Id{Id: "publicSubnet"}})
// verify invalid resource type results in nil
assert.Nil(t,
encoder.resourceToProto(nmdata.Resource{Type: "boom", ID: "boom"}))
// verify invalid networkresource id results in nil
assert.Nil(t,
encoder.resourceToProto(nmdata.Resource{Type: "host", ID: "boom"}))
}
func TestEncodeNetworkMapEnvelope_PolicyExpansion(t *testing.T) {
c := newTestComponents()
@@ -377,12 +402,12 @@ func TestEncodeNetworkMapEnvelope_MalformedWgKey(t *testing.T) {
func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
c := newTestComponents()
v6Only := &types.ComponentPeer{
ID: "peer-v6",
Key: testWgKeyA,
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 9}),
DNSLabel: "peerv6",
AgentVersion: "0.40.0",
v6Only := &nmdata.Peer{
ID: "peer-v6",
Key: testWgKeyA,
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 9}),
DNSLabel: "peerv6",
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
}
c.Peers["peer-v6"] = v6Only
@@ -401,11 +426,11 @@ func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
func TestEncodeNetworkMapEnvelope_PeerWithoutIP(t *testing.T) {
c := newTestComponents()
c.Peers["peer-noip"] = &types.ComponentPeer{
ID: "peer-noip",
Key: testWgKeyA,
DNSLabel: "peernoip",
AgentVersion: "0.40.0",
c.Peers["peer-noip"] = &nmdata.Peer{
ID: "peer-noip",
Key: testWgKeyA,
DNSLabel: "peernoip",
Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
}
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -423,7 +448,7 @@ func TestEncodeNetworkMapEnvelope_PeerWithoutIP(t *testing.T) {
func TestEncodeNetworkMapEnvelope_EmptyInput(t *testing.T) {
c := &types.NetworkMapComponents{
Network: &types.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
Network: &nmdata.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
}
env := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c})
@@ -440,9 +465,9 @@ func TestEncodeNetworkMapEnvelope_EmptyInput(t *testing.T) {
func TestEncodeNetworkMapEnvelope_PeerLoginExpirationFields(t *testing.T) {
c := newTestComponents()
now := time.Date(2024, 1, 2, 3, 4, 5, 0, time.UTC)
c.Peers["peer-a"].AddedWithSSOLogin = true
c.Peers["peer-a"].UserID = "user-1"
c.Peers["peer-a"].LoginExpirationEnabled = true
c.Peers["peer-a"].LastLogin = now
c.Peers["peer-a"].LastLogin = &now
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -472,7 +497,7 @@ func TestEncodeNetworkMapEnvelope_PeerLoginExpirationFields(t *testing.T) {
func TestEncodeNetworkMapEnvelope_RoutesRoundTrip(t *testing.T) {
c := newTestComponents()
c.Routes = []*nbroute.Route{
c.Routes = []*nmdata.Route{
{
ID: "route-peer",
PublicID: "100",
@@ -519,7 +544,7 @@ func TestEncodeNetworkMapEnvelope_RoutesRoundTrip(t *testing.T) {
func TestEncodeNetworkMapEnvelope_RouteWithMissingPeerLeavesIndexUnset(t *testing.T) {
c := newTestComponents()
c.Routes = []*nbroute.Route{{
c.Routes = []*nmdata.Route{{
ID: "route-x",
PublicID: "100",
Peer: "peer-not-in-components",
@@ -539,21 +564,21 @@ func TestEncodeNetworkMapEnvelope_ResourceOnlyPolicyShippedAndIndexed(t *testing
// Policy that exists ONLY in ResourcePoliciesMap, not in c.Policies. This
// is the I1 case — without unionPolicies the encoder would silently
// drop it from the wire.
resourceOnlyPolicy := &types.Policy{
resourceOnlyPolicy := &nmdata.Policy{
ID: "pol-resource", PublicID: "99", Enabled: true,
Rules: []*types.PolicyRule{{
ID: "rule-r", Enabled: true, Action: types.PolicyTrafficActionAccept,
Protocol: types.PolicyRuleProtocolTCP,
Rules: []*nmdata.PolicyRule{{
ID: "rule-r", Enabled: true, Action: string(types.PolicyTrafficActionAccept),
Protocol: string(types.PolicyRuleProtocolTCP),
Sources: []string{"group-src"},
Destinations: []string{"group-dst"},
}},
}
c.ResourcePoliciesMap = map[string][]*types.Policy{
c.ResourcePoliciesMap = map[string][]*nmdata.Policy{
"resource-x": {c.Policies[0], resourceOnlyPolicy}, // shared + resource-only
}
// Resource must appear in components.NetworkResources with a seq id —
// encoder uses that to translate the xid map key to uint32.
c.NetworkResources = []*types.ComponentResource{
c.NetworkResources = []*nmdata.NetworkResource{
{ID: "resource-x", PublicID: "77", Name: "res-x", Enabled: true},
}
@@ -579,10 +604,10 @@ func TestEncodeNetworkMapEnvelope_ResourceOnlyPolicyShippedAndIndexed(t *testing
func TestEncodeNetworkMapEnvelope_NameServerGroups(t *testing.T) {
c := newTestComponents()
c.NameServerGroups = []*nbdns.NameServerGroup{{
c.NameServerGroups = []*nmdata.NameServerGroup{{
ID: "nsg-1", PublicID: "50", Name: "Main", Description: "primary",
NameServers: []nbdns.NameServer{{
IP: netip.MustParseAddr("8.8.8.8"), NSType: nbdns.UDPNameServerType, Port: 53,
NameServers: []nmdata.NameServer{{
IP: netip.MustParseAddr("8.8.8.8"), NSType: int(nbdns.UDPNameServerType), Port: 53,
}},
Groups: []string{"group-src", "group-not-persisted"},
Primary: true, Enabled: true,
@@ -621,11 +646,11 @@ func TestEncodeNetworkMapEnvelope_PostureFailedPeers(t *testing.T) {
func TestEncodeNetworkMapEnvelope_RoutersMap(t *testing.T) {
c := newTestComponents()
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
c.RoutersMap = map[string]map[string]*types.ComponentRouter{
c.RoutersMap = map[string]map[string]*nmdata.NetworkRouter{
"net-1": {
"peer-c": {
PublicID: "200",
Peer: "peer-c", Masquerade: true, Metric: 10, Enabled: true,
PublicID: "200",
Masquerade: true, Metric: 10, Enabled: true,
},
},
}
@@ -651,14 +676,14 @@ func TestEncodeNetworkMapEnvelope_RouterPeerNotInComponentsPeers(t *testing.T) {
// peer_index reference must still resolve.
c := newTestComponents()
delete(c.Peers, "peer-c")
routerPeer := &types.ComponentPeer{
routerPeer := &nmdata.Peer{
ID: "peer-c", Key: testWgKeyC, IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc", AgentVersion: "0.40.0",
DNSLabel: "peerc", Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
}
c.RouterPeers = map[string]*types.ComponentPeer{"peer-c": routerPeer}
c.RouterPeers = map[string]*nmdata.Peer{"peer-c": routerPeer}
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
c.RoutersMap = map[string]map[string]*types.ComponentRouter{
"net-1": {"peer-c": {PublicID: "1", Peer: "peer-c", Enabled: true}},
c.RoutersMap = map[string]map[string]*nmdata.NetworkRouter{
"net-1": {"peer-c": {PublicID: "1", Enabled: true}},
}
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -691,9 +716,9 @@ func TestToProxyPatch_EmptyInputReturnsNil(t *testing.T) {
func TestToProxyPatch_PopulatesAllFields(t *testing.T) {
nm := &types.NetworkMap{
Peers: []*types.ComponentPeer{{
Peers: []*nmdata.Peer{{
ID: "ext-peer", Key: testWgKeyA, IP: netip.AddrFrom4([4]byte{100, 64, 0, 9}),
DNSLabel: "extpeer", AgentVersion: "0.40.0",
DNSLabel: "extpeer", Meta: nmdata.PeerSystemMeta{WtVersion: "0.40.0"},
}},
FirewallRules: []*types.FirewallRule{{
PeerIP: "100.64.0.9", Action: "accept", Direction: 0, Protocol: "tcp",
@@ -765,7 +790,7 @@ func TestEncodeNetworkMapEnvelope_NilComponentsGracefulDegrade(t *testing.T) {
func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
c := &types.NetworkMapComponents{
Network: &types.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
Network: &nmdata.Network{Identifier: "x", Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)}},
// AccountSettings deliberately nil
}
@@ -779,8 +804,8 @@ func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
func emptyNetworkMapComponents() *types.NetworkMapComponents {
return types.EmptyNetworkMapComponents(
&types.NetworkMapComponents{
PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}},
Network: &types.Network{
PeerID: "peer-id", Peers: map[string]*nmdata.Peer{"peer-id": {}},
Network: &nmdata.Network{
Identifier: "net-empty",
Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)},
Serial: 9,

View File

@@ -7,11 +7,11 @@ import (
"github.com/netbirdio/netbird/client/ssh/auth"
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/posture"
"github.com/netbirdio/netbird/management/server/types"
sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc"
"github.com/netbirdio/netbird/shared/management/networkmap"
nmdata "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
)
@@ -31,14 +31,14 @@ func ToComponentSyncResponse(
config *nbconfig.Config,
httpConfig *nbconfig.HttpServerConfig,
deviceFlowConfig *nbconfig.DeviceAuthorizationFlow,
peer *nbpeer.Peer,
peer *nmdata.Peer,
turnCredentials *Token,
relayCredentials *Token,
components *types.NetworkMapComponents,
proxyPatch *types.NetworkMap,
dnsName string,
checks []*posture.Checks,
settings *types.Settings,
settings *nmdata.AccountSettingsInfo,
extraSettings *types.ExtraSettings,
peerGroups []string,
dnsFwdPort int64,
@@ -145,7 +145,7 @@ func toProxyPatch(nm *types.NetworkMap, dnsName string, includeIPv6, useSourcePr
//
// The full SSH AuthorizedUsers map is still produced by the client when it
// runs Calculate() over the envelope.
func computeSSHEnabledForPeer(c *types.NetworkMapComponents, peer *nbpeer.Peer) bool {
func computeSSHEnabledForPeer(c *types.NetworkMapComponents, peer *nmdata.Peer) bool {
if c == nil || peer == nil {
return false
}
@@ -170,25 +170,25 @@ func computeSSHEnabledForPeer(c *types.NetworkMapComponents, peer *nbpeer.Peer)
// ruleEnablesSSHForPeer returns true when rule is active, targets peer, and
// either explicitly authorises SSH or covers the legacy TCP/22 path while the
// peer itself has SSH enabled locally.
func ruleEnablesSSHForPeer(c *types.NetworkMapComponents, rule *types.PolicyRule, peer *nbpeer.Peer) bool {
func ruleEnablesSSHForPeer(c *types.NetworkMapComponents, rule *nmdata.PolicyRule, peer *nmdata.Peer) bool {
if rule == nil || !rule.Enabled {
return false
}
if !peerInDestinations(c, rule, peer.ID) {
return false
}
if rule.Protocol == types.PolicyRuleProtocolNetbirdSSH {
if rule.Protocol == string(types.PolicyRuleProtocolNetbirdSSH) {
return true
}
return peer.SSHEnabled && types.PolicyRuleImpliesLegacySSH(rule)
return peer.SSHEnabled && nmdata.PolicyRuleImpliesLegacySSH(rule)
}
// peerInDestinations reports whether peerID is in any of rule.Destinations'
// groups (or matches DestinationResource if it's a peer-typed resource —
// for non-peer types Calculate falls through to group lookup, so we mirror
// that exactly to avoid silent divergence).
func peerInDestinations(c *types.NetworkMapComponents, rule *types.PolicyRule, peerID string) bool {
if rule.DestinationResource.Type == types.ResourceTypePeer && rule.DestinationResource.ID != "" {
func peerInDestinations(c *types.NetworkMapComponents, rule *nmdata.PolicyRule, peerID string) bool {
if rule.DestinationResource.Type == string(types.ResourceTypePeer) && rule.DestinationResource.ID != "" {
return rule.DestinationResource.ID == peerID
}
for _, groupID := range rule.Destinations {

View File

@@ -5,8 +5,8 @@ import (
"github.com/stretchr/testify/assert"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
// TestComputeSSHEnabledForPeer covers both Calculate-mirroring branches:
@@ -17,16 +17,15 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
const targetPeerID = "target"
const targetGroupID = "g_dst"
mkComponents := func(rule *types.PolicyRule, sshEnabled bool) (*types.NetworkMapComponents, *nbpeer.Peer) {
peer := &nbpeer.Peer{ID: targetPeerID, SSHEnabled: sshEnabled}
group := &types.ComponentGroup{ID: targetGroupID, Name: "dst", Peers: []string{targetPeerID}}
mkComponents := func(rule *nmdata.PolicyRule, sshEnabled bool) (*types.NetworkMapComponents, *nmdata.Peer) {
peer := &nmdata.Peer{ID: targetPeerID, SSHEnabled: sshEnabled}
return &types.NetworkMapComponents{
Peers: map[string]*types.ComponentPeer{targetPeerID: peer.ToComponent()},
Groups: map[string]*types.ComponentGroup{targetGroupID: group},
Policies: []*types.Policy{{
Peers: map[string]*nmdata.Peer{targetPeerID: peer},
Groups: map[string]*nmdata.Group{targetGroupID: {Name: "dst", Peers: []string{targetPeerID}}},
Policies: []*nmdata.Policy{{
ID: "p",
Enabled: true,
Rules: []*types.PolicyRule{rule},
Rules: []*nmdata.PolicyRule{rule},
}},
}, peer
}
@@ -34,14 +33,14 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
cases := []struct {
name string
peerSSH bool
rule types.PolicyRule
rule nmdata.PolicyRule
wantEnabled bool
}{
{
name: "explicit-netbird-ssh-activates-regardless-of-peer-ssh",
peerSSH: false,
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -49,8 +48,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-tcp-22-with-peer-ssh",
peerSSH: true,
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22"},
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22"},
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -58,8 +57,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-tcp-22-without-peer-ssh-disabled",
peerSSH: false,
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22"},
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22"},
Destinations: []string{targetGroupID},
},
wantEnabled: false,
@@ -67,8 +66,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-tcp-22022-with-peer-ssh",
peerSSH: true,
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"22022"},
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"22022"},
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -76,8 +75,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-all-protocol-with-peer-ssh",
peerSSH: true,
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolALL,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolALL),
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -85,10 +84,10 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "implicit-port-range-covers-22",
peerSSH: true,
rule: types.PolicyRule{
rule: nmdata.PolicyRule{
Enabled: true,
Protocol: types.PolicyRuleProtocolTCP,
PortRanges: []types.RulePortRange{{Start: 20, End: 30}},
Protocol: string(types.PolicyRuleProtocolTCP),
PortRanges: []nmdata.RulePortRange{{Start: 20, End: 30}},
Destinations: []string{targetGroupID},
},
wantEnabled: true,
@@ -96,8 +95,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "tcp-80-no-ssh",
peerSSH: true,
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolTCP, Ports: []string{"80"},
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolTCP), Ports: []string{"80"},
Destinations: []string{targetGroupID},
},
wantEnabled: false,
@@ -105,8 +104,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "disabled-rule-skipped",
peerSSH: true,
rule: types.PolicyRule{
Enabled: false, Protocol: types.PolicyRuleProtocolNetbirdSSH,
rule: nmdata.PolicyRule{
Enabled: false, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
Destinations: []string{targetGroupID},
},
wantEnabled: false,
@@ -114,8 +113,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "peer-not-in-destinations",
peerSSH: true,
rule: types.PolicyRule{
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
rule: nmdata.PolicyRule{
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
Destinations: []string{"g_other"}, // target not in this group
},
wantEnabled: false,
@@ -123,21 +122,21 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
{
name: "peer-typed-destination-resource-matches",
peerSSH: false,
rule: types.PolicyRule{
rule: nmdata.PolicyRule{
Enabled: true,
Protocol: types.PolicyRuleProtocolNetbirdSSH,
DestinationResource: types.Resource{ID: targetPeerID, Type: types.ResourceTypePeer},
Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
DestinationResource: nmdata.Resource{ID: targetPeerID, Type: string(types.ResourceTypePeer)},
},
wantEnabled: true,
},
{
name: "non-peer-destination-resource-falls-through-to-groups",
peerSSH: false,
rule: types.PolicyRule{
rule: nmdata.PolicyRule{
Enabled: true,
Protocol: types.PolicyRuleProtocolNetbirdSSH,
DestinationResource: types.Resource{ID: targetPeerID, Type: "host"}, // wrong type
Destinations: []string{targetGroupID}, // saved by group fallback
Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
DestinationResource: nmdata.Resource{ID: targetPeerID, Type: "host"}, // wrong type
Destinations: []string{targetGroupID}, // saved by group fallback
},
wantEnabled: true,
},
@@ -156,16 +155,16 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
// belt-and-suspenders presence guard mirroring Calculate's
// getAllPeersFromGroups invariant.
func TestComputeSSHEnabledForPeer_TargetMissingFromComponents(t *testing.T) {
peer := &nbpeer.Peer{ID: "missing", SSHEnabled: true}
peer := &nmdata.Peer{ID: "missing", SSHEnabled: true}
c := &types.NetworkMapComponents{
Peers: map[string]*types.ComponentPeer{}, // target peer NOT present
Groups: map[string]*types.ComponentGroup{
"g": {ID: "g", Peers: []string{"missing"}},
Peers: map[string]*nmdata.Peer{}, // target peer NOT present
Groups: map[string]*nmdata.Group{
"g": {Peers: []string{"missing"}},
},
Policies: []*types.Policy{{
Policies: []*nmdata.Policy{{
ID: "p", Enabled: true,
Rules: []*types.PolicyRule{{
Enabled: true, Protocol: types.PolicyRuleProtocolNetbirdSSH,
Rules: []*nmdata.PolicyRule{{
Enabled: true, Protocol: string(types.PolicyRuleProtocolNetbirdSSH),
Destinations: []string{"g"},
}},
}},
@@ -179,6 +178,6 @@ func TestComputeSSHEnabledForPeer_TargetMissingFromComponents(t *testing.T) {
// exported indirectly via ToComponentSyncResponse and may receive nil
// components on graceful-degrade paths.
func TestComputeSSHEnabledForPeer_NilInputs(t *testing.T) {
assert.False(t, computeSSHEnabledForPeer(nil, &nbpeer.Peer{ID: "x"}))
assert.False(t, computeSSHEnabledForPeer(nil, &nmdata.Peer{ID: "x"}))
assert.False(t, computeSSHEnabledForPeer(&types.NetworkMapComponents{}, nil))
}

View File

@@ -18,10 +18,10 @@ import (
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache"
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/posture"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/netiputil"
)
@@ -47,7 +47,7 @@ func init() {
// nil when no server config is set (the fan-out network-map path) because clients treat any
// non-nil config as authoritative: a config without a relay section is interpreted as relay
// disabled and wipes the clients' relay URLs.
func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken *Token, extraSettings *types.ExtraSettings, settings *types.Settings) *proto.NetbirdConfig {
func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken *Token, extraSettings *types.ExtraSettings, settings *nmdata.AccountSettingsInfo) *proto.NetbirdConfig {
if config == nil {
return nil
}
@@ -119,7 +119,7 @@ func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken
return nbConfig
}
func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool, forceRoutingPeerDNS bool) *proto.PeerConfig {
func toPeerConfig(peer *nmdata.Peer, network *nmdata.Network, dnsName string, settings *nmdata.AccountSettingsInfo, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool, forceRoutingPeerDNS bool) *proto.PeerConfig {
netmask, _ := network.Net.Mask.Size()
fqdn := peer.FQDN(dnsName)
@@ -154,7 +154,7 @@ func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, set
return peerConfig
}
func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, peer *nbpeer.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*posture.Checks, dnsCache *cache.DNSConfigCache, settings *types.Settings, extraSettings *types.ExtraSettings, peerGroups []string, dnsFwdPort int64) *proto.SyncResponse {
func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, peer *nmdata.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*posture.Checks, dnsCache *cache.DNSConfigCache, settings *nmdata.AccountSettingsInfo, extraSettings *types.ExtraSettings, peerGroups []string, dnsFwdPort int64) *proto.SyncResponse {
// IPv6 data in AllowedIPs and SourcePrefixes wildcard expansion depends on
// whether the target peer supports IPv6. Routes and firewall rules are already
// filtered at the source (network map builder).

View File

@@ -278,7 +278,7 @@ func TestToNetbirdConfig_RelayInvariant(t *testing.T) {
settings := &types.Settings{MetricsPushEnabled: true}
t.Run("nil server config returns nil config", func(t *testing.T) {
nbCfg := toNetbirdConfig(nil, nil, nil, nil, settings)
nbCfg := toNetbirdConfig(nil, nil, nil, nil, types.TwinAccountSettings(settings))
assert.Nil(t, nbCfg, "fan-out updates must not carry a partial NetbirdConfig even when settings are present")
})
@@ -293,7 +293,7 @@ func TestToNetbirdConfig_RelayInvariant(t *testing.T) {
}
relayToken := &Token{Payload: "token-payload", Signature: "token-signature"}
nbCfg := toNetbirdConfig(cfg, nil, relayToken, nil, settings)
nbCfg := toNetbirdConfig(cfg, nil, relayToken, nil, types.TwinAccountSettings(settings))
require.NotNil(t, nbCfg)
require.NotNil(t, nbCfg.Relay, "non-nil NetbirdConfig must include the relay section")
assert.Equal(t, cfg.Relay.Addresses, nbCfg.Relay.Urls, "relay URLs should match the server config")
@@ -329,7 +329,7 @@ func TestToPeerConfig_RoutingPeerDNSResolution(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
settings := &types.Settings{RoutingPeerDNSResolutionEnabled: tt.globalFlag}
cfg := toPeerConfig(newPeer(tt.embedded), network, "netbird.selfhosted", settings, nil, nil, false, tt.forceParam)
cfg := toPeerConfig(types.TwinPeer(newPeer(tt.embedded)), types.TwinNetwork(network), "netbird.selfhosted", types.TwinAccountSettings(settings), nil, nil, false, tt.forceParam)
assert.Equal(t, tt.wantEnabled, cfg.RoutingPeerDnsResolutionEnabled,
"RoutingPeerDnsResolutionEnabled should reflect global || embedded || forced")
})

View File

@@ -29,6 +29,7 @@ import (
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/peers"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
@@ -36,7 +37,6 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/management/server/users"
proxyauth "github.com/netbirdio/netbird/proxy/auth"

View File

@@ -920,8 +920,8 @@ func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, ne
// if peer has reached this point then it has logged in
loginResp := &proto.LoginResponse{
NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil, settings),
PeerConfig: toPeerConfig(peer, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH, false),
NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil, types.TwinAccountSettings(settings)),
PeerConfig: toPeerConfig(types.TwinPeer(peer), types.TwinNetwork(network), s.networkMapController.GetDNSDomain(settings), types.TwinAccountSettings(settings), s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH, false),
Checks: toProtocolChecks(ctx, postureChecks),
}
@@ -1052,9 +1052,9 @@ func (s *Server) sendInitialSync(ctx context.Context, peerKey wgtypes.Key, peer
log.WithContext(ctx).Errorf("failed to build components for peer %s on initial sync: %v", peer.ID, err)
return status.Errorf(codes.Internal, "failed to build initial sync envelope")
}
plainResp = ToComponentSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, freshPeer, turnToken, relayToken, components, proxyPatch, dnsName, freshPostureChecks, settings, settings.Extra, peerGroups, freshDnsFwdPort)
plainResp = ToComponentSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, types.TwinPeer(freshPeer), turnToken, relayToken, components, proxyPatch, dnsName, freshPostureChecks, types.TwinAccountSettings(settings), settings.Extra, peerGroups, freshDnsFwdPort)
} else {
plainResp = ToSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, peer, turnToken, relayToken, networkMap, dnsName, postureChecks, nil, settings, settings.Extra, peerGroups, dnsFwdPort)
plainResp = ToSyncResponse(ctx, s.config, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, types.TwinPeer(peer), turnToken, relayToken, networkMap, dnsName, postureChecks, nil, types.TwinAccountSettings(settings), settings.Extra, peerGroups, dnsFwdPort)
}
key, err := s.secretsManager.GetWGKey()

View File

@@ -3330,7 +3330,7 @@ func createManager(t testing.TB) (*DefaultAccountManager, *update_channel.PeersU
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := NewAccountRequestBuffer(ctx, store)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{})
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nil)
manager, err := BuildManager(ctx, &config.Config{}, store, networkMapController, job.NewJobManager(nil, store, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
if err != nil {
return nil, nil, err

View File

@@ -102,11 +102,20 @@ func TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection(t *t
require.NotEmpty(t, before.Subdomain, "subdomain pinned at bootstrap")
assert.False(t, before.EnablePromptCollection, "prompt collection defaults off")
// Attempt to flip toggles AND smuggle a different cluster/subdomain — the
// immutable fields must be ignored.
// A cluster different from the one pinned at bootstrap must be rejected
// outright — never silently swapped or ignored.
_, err = mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{
AccountID: accountID,
Cluster: "attacker.cluster",
EnableLogCollection: true,
})
require.Error(t, err, "UpdateSettings with a mismatched cluster must fail")
// Flipping the toggles works with the pinned cluster echoed back (and
// with it omitted); the subdomain is never taken from the request.
updated, err := mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{
AccountID: accountID,
Cluster: "attacker.cluster",
Cluster: clusterAddr,
Subdomain: "evil",
EnableLogCollection: true,
EnablePromptCollection: true,

View File

@@ -234,7 +234,7 @@ func createDNSManager(t *testing.T) (*DefaultAccountManager, error) {
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := NewAccountRequestBuffer(ctx, store)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.test", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{})
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.test", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nil)
return BuildManager(context.Background(), nil, store, networkMapController, job.NewJobManager(nil, store, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
}

View File

@@ -6,6 +6,8 @@ import (
"fmt"
"slices"
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/rs/xid"
log "github.com/sirupsen/logrus"
@@ -744,6 +746,14 @@ func validateDeleteGroup(ctx context.Context, transaction store.Store, group *ty
return &GroupLinkError{"network router", linkedRouter.ID}
}
if isLinked, linkedService := isGroupLinkedToReverseProxyService(ctx, transaction, group.AccountID, group.ID); isLinked {
return &GroupLinkError{"reverse proxy service", linkedService.Domain}
}
if isLinked, linkedPolicy := isGroupLinkedToAgentNetworkPolicy(ctx, transaction, group.AccountID, group.ID); isLinked {
return &GroupLinkError{"agent network policy", linkedPolicy.Name}
}
return checkGroupLinkedToSettings(ctx, transaction, group)
}
@@ -875,6 +885,46 @@ func isGroupLinkedToNetworkRouter(ctx context.Context, transaction store.Store,
return false, nil
}
// isGroupLinkedToReverseProxyService checks if a group is used as an access group
// of a private reverse proxy service or as a bearer-auth distribution group.
func isGroupLinkedToReverseProxyService(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *service.Service) {
services, err := transaction.GetAccountServices(ctx, store.LockingStrengthNone, accountID)
if err != nil {
log.WithContext(ctx).Errorf("error retrieving reverse proxy services while checking group linkage: %v", err)
return false, nil
}
for _, svc := range services {
if svc.Private && slices.Contains(svc.AccessGroups, groupID) {
return true, svc
}
if svc.Auth.BearerAuth != nil && svc.Auth.BearerAuth.Enabled && slices.Contains(svc.Auth.BearerAuth.DistributionGroups, groupID) {
return true, svc
}
}
return false, nil
}
// isGroupLinkedToAgentNetworkPolicy checks if a group is used as a source group by any
// agent network policy in the account.
func isGroupLinkedToAgentNetworkPolicy(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *agentNetworkTypes.Policy) {
policies, err := transaction.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID)
if err != nil {
log.WithContext(ctx).Errorf("error retrieving agent network policies while checking group linkage: %v", err)
return false, nil
}
for _, policy := range policies {
if policy == nil {
continue
}
if slices.Contains(policy.SourceGroups, groupID) {
return true, policy
}
}
return false, nil
}
// areGroupChangesAffectPeers checks if any changes to the specified groups will affect peers.
// It fetches each collection once and checks all groupIDs against them in memory.
func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) {

View File

@@ -18,6 +18,8 @@ import (
"golang.org/x/exp/maps"
nbdns "github.com/netbirdio/netbird/dns"
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/management/server/groups"
"github.com/netbirdio/netbird/management/server/networks"
"github.com/netbirdio/netbird/management/server/networks/resources"
@@ -125,6 +127,21 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) {
"grp-for-integration",
"only service users with admin power can delete integration group",
},
{
"agent network policy",
"grp-for-agent-network-policy",
"agent network policy",
},
{
"reverse proxy private service access group",
"grp-for-rp-private",
"reverse proxy service",
},
{
"reverse proxy bearer distribution group",
"grp-for-rp-bearer",
"reverse proxy service",
},
}
for _, testCase := range testCases {
@@ -218,6 +235,17 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) {
groupIDs: []string{"grp-for-integration"},
expectedReasons: []string{"only service users with admin power can delete integration group"},
},
{
name: "agent network policy",
groupIDs: []string{"grp-for-agent-network-policy"},
expectedReasons: []string{"agent network policy"},
},
{
name: "reverse proxy services",
groupIDs: []string{"grp-for-rp-private", "grp-for-rp-bearer"},
expectedReasons: []string{"reverse proxy service", "reverse proxy service"},
expectedNotDeleted: []string{"grp-for-rp-private", "grp-for-rp-bearer"},
},
{
name: "successfully delete multiple groups",
groupIDs: []string{"group-1", "group-2"},
@@ -285,6 +313,65 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) {
}
}
func TestDefaultAccountManager_DeleteGroupUnlinkedFromReverseProxyService(t *testing.T) {
am, _, err := createManager(t)
require.NoError(t, err, "Failed to create account manager")
_, account, err := initTestGroupAccount(am)
require.NoError(t, err, "Failed to init testing account")
deletableGroups := []*types.Group{
{
ID: "grp-rp-bearer-disabled",
AccountID: account.Id,
Name: "Group only in a disabled bearer auth",
Issued: types.GroupIssuedAPI,
Peers: make([]string, 0),
},
{
ID: "grp-rp-nonprivate-access",
AccountID: account.Id,
Name: "Group only in a non-private service's access groups",
Issued: types.GroupIssuedAPI,
Peers: make([]string, 0),
},
}
for _, group := range deletableGroups {
require.NoError(t, am.CreateGroup(context.Background(), account.Id, groupAdminUserID, group))
}
// Disabled bearer auth and stale access groups on a non-private service
// are inert configuration and must not block group deletion.
services := []*rpservice.Service{
{
ID: "rp-svc-bearer-disabled",
AccountID: account.Id,
Domain: "bearer-disabled.services.example.com",
Auth: rpservice.AuthConfig{
BearerAuth: &rpservice.BearerAuthConfig{
Enabled: false,
DistributionGroups: []string{"grp-rp-bearer-disabled"},
},
},
},
{
ID: "rp-svc-nonprivate-access",
AccountID: account.Id,
Domain: "nonprivate.services.example.com",
Private: false,
AccessGroups: []string{"grp-rp-nonprivate-access"},
},
}
for _, svc := range services {
require.NoError(t, am.Store.CreateService(context.Background(), svc))
}
for _, group := range deletableGroups {
err = am.DeleteGroup(context.Background(), account.Id, groupAdminUserID, group.ID)
assert.NoError(t, err, "group %s is not referenced by an active reverse proxy gate and should be deletable", group.ID)
}
}
func TestDefaultAccountManager_DeleteGroupLinkedToFlowGroup(t *testing.T) {
am, _, err := createManager(t)
require.NoError(t, err)
@@ -406,6 +493,30 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t
Peers: make([]string, 0),
}
groupForAgentNetworkPolicy := &types.Group{
ID: "grp-for-agent-network-policy",
AccountID: "account-id",
Name: "Group for agent network policies",
Issued: types.GroupIssuedAPI,
Peers: make([]string, 0),
}
groupForRPPrivate := &types.Group{
ID: "grp-for-rp-private",
AccountID: "account-id",
Name: "Group for private reverse proxy service",
Issued: types.GroupIssuedAPI,
Peers: make([]string, 0),
}
groupForRPBearer := &types.Group{
ID: "grp-for-rp-bearer",
AccountID: "account-id",
Name: "Group for bearer reverse proxy service",
Issued: types.GroupIssuedAPI,
Peers: make([]string, 0),
}
routeResource := &route.Route{
ID: "example route",
Groups: []string{groupForRoute.ID},
@@ -461,6 +572,66 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForSetupKeys)
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForUsers)
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForIntegration)
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForAgentNetworkPolicy)
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPPrivate)
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPBearer)
agentNetworkPolicy := &agentNetworkTypes.Policy{
ID: "example agent network policy",
AccountID: accountID,
Name: "Example agent network policy",
Enabled: true,
SourceGroups: []string{groupForAgentNetworkPolicy.ID},
}
if err := am.Store.SaveAgentNetworkPolicy(context.Background(), agentNetworkPolicy); err != nil {
return nil, nil, err
}
// The decoy services are created first so the linkage check has to scan
// past services that do not reference the groups under test.
rpServices := []*rpservice.Service{
{
ID: "rp-svc-private-decoy",
AccountID: accountID,
Domain: "private-decoy.services.example.com",
Private: true,
AccessGroups: []string{"unrelated-group"},
},
{
ID: "rp-svc-bearer-decoy",
AccountID: accountID,
Domain: "bearer-decoy.services.example.com",
Auth: rpservice.AuthConfig{
BearerAuth: &rpservice.BearerAuthConfig{
Enabled: true,
DistributionGroups: []string{"unrelated-group"},
},
},
},
{
ID: "rp-svc-private",
AccountID: accountID,
Domain: "private.services.example.com",
Private: true,
AccessGroups: []string{groupForRPPrivate.ID},
},
{
ID: "rp-svc-bearer",
AccountID: accountID,
Domain: "bearer.services.example.com",
Auth: rpservice.AuthConfig{
BearerAuth: &rpservice.BearerAuthConfig{
Enabled: true,
DistributionGroups: []string{groupForRPBearer.ID},
},
},
},
}
for _, svc := range rpServices {
if err := am.Store.CreateService(context.Background(), svc); err != nil {
return nil, nil, err
}
}
acc, err := am.Store.GetAccount(context.Background(), account.Id)
if err != nil {

View File

@@ -6,7 +6,6 @@ import (
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/activity"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
@@ -31,10 +30,6 @@ type managerImpl struct {
accountManager account.Manager
}
func eventMetaResource(group *types.Group, resource *resourceTypes.NetworkResource) map[string]any {
return map[string]any{"name": group.Name, "id": group.ID, "resource_name": resource.Name, "resource_id": resource.ID, "resource_type": resource.Type}
}
type mockManager struct {
}
@@ -114,7 +109,7 @@ func (m *managerImpl) AddResourceToGroupInTransaction(ctx context.Context, trans
}
event := func() {
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceAddedToGroup, eventMetaResource(group, networkResource))
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceAddedToGroup, group.EventMetaResource(types.TwinNetworkResource(networkResource)))
}
return event, nil
@@ -138,7 +133,7 @@ func (m *managerImpl) RemoveResourceFromGroupInTransaction(ctx context.Context,
}
event := func() {
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceRemovedFromGroup, eventMetaResource(group, networkResource))
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceRemovedFromGroup, group.EventMetaResource(types.TwinNetworkResource(networkResource)))
}
return event, nil

View File

@@ -446,7 +446,7 @@ func (h *Handler) GetAccessiblePeers(w http.ResponseWriter, r *http.Request) {
netMap := account.GetPeerNetworkMapFromComponents(ctx, peerID, dns.CustomZone{}, nil, validPeers, account.GetResourcePoliciesMap(), account.GetResourceRoutersMap(), nil, account.GetActiveGroupUsers())
util.WriteJSONObject(ctx, w, toAccessiblePeers(netMap, account.Peers, dnsDomain))
util.WriteJSONObject(ctx, w, toAccessiblePeers(account.Peers, netMap, dnsDomain))
}
func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request) {
@@ -534,20 +534,22 @@ func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request)
util.WriteJSONObject(r.Context(), w, resp)
}
// toAccessiblePeers rehydrates the calculated map's component peers into the
// account's full peer objects, which carry the location/status/meta fields
// the API response needs.
func toAccessiblePeers(netMap *types.NetworkMap, accountPeers map[string]*nbpeer.Peer, dnsDomain string) []api.AccessiblePeer {
// toAccessiblePeers resolves the twin peers in netMap back to the full account
// peers (by ID) so the API response keeps Status/Name/OS/GeoNameID, which the
// slim netmap twins intentionally don't carry.
func toAccessiblePeers(accountPeers map[string]*nbpeer.Peer, netMap *types.NetworkMap, dnsDomain string) []api.AccessiblePeer {
accessiblePeers := make([]api.AccessiblePeer, 0, len(netMap.Peers)+len(netMap.OfflinePeers))
add := func(peers []*types.ComponentPeer) {
for _, p := range peers {
if peer := accountPeers[p.ID]; peer != nil {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(peer, dnsDomain))
}
appendByID := func(id string) {
if p, ok := accountPeers[id]; ok && p != nil {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
}
}
add(netMap.Peers)
add(netMap.OfflinePeers)
for _, p := range netMap.Peers {
appendByID(p.ID)
}
for _, p := range netMap.OfflinePeers {
appendByID(p.ID)
}
return accessiblePeers
}

View File

@@ -96,7 +96,7 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee
}
requestBuffer := server.NewAccountRequestBuffer(ctx, store)
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{})
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{}, nil)
am, err := server.BuildManager(ctx, nil, store, networkMapController, jobManager, nil, "", &activity.InMemoryEventStore{}, geoMock, false, validatorMock, metrics, proxyController, settingsManager, permissionsManager, false, cacheStore)
if err != nil {
t.Fatalf("Failed to create manager: %v", err)
@@ -226,7 +226,7 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin
}
requestBuffer := server.NewAccountRequestBuffer(ctx, store)
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{})
networkMapController := controller.NewController(ctx, store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsManager, "", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peersManager), &config.Config{}, nil)
am, err := server.BuildManager(ctx, nil, store, networkMapController, jobManager, nil, "", &activity.InMemoryEventStore{}, geoMock, false, validatorMock, metrics, proxyController, settingsManager, permissionsManager, false, cacheStore)
if err != nil {
t.Fatalf("Failed to create manager: %v", err)

View File

@@ -92,7 +92,7 @@ func createManagerWithEmbeddedIdP(t testing.TB) (*DefaultAccountManager, *update
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := NewAccountRequestBuffer(ctx, testStore)
networkMapController := controller.NewController(ctx, testStore, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(testStore, peersManager), &config.Config{})
networkMapController := controller.NewController(ctx, testStore, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(testStore, peersManager), &config.Config{}, nil)
manager, err := BuildManager(ctx, &config.Config{}, testStore, networkMapController, job.NewJobManager(nil, testStore, peersManager), idpManager, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
if err != nil {
return nil, nil, err

View File

@@ -11,6 +11,7 @@ import (
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
// UpdateIntegratedValidator updates the integrated validator groups for a specified account.
@@ -109,7 +110,7 @@ func (am *DefaultAccountManager) GetValidatedPeers(ctx context.Context, accountI
return nil, nil, err
}
validPeers, err := am.integratedPeerValidator.GetValidatedPeers(ctx, accountID, groups, peers, settings.Extra)
validPeers, err := am.integratedPeerValidator.GetValidatedPeers(ctx, accountID, types.TwinGroups(groups), types.TwinPeers(peers), settings.Extra)
if err != nil {
return nil, nil, err
}
@@ -138,7 +139,7 @@ func (a MockIntegratedValidator) ValidatePeer(_ context.Context, update *nbpeer.
return update, false, nil
}
func (a MockIntegratedValidator) GetValidatedPeers(_ context.Context, accountID string, groups []*types.Group, peers []*nbpeer.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error) {
func (a MockIntegratedValidator) GetValidatedPeers(_ context.Context, accountID string, groups []*nmdata.Group, peers []*nmdata.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error) {
validatedPeers := make(map[string]struct{})
for _, peer := range peers {
validatedPeers[peer.ID] = struct{}{}

View File

@@ -5,6 +5,7 @@ import (
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
)
@@ -14,7 +15,7 @@ type IntegratedValidator interface {
ValidatePeer(ctx context.Context, update *nbpeer.Peer, peer *nbpeer.Peer, userID string, accountID string, dnsDomain string, peersGroup []string, extraSettings *types.ExtraSettings) (*nbpeer.Peer, bool, error)
PreparePeer(ctx context.Context, accountID string, peer *nbpeer.Peer, peersGroup []string, extraSettings *types.ExtraSettings, temporary bool) *nbpeer.Peer
IsNotValidPeer(ctx context.Context, accountID string, peer *nbpeer.Peer, peersGroup []string, extraSettings *types.ExtraSettings) (bool, bool, error)
GetValidatedPeers(ctx context.Context, accountID string, groups []*types.Group, peers []*nbpeer.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error)
GetValidatedPeers(ctx context.Context, accountID string, groups []*nmdata.Group, peers []*nmdata.Peer, extraSettings *types.ExtraSettings) (map[string]struct{}, error)
GetInvalidPeers(ctx context.Context, accountID string, extraSettings *types.ExtraSettings) (map[string]string, error)
PeerDeleted(ctx context.Context, accountID, peerID string, extraSettings *types.ExtraSettings) error
SetPeerInvalidationListener(fn func(accountID string, peerIDs []string))

View File

@@ -10,6 +10,7 @@ import (
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
)
@@ -35,7 +36,7 @@ func (v *IntegratedValidatorImpl) IsNotValidPeer(_ context.Context, _ string, _
return false, false, nil
}
func (v *IntegratedValidatorImpl) GetValidatedPeers(_ context.Context, _ string, _ []*types.Group, peers []*nbpeer.Peer, _ *types.ExtraSettings) (map[string]struct{}, error) {
func (v *IntegratedValidatorImpl) GetValidatedPeers(_ context.Context, _ string, _ []*nmdata.Group, peers []*nmdata.Peer, _ *types.ExtraSettings) (map[string]struct{}, error) {
validatedPeers := make(map[string]struct{})
for _, p := range peers {
validatedPeers[p.ID] = struct{}{}

View File

@@ -376,7 +376,7 @@ func startManagementForTest(t *testing.T, testFile string, config *config.Config
return nil, nil, "", cleanup, err
}
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeralMgr, config)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeralMgr, config, nil)
accountManager, err := BuildManager(ctx, nil, store, networkMapController, jobManager, nil, "",
eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)

View File

@@ -216,7 +216,7 @@ func startServer(
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := server.NewAccountRequestBuffer(ctx, str)
networkMapController := controller.NewController(ctx, str, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(str, peers.NewManager(str, permissionsManager)), config)
networkMapController := controller.NewController(ctx, str, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(str, peers.NewManager(str, permissionsManager)), config, nil)
accountManager, err := server.BuildManager(
context.Background(),

View File

@@ -803,7 +803,7 @@ func createNSManager(t *testing.T) (*DefaultAccountManager, error) {
updateManager := update_channel.NewPeersUpdateManager(metrics)
requestBuffer := NewAccountRequestBuffer(ctx, store)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{})
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), ephemeral_manager.NewEphemeralManager(store, peers.NewManager(store, permissionsManager)), &config.Config{}, nil)
return BuildManager(context.Background(), nil, store, networkMapController, job.NewJobManager(nil, store, peersManager), nil, "", eventStore, nil, false, MockIntegratedValidator{}, metrics, port_forwarding.NewControllerMock(), settingsMockManager, permissionsManager, false, cacheStore)
}

View File

@@ -14,7 +14,6 @@ import (
nbDomain "github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/http/api"
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
)
type NetworkResourceType string
@@ -65,27 +64,6 @@ func NewNetworkResource(accountID, networkID, name, description, address string,
}, nil
}
// ToComponent converts the resource to its self-contained components
// representation. Returns nil for a nil resource.
func (n *NetworkResource) ToComponent() *sharedTypes.ComponentResource {
if n == nil {
return nil
}
return &sharedTypes.ComponentResource{
ID: n.ID,
PublicID: n.PublicID,
NetworkID: n.NetworkID,
AccountID: n.AccountID,
Name: n.Name,
Description: n.Description,
Type: sharedTypes.ComponentResourceType(n.Type),
Address: n.Address,
Domain: n.Domain,
Prefix: n.Prefix,
Enabled: n.Enabled,
}
}
func (n *NetworkResource) ToAPIResponse(groups []api.GroupMinimum) *api.NetworkResource {
addr := n.Prefix.String()
if n.Type == Domain {

View File

@@ -7,7 +7,6 @@ import (
"github.com/netbirdio/netbird/management/server/networks/types"
"github.com/netbirdio/netbird/shared/management/http/api"
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
)
type NetworkRouter struct {
@@ -22,36 +21,6 @@ type NetworkRouter struct {
Enabled bool
}
// ToComponent converts the router to its self-contained components
// representation. Returns nil for a nil router.
func (n *NetworkRouter) ToComponent() *sharedTypes.ComponentRouter {
if n == nil {
return nil
}
return &sharedTypes.ComponentRouter{
NetworkID: n.NetworkID,
PublicID: n.PublicID,
Peer: n.Peer,
PeerGroups: n.PeerGroups,
Masquerade: n.Masquerade,
Metric: n.Metric,
Enabled: n.Enabled,
}
}
// ToComponentMap converts a peer-keyed router map to its components
// representation.
func ToComponentMap(routers map[string]*NetworkRouter) map[string]*sharedTypes.ComponentRouter {
if routers == nil {
return nil
}
out := make(map[string]*sharedTypes.ComponentRouter, len(routers))
for id, r := range routers {
out[id] = r.ToComponent()
}
return out
}
func NewNetworkRouter(accountID string, networkID string, peer string, peerGroups []string, masquerade bool, metric int, enabled bool) (*NetworkRouter, error) {
r := &NetworkRouter{
ID: xid.New().String(),

View File

@@ -21,6 +21,7 @@ import (
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/management/server/posture"
"github.com/netbirdio/netbird/management/server/store"
@@ -1588,7 +1589,7 @@ func affectedPeerIDsFromNetworkMap(nmap *types.NetworkMap, selfPeerID string) []
}
seen := make(map[string]struct{}, len(nmap.Peers)+len(nmap.OfflinePeers))
ids := make([]string, 0, len(nmap.Peers)+len(nmap.OfflinePeers))
add := func(peers []*types.ComponentPeer) {
add := func(peers []*nmdata.Peer) {
for _, p := range peers {
if p == nil || p.ID == "" || p.ID == selfPeerID {
continue

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