Compare commits

..

1 Commits

Author SHA1 Message Date
Viktor Liu
b64080400d Enforce reverse proxy group access before minting and when honouring a session cookie 2026-08-18 15:04:16 +02:00
391 changed files with 4829 additions and 27310 deletions

View File

@@ -730,11 +730,6 @@ jobs:
- name: Install modules
run: go mod tidy
- name: Run Mage
uses: magefile/mage-action@a662bd8c29d8106879588cfff83b2faf6e6f59db # v4.0.0
with:
install-only: true
- name: check git status
run: git --no-pager diff --exit-code
@@ -743,7 +738,9 @@ jobs:
CGO_ENABLED=1 GOARCH=${{ matrix.arch }} \
NETBIRD_STORE_ENGINE=${{ matrix.store }} \
CI=true \
mage integrationtest:all -gotestflags="-coverprofile=coverage.txt"
go test -tags=integration -coverprofile=coverage.txt \
-exec 'sudo --preserve-env=CI,NETBIRD_STORE_ENGINE' \
-timeout 20m ./management/server/http/...
- name: Upload coverage reports to Codecov
if: matrix.arch == 'amd64'

View File

@@ -26,8 +26,6 @@ import (
"github.com/netbirdio/netbird/client/internal/routemanager"
"github.com/netbirdio/netbird/client/internal/stdnet"
"github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/netstate"
"github.com/netbirdio/netbird/client/netsweep"
"github.com/netbirdio/netbird/client/system"
"github.com/netbirdio/netbird/formatter"
"github.com/netbirdio/netbird/route"
@@ -42,6 +40,11 @@ const (
AnonymizeLevelStrict = nbAnonymize.LevelStrictString
)
// ConnectionListener export internal Listener for mobile
type ConnectionListener interface {
peer.Listener
}
// TunAdapter export internal TunAdapter for mobile
type TunAdapter interface {
device.TunAdapter
@@ -82,13 +85,6 @@ type Client struct {
deviceName string
uiVersion string
networkChangeListener listener.NetworkChangeListener
// netState outlives engine restarts: it mirrors the OS connectivity, not
// the engine lifecycle. Run and RunWithoutLogin inject it into each new
// ConnectClient, which distributes it to every reconnection loop.
netState *netstate.State
// sweeper also outlives engine restarts; NotifyNetworkChange sweeps it.
sweeper *netsweep.Sweeper
stateMu sync.RWMutex
connectClient *internal.ConnectClient
@@ -152,7 +148,6 @@ func NewClient(androidSDKVersion int, deviceName string, uiVersion string, tunAd
execWorkaround(androidSDKVersion)
net.SetAndroidProtectSocketFn(tunAdapter.ProtectSocket)
system.SetIFaceDiscover(iFaceDiscover)
return &Client{
deviceName: deviceName,
uiVersion: uiVersion,
@@ -161,8 +156,6 @@ func NewClient(androidSDKVersion int, deviceName string, uiVersion string, tunAd
recorder: peer.NewRecorder(""),
ctxCancelLock: &sync.Mutex{},
networkChangeListener: networkChangeListener,
netState: netstate.New(),
sweeper: netsweep.New(),
}
}
@@ -203,8 +196,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,
internal.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper))
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder)
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
@@ -245,8 +237,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,
internal.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper))
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder)
c.setState(cfg, cacheDir, cfgFile, connectClient)
return connectClient.RunOnAndroid(c.tunAdapter, c.iFaceDiscover, c.networkChangeListener, slices.Clone(dns.items), dnsReadyListener, stateFile, cacheDir)
}
@@ -294,24 +285,6 @@ func (c *Client) GetTunSettings() (*TunSettings, error) {
}, nil
}
// SetNetworkAvailable feeds OS-reported network availability into the client.
// While unavailable, the internal reconnect loops suspend their attempts and
// the connection listener reports NoNetwork instead of Connecting; when
// availability returns, the loops resume immediately with a fresh backoff.
func (c *Client) SetNetworkAvailable(available bool) {
c.netState.Set(available)
c.recorder.SetNetworkAvailable(available)
}
// NotifyNetworkChange marks the management, signal and relay connections
// stale after the OS switched networks and schedules a sweep that cuts
// whatever has not redialed on the new network by then. The engine and the
// TUN device stay untouched.
func (c *Client) NotifyNetworkChange() {
c.sweeper.MarkNetworkChange()
log.Infof("network change: connections marked stale")
}
// DebugBundle generates a debug bundle, uploads it, and returns the upload key.
// It works both with and without a running engine. anonymizeLevel is "default"
// or "strict"; strict also anonymizes internal IP ranges, peer names, and
@@ -552,11 +525,7 @@ func (c *Client) OnUpdatedHostDNS(list *DNSList) error {
// SetConnectionListener set the network connection listener
func (c *Client) SetConnectionListener(listener ConnectionListener) {
if listener == nil {
c.recorder.RemoveConnectionListener()
return
}
c.recorder.SetConnectionListener(connectionListenerAdapter{listener})
c.recorder.SetConnectionListener(listener)
}
// RemoveConnectionListener remove connection listener

View File

@@ -1,41 +0,0 @@
//go:build android
package android
import (
"github.com/netbirdio/netbird/client/internal/peer"
)
// Client state values delivered via ConnectionListener.OnStateChanged,
// re-exported as basic constants so gomobile emits them into the generated
// Java bindings. They mirror peer.ClientState*: append-only, never reorder.
const (
ClientStateDisconnected = int(peer.ClientStateDisconnected)
ClientStateConnected = int(peer.ClientStateConnected)
ClientStateConnecting = int(peer.ClientStateConnecting)
ClientStateDisconnecting = int(peer.ClientStateDisconnecting)
ClientStateNoNetwork = int(peer.ClientStateNoNetwork)
)
// ConnectionListener export internal Listener for mobile. It mirrors
// peer.Listener with OnStateChanged taking a plain int (one of the
// ClientState* constants), because gomobile cannot bind named types.
type ConnectionListener interface {
OnStateChanged(state int)
OnConnected()
OnDisconnected()
OnConnecting()
OnDisconnecting()
OnAddressChanged(string, string)
OnPeersListChanged(int)
}
// connectionListenerAdapter adapts the gomobile-facing ConnectionListener to
// peer.Listener, converting the typed state to the int the binding carries.
type connectionListenerAdapter struct {
ConnectionListener
}
func (a connectionListenerAdapter) OnStateChanged(state peer.ClientState) {
a.ConnectionListener.OnStateChanged(int(state))
}

View File

@@ -191,49 +191,40 @@ func (a *Auth) login(urlOpener URLOpener, isAndroidTV bool) error {
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, profileLoginHint(a.cfgPath))
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, isAndroidTV)
if err != nil {
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
}
return runOAuthFlow(a.ctx, oAuthFlow, urlOpener, nil)
}
// profileLoginHint returns the stored account email for the profile at cfgPath.
// An empty hint is deliberate, not a fallback: a fresh profile leaves the
// choice to the IdP. Switching accounts is done by switching or removing
// profiles, not by logging out — logout keeps the email.
func profileLoginHint(cfgPath string) string {
if cfgPath == "" {
return ""
// An empty hint is deliberate, not a fallback: a fresh profile leaves the
// choice to the IdP. Switching accounts is done by switching or removing
// profiles, not by logging out — logout keeps the email.
if a.cfgPath != "" {
if hint := readProfileEmail(a.cfgPath); hint != "" {
if setter, ok := oAuthFlow.(loginHintSetter); ok {
setter.SetLoginHint(hint)
}
}
}
return readProfileEmail(cfgPath)
}
// runOAuthFlow drives an already acquired OAuth flow to a token: requests the
// flow info, presents the verification URL through the opener and waits for
// the browser round-trip. Open is called synchronously — it is what marks the
// surface as opened on the client side, and a fast token's OnLoginSuccess is
// a no-op until it has, so the dismissal would be dropped rather than
// delayed. Openers must therefore not block: they post their UI work and
// return. onWaiting, when set, runs after the URL is shown, right before the
// blocking wait.
func runOAuthFlow(ctx context.Context, flow auth.OAuthFlow, urlOpener URLOpener, onWaiting func()) (*auth.TokenInfo, error) {
flowInfo, err := flow.RequestAuthInfo(ctx)
flowInfo, err := oAuthFlow.RequestAuthInfo(context.TODO())
if err != nil {
return nil, fmt.Errorf("request auth info: %w", err)
return nil, fmt.Errorf("getting a request OAuth flow info failed: %v", err)
}
urlOpener.Open(flowInfo.VerificationURIComplete, flowInfo.UserCode)
go urlOpener.Open(flowInfo.VerificationURIComplete, flowInfo.UserCode)
if onWaiting != nil {
onWaiting()
}
tokenInfo, err := flow.WaitToken(ctx, flowInfo)
tokenInfo, err := oAuthFlow.WaitToken(a.ctx, flowInfo)
if err != nil {
return nil, fmt.Errorf("wait for token: %w", err)
return nil, fmt.Errorf("waiting for browser login failed: %v", err)
}
return &tokenInfo, nil

View File

@@ -1,38 +0,0 @@
//go:build android
package android
import (
"fmt"
"github.com/netbirdio/netbird/client/internal/profilemanager"
)
type prefsStore interface {
Get(namespace string, v any) (bool, error)
Put(namespace string, v any) error
}
type profilePrefs struct {
prefs *profilemanager.Prefs
}
func newProfilePrefs(configDir, profileID string) (*profilePrefs, error) {
if configDir == "" || profileID == "" {
return nil, fmt.Errorf("profile prefs require a config dir and profile ID")
}
pm := NewProfileManager(configDir)
prefs, err := pm.serviceMgr.ProfilePrefs(profilemanager.ID(profileID), androidUsername)
if err != nil {
return nil, fmt.Errorf("resolve profile prefs: %w", err)
}
return &profilePrefs{prefs: prefs}, nil
}
func (p *profilePrefs) Get(namespace string, v any) (bool, error) {
return p.prefs.Get(namespace, v)
}
func (p *profilePrefs) Put(namespace string, v any) error {
return p.prefs.Put(namespace, v)
}

View File

@@ -1,649 +0,0 @@
//go:build android
package android
import (
"context"
"errors"
"fmt"
"io"
"net"
"strconv"
"strings"
"sync"
"time"
log "github.com/sirupsen/logrus"
gossh "golang.org/x/crypto/ssh"
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/profilemanager"
nbssh "github.com/netbirdio/netbird/client/ssh"
"github.com/netbirdio/netbird/client/ssh/detection"
)
const (
sshDialTimeout = 30 * time.Second
sshDetectionTimeout = 5 * time.Second
)
// PasswordRequiredMarker tells Java to prompt for a password and retry. It is
// a string because gomobile flattens errors to their message, so a sentinel
// value would not survive the binding.
const PasswordRequiredMarker = "netbird-ssh-password-required"
// HostKeyUnknownMarker tells Java to show the fingerprint and, on confirmation,
// retry with TrustHostKey set. The presented fingerprint is appended after the
// marker so the prompt can display it and the retry can guard against a key
// that changed between the two connects. Only regular (non-NetBird) servers
// reach this: NetBird peers verify against the registry.
const HostKeyUnknownMarker = "netbird-ssh-hostkey-unknown"
var (
errPasswordRequired = errors.New(PasswordRequiredMarker)
errClientClosed = errors.New("ssh client closed")
)
// errHostKeyUnknown carries the presented fingerprint so Connect can build the
// marker message the Java side parses.
type errHostKeyUnknown struct {
fingerprint string
}
func (e *errHostKeyUnknown) Error() string {
return HostKeyUnknownMarker + ":" + e.fingerprint
}
// SSHTerminalListener receives SSH session events. It is implemented in Java.
//
// All callbacks are invoked from goroutines and may run concurrently with each
// other; the implementation must be safe to call from any thread.
type SSHTerminalListener interface {
OnConnected()
OnData(data []byte)
OnClose(reason string)
OnError(message string)
}
// SSHClient is a NetBird-aware SSH client exposed to Java via gomobile.
//
// It dials through the running NetBird tunnel and runs a standard SSH session
// on top with PTY enabled. Host-key verification uses the NetBird-provided
// peer SSH host keys, identical to the desktop client.
type SSHClient struct {
nb *Client
mu sync.Mutex
listener SSHTerminalListener
urlOpener URLOpener
sshClient *gossh.Client
session *gossh.Session
stdin io.WriteCloser
closed bool
// gen identifies the current connection attempt. Connect and Close bump it,
// so an in-flight dial or a reader left over from a previous connection
// finds itself stale and stays silent instead of publishing OnConnected or
// OnClose for a connection the caller already abandoned.
gen uint64
dialCancel context.CancelFunc
// knownHostsConfigDir and knownHostsProfile locate the TOFU store for
// regular SSH servers in the profile's preferences. Java supplies them,
// since an overlay IP is a different host under a different profile. Empty
// until set: without them a regular server cannot be verified and Connect
// refuses one.
knownHostsConfigDir string
knownHostsProfile string
// trustHostKey carries the fingerprint the user confirmed on a previous
// attempt, so the retry accepts exactly that key and persists it.
trustHostKey string
}
// NewSSHClient creates a new SSH client bound to the running NetBird Client.
func NewSSHClient(c *Client) *SSHClient {
return &SSHClient{nb: c}
}
// SetListener registers the Java listener. Must be called before Connect to
// receive any events.
func (s *SSHClient) SetListener(l SSHTerminalListener) {
s.mu.Lock()
s.listener = l
s.mu.Unlock()
}
// SetURLOpener registers the Java URL opener used to display the device-code
// authorization page in a Custom Tabs window when the target peer requires
// JWT authentication. Must be set before Connect to be effective.
func (s *SSHClient) SetURLOpener(opener URLOpener) {
s.mu.Lock()
s.urlOpener = opener
s.mu.Unlock()
}
// SetKnownHostsStore points the TOFU host-key store at a profile's preferences.
// Must be set before connecting to a regular SSH server; without it such a
// server cannot be verified and Connect refuses one.
func (s *SSHClient) SetKnownHostsStore(configDir, profileID string) {
s.mu.Lock()
s.knownHostsConfigDir = configDir
s.knownHostsProfile = profileID
s.mu.Unlock()
}
// TrustHostKey records the fingerprint the user confirmed for a regular server,
// so the next Connect accepts that exact key and adds it to the known-hosts
// store. Passing a fingerprint that no longer matches makes the connect fail
// rather than trust a key that changed since the prompt.
func (s *SSHClient) TrustHostKey(fingerprint string) {
s.mu.Lock()
s.trustHostKey = fingerprint
s.mu.Unlock()
}
// Connect dials the SSH server through the NetBird tunnel and performs the
// SSH handshake. It auto-detects the server type via SSH banner inspection
// and selects the appropriate authentication path:
//
// - NetBird-SSH server requiring JWT: launches the OAuth 2.0 device-code
// flow, opens the verification URL through the registered URLOpener, and
// uses the resulting token as the SSH password. Host-key verification
// uses the NetBird peer registry.
// - NetBird-SSH server without JWT: authenticates with the NetBird SSH
// private key. Host-key verification uses the NetBird peer registry.
// - Regular SSH server (e.g. OpenSSH): authenticates with the NetBird key
// first (so a user-installed NetBird public key works), then falls back
// to the supplied password if non-empty. Host-key verification is
// trust-on-first-use against the per-profile known-hosts store.
//
// The password parameter is only consulted for regular SSH servers.
func (s *SSHClient) Connect(host string, port int, user, password string) error {
if port < 1 || port > 65535 {
return fmt.Errorf("invalid port: %d", port)
}
cfg, cfgPath, cc := s.nb.authSnapshot()
if cc == nil {
return errors.New("netbird client not running")
}
if cfg == nil {
return errors.New("netbird config not loaded")
}
engine := cc.Engine()
if engine == nil {
return errors.New("netbird engine not available")
}
s.mu.Lock()
s.gen++
gen := s.gen
s.mu.Unlock()
serverType := detectServerType(host, port)
log.Debugf("SSH server type: %s", serverType)
authMethods, hostKeyCallback, err := s.buildAuth(cfg, cfgPath, engine, serverType, password)
if err != nil {
return err
}
clientConfig := &gossh.ClientConfig{
User: user,
Auth: authMethods,
HostKeyCallback: hostKeyCallback,
Timeout: sshDialTimeout,
}
err = s.dialAndHandshake(gen, host, port, clientConfig)
// An unknown host key is a prompt, not a failure: return the marker intact
// (rootCause would unwrap it) so Java can show the fingerprint and retry.
var unknownHost *errHostKeyUnknown
if errors.As(err, &unknownHost) {
return errors.New(unknownHost.Error())
}
// A regular server may still accept a password, so let the caller ask for
// one instead of failing. NetBird servers never use a password, so a
// failure there is genuine.
if err != nil && serverType != detection.ServerTypeNetBirdJWT &&
serverType != detection.ServerTypeNetBirdNoJWT && isAuthFailure(err) &&
passwordCouldHelp(err, password != "") {
return errPasswordRequired
}
if err != nil {
return rootCause(err)
}
return nil
}
// StartSession requests a PTY and starts an interactive shell. Output from
// the session is forwarded to the listener via OnData.
func (s *SSHClient) StartSession(cols, rows int) error {
err := s.startSession(cols, rows)
if err != nil {
log.Infof("SSH: start session failed: %v", err)
return rootCause(err)
}
return nil
}
// Write sends data to the SSH session stdin.
func (s *SSHClient) Write(data []byte) error {
s.mu.Lock()
stdin := s.stdin
s.mu.Unlock()
if stdin == nil {
return errors.New("ssh session not started")
}
if _, err := stdin.Write(data); err != nil {
return fmt.Errorf("write stdin: %w", err)
}
return nil
}
// Resize updates the PTY window size.
func (s *SSHClient) Resize(cols, rows int) error {
s.mu.Lock()
session := s.session
s.mu.Unlock()
if session == nil {
return errors.New("ssh session not started")
}
return session.WindowChange(rows, cols)
}
// Reset makes a closed client usable for another Connect: Close leaves the
// one-shot guard set, and clearing it lets the same client back a reconnect.
func (s *SSHClient) Reset() {
s.mu.Lock()
defer s.mu.Unlock()
s.closed = false
}
// Close terminates the SSH session and underlying connection. Safe to call
// multiple times.
func (s *SSHClient) Close() error {
s.mu.Lock()
s.gen++
if s.dialCancel != nil {
s.dialCancel()
s.dialCancel = nil
}
sshClient := s.sshClient
session := s.session
stdin := s.stdin
s.sshClient = nil
s.session = nil
s.stdin = nil
notify := !s.closed
s.closed = true
listener := s.listener
s.mu.Unlock()
if stdin != nil {
if err := stdin.Close(); err != nil {
log.Debugf("ssh: stdin close: %v", err)
}
}
if session != nil {
if err := session.Close(); err != nil && !errors.Is(err, io.EOF) {
log.Debugf("ssh: session close: %v", err)
}
}
var firstErr error
if sshClient != nil {
if err := sshClient.Close(); err != nil {
firstErr = err
}
}
if notify && listener != nil {
listener.OnClose("closed by client")
}
return firstErr
}
func (s *SSHClient) startSession(cols, rows int) error {
log.Debugf("SSH: starting session %dx%d", cols, rows)
s.mu.Lock()
sshClient := s.sshClient
gen := s.gen
s.mu.Unlock()
if sshClient == nil {
return errors.New("ssh client not connected")
}
pty, err := nbssh.StartPTYSession(sshClient, cols, rows)
if err != nil {
return err
}
s.mu.Lock()
if gen != s.gen {
s.mu.Unlock()
closeQuiet(pty.Session, "stale session")
return errClientClosed
}
s.session = pty.Session
s.stdin = pty.Stdin
s.mu.Unlock()
readerDone := make(chan string, 2)
go func() { readerDone <- s.readLoop(pty.Stdout, "stdout") }()
go func() { readerDone <- s.readLoop(pty.Stderr, "stderr") }()
go func() {
reason := <-readerDone
if second := <-readerDone; reason == "" {
reason = second
}
s.notifyClose(gen, reason)
}()
log.Debug("SSH: session started, shell running")
return nil
}
func (s *SSHClient) buildAuth(cfg *profilemanager.Config, cfgPath string, engine *internal.Engine,
serverType detection.ServerType, password string) ([]gossh.AuthMethod, gossh.HostKeyCallback, error) {
switch serverType {
case detection.ServerTypeNetBirdJWT:
token, err := s.requestJWTToken(cfg, cfgPath)
if err != nil {
return nil, nil, fmt.Errorf("jwt: %w", err)
}
auths := []gossh.AuthMethod{gossh.Password(token)}
return auths, nbssh.CreateHostKeyCallback(nbssh.PeerKeyLookup(engine.GetPeerSSHKey)), nil
case detection.ServerTypeNetBirdNoJWT:
if cfg.SSHKey == "" {
return nil, nil, errors.New("no NetBird SSH key available")
}
signer, err := gossh.ParsePrivateKey([]byte(cfg.SSHKey))
if err != nil {
return nil, nil, fmt.Errorf("parse netbird ssh key: %w", err)
}
auths := []gossh.AuthMethod{gossh.PublicKeys(signer)}
return auths, nbssh.CreateHostKeyCallback(nbssh.PeerKeyLookup(engine.GetPeerSSHKey)), nil
case detection.ServerTypeRegular:
var auths []gossh.AuthMethod
if cfg.SSHKey != "" {
if signer, err := gossh.ParsePrivateKey([]byte(cfg.SSHKey)); err == nil {
auths = append(auths, gossh.PublicKeys(signer))
} else {
log.Debugf("ssh: parse netbird key for regular auth: %v", err)
}
}
if password != "" {
pw := password
auths = append(auths, gossh.Password(pw))
auths = append(auths, gossh.KeyboardInteractive(func(_, _ string, questions []string, _ []bool) ([]string, error) {
answers := make([]string, len(questions))
for i := range questions {
answers[i] = pw
}
return answers, nil
}))
}
if len(auths) == 0 {
// Nothing to offer at all: ask for a password rather than failing,
// so the caller can retry once the user supplies one.
return nil, nil, errPasswordRequired
}
callback, err := s.tofuHostKeyCallback()
if err != nil {
return nil, nil, err
}
return auths, callback, nil
default:
return nil, nil, fmt.Errorf("unsupported SSH server type: %v", serverType)
}
}
// tofuHostKeyCallback verifies a regular server's host key against the
// per-profile known-hosts store. An unknown host returns errHostKeyUnknown so
// Java can show the fingerprint and, once confirmed, retry with the key
// trusted; a changed key is rejected outright, as OpenSSH does. When the user
// has confirmed a fingerprint, the callback accepts exactly that key and
// appends it to the store.
func (s *SSHClient) tofuHostKeyCallback() (gossh.HostKeyCallback, error) {
s.mu.Lock()
configDir := s.knownHostsConfigDir
profileID := s.knownHostsProfile
trusted := s.trustHostKey
s.mu.Unlock()
if configDir == "" || profileID == "" {
return nil, errors.New("no known-hosts store configured for regular SSH")
}
store, err := openKnownHostsStore(configDir, profileID)
if err != nil {
return nil, fmt.Errorf("load known-hosts store: %w", err)
}
return func(hostname string, remote net.Addr, key gossh.PublicKey) error {
verdict, err := store.verify(hostname, remote, key)
if err != nil {
return err
}
if verdict == hostKeyMatched {
return nil
}
if verdict == hostKeyChanged {
return fmt.Errorf("SSH host key changed for %s (possible attack)", hostname)
}
fingerprint := gossh.FingerprintSHA256(key)
if trusted == "" {
return &errHostKeyUnknown{fingerprint: fingerprint}
}
if trusted != fingerprint {
return fmt.Errorf("SSH host key changed since it was confirmed for %s", hostname)
}
if err := store.append(hostname, remote, key); err != nil {
return fmt.Errorf("persist trusted host key: %w", err)
}
// The confirmation is spent: now that the key is stored, a later
// reconnect must verify against the file, not re-accept this fingerprint.
s.mu.Lock()
s.trustHostKey = ""
s.mu.Unlock()
return nil
}, nil
}
func (s *SSHClient) requestJWTToken(cfg *profilemanager.Config, cfgPath string) (string, error) {
s.mu.Lock()
urlOpener := s.urlOpener
s.mu.Unlock()
if urlOpener == nil {
return "", errors.New("URL opener not configured for JWT auth")
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
defer cancel()
flow, err := auth.NewOAuthFlow(ctx, cfg, false, true, profileLoginHint(cfgPath))
if err != nil {
return "", fmt.Errorf("create oauth flow: %w", err)
}
// The status callback covers the browser round-trip, which would
// otherwise leave the terminal blank.
tokenInfo, err := runOAuthFlow(ctx, flow, urlOpener, func() {
s.notifyStatus("Waiting for browser authentication...")
})
if err != nil {
return "", err
}
token := tokenInfo.GetTokenToUse()
if token == "" {
return "", errors.New("empty token returned by IdP")
}
// Tells the client the browser round-trip is over so it can dismiss the
// surface it opened, the same way the login and session-extend flows do.
// Without it the Custom Tab stays in front of the terminal even though the
// token has already been collected.
urlOpener.OnLoginSuccess()
return token, nil
}
func (s *SSHClient) dialAndHandshake(gen uint64, host string, port int, clientConfig *gossh.ClientConfig) error {
addr := net.JoinHostPort(host, strconv.Itoa(port))
ctx, cancel := context.WithTimeout(context.Background(), sshDialTimeout)
defer cancel()
s.mu.Lock()
if gen != s.gen {
s.mu.Unlock()
return errClientClosed
}
s.dialCancel = cancel
s.mu.Unlock()
var dialer net.Dialer
conn, err := dialer.DialContext(ctx, "tcp", addr)
if err != nil {
return fmt.Errorf("dial %s: %w", addr, err)
}
client, err := nbssh.Handshake(ctx, conn, addr, clientConfig)
if err != nil {
return err
}
s.mu.Lock()
if gen != s.gen {
s.mu.Unlock()
closeQuiet(client, "stale ssh client")
return errClientClosed
}
s.sshClient = client
listener := s.listener
s.mu.Unlock()
if listener != nil {
listener.OnConnected()
}
return nil
}
func (s *SSHClient) readLoop(r io.Reader, name string) string {
buf := make([]byte, 4096)
for {
n, err := r.Read(buf)
if n > 0 {
s.mu.Lock()
listener := s.listener
s.mu.Unlock()
if listener != nil {
chunk := make([]byte, n)
copy(chunk, buf[:n])
listener.OnData(chunk)
}
}
if err != nil {
// EOF is a normal shell exit, so report it without a reason.
if errors.Is(err, io.EOF) {
return ""
}
log.Debugf("ssh %s read: %v", name, err)
return rootCause(err).Error()
}
}
}
// notifyStatus writes a progress line to the terminal through the normal
// output path, so long steps are visible while nothing else is arriving.
func (s *SSHClient) notifyStatus(text string) {
s.mu.Lock()
listener := s.listener
s.mu.Unlock()
if listener != nil {
listener.OnData([]byte("\r\n\x1b[33m" + text + "\x1b[0m\r\n"))
}
}
func (s *SSHClient) notifyClose(gen uint64, reason string) {
s.mu.Lock()
if gen != s.gen || s.closed {
s.mu.Unlock()
return
}
s.closed = true
listener := s.listener
s.mu.Unlock()
if listener != nil {
listener.OnClose(reason)
}
}
func closeQuiet(c io.Closer, label string) {
if c == nil {
return
}
if err := c.Close(); err != nil && !errors.Is(err, io.EOF) {
log.Debugf("ssh: close %s: %v", label, err)
}
}
func detectServerType(host string, port int) detection.ServerType {
ctx, cancel := context.WithTimeout(context.Background(), sshDetectionTimeout)
defer cancel()
dialer := &net.Dialer{}
serverType, err := detection.DetectSSHServerType(ctx, dialer, host, port)
if err != nil {
log.Debugf("ssh: server detection failed: %v (assuming regular SSH)", err)
return detection.ServerTypeRegular
}
return serverType
}
// rootCause returns the innermost error of a %w chain, so the terminal shows
// "i/o timeout" rather than every layer that added context on the way up.
func rootCause(err error) error {
for {
// A joined error has no single root, so keep it as-is.
if _, ok := err.(interface{ Unwrap() []error }); ok {
return err
}
next := errors.Unwrap(err)
if next == nil {
return err
}
err = next
}
}
// isAuthFailure distinguishes credential rejection from dial, timeout and
// host-key errors, which retrying with a password would not fix.
func isAuthFailure(err error) bool {
if errors.Is(err, errPasswordRequired) {
return true
}
var partial *gossh.PartialSuccessError
if errors.As(err, &partial) {
return true
}
return strings.Contains(err.Error(), "unable to authenticate")
}
// passwordCouldHelp reports whether prompting for a password again can change
// the outcome. gossh lists a method under "attempted methods" only when the
// server offered it, so a supplied password that was never attempted means the
// server does not accept passwords and the real error should surface instead.
func passwordCouldHelp(err error, passwordOffered bool) bool {
if !passwordOffered {
return true
}
msg := err.Error()
return strings.Contains(msg, "password") || strings.Contains(msg, "keyboard-interactive")
}

View File

@@ -1,168 +0,0 @@
//go:build android
package android
import (
"bytes"
"net"
"strconv"
"strings"
"sync"
gossh "golang.org/x/crypto/ssh"
"golang.org/x/crypto/ssh/knownhosts"
)
const knownHostsNamespace = "ssh"
const (
hostKeyUnknown hostKeyVerdict = iota
hostKeyMatched
hostKeyChanged
)
var knownHostsMu sync.Mutex
type hostKeyVerdict uint8
type knownHostsSection struct {
KnownHosts []string `json:"knownHosts"`
}
type knownHostsStore struct {
prefs prefsStore
}
// RemoveKnownHost deletes every known-hosts entry for host:port from the
// profile's store, so a host trusted for a session that is being deleted does
// not linger. Java calls this only once no session targets that host, so a
// shared host stays trusted. A missing entry is not an error: the goal state
// is "absent".
func RemoveKnownHost(configDir, profileID, host string, port int) error {
store, err := openKnownHostsStore(configDir, profileID)
if err != nil {
return err
}
return store.removeHost(host, port)
}
func openKnownHostsStore(configDir, profileID string) (*knownHostsStore, error) {
prefs, err := newProfilePrefs(configDir, profileID)
if err != nil {
return nil, err
}
return &knownHostsStore{prefs: prefs}, nil
}
func (st *knownHostsStore) verify(hostname string, remote net.Addr, key gossh.PublicKey) (hostKeyVerdict, error) {
lines, err := st.lines()
if err != nil {
return hostKeyUnknown, err
}
targets := knownHostsTargets(hostname, remote)
verdict := hostKeyUnknown
for _, line := range lines {
pubKey, ok := knownHostsLineKey(line, targets)
if !ok {
continue
}
if pubKey.Type() == key.Type() && bytes.Equal(pubKey.Marshal(), key.Marshal()) {
return hostKeyMatched, nil
}
verdict = hostKeyChanged
}
return verdict, nil
}
func (st *knownHostsStore) append(hostname string, remote net.Addr, key gossh.PublicKey) error {
line := knownhosts.Line(knownHostsTargets(hostname, remote), key)
knownHostsMu.Lock()
defer knownHostsMu.Unlock()
lines, err := st.lines()
if err != nil {
return err
}
return st.prefs.Put(knownHostsNamespace, knownHostsSection{KnownHosts: append(lines, line)})
}
func (st *knownHostsStore) removeHost(host string, port int) error {
target := knownhosts.Normalize(net.JoinHostPort(host, strconv.Itoa(port)))
knownHostsMu.Lock()
defer knownHostsMu.Unlock()
lines, err := st.lines()
if err != nil {
return err
}
kept := make([]string, 0, len(lines))
for _, line := range lines {
if knownHostsLineMatches(line, target) {
continue
}
kept = append(kept, line)
}
if len(kept) == len(lines) {
return nil
}
return st.prefs.Put(knownHostsNamespace, knownHostsSection{KnownHosts: kept})
}
func (st *knownHostsStore) lines() ([]string, error) {
var section knownHostsSection
if _, err := st.prefs.Get(knownHostsNamespace, &section); err != nil {
return nil, err
}
return section.KnownHosts, nil
}
func knownHostsTargets(hostname string, remote net.Addr) []string {
targets := []string{knownhosts.Normalize(hostname)}
if remote != nil {
if normalized := knownhosts.Normalize(remote.String()); normalized != targets[0] {
targets = append(targets, normalized)
}
}
return targets
}
func knownHostsLineKey(line string, targets []string) (gossh.PublicKey, bool) {
trimmed := strings.TrimSpace(line)
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
return nil, false
}
_, hosts, pubKey, _, _, err := gossh.ParseKnownHosts([]byte(trimmed))
if err != nil {
return nil, false
}
for _, host := range hosts {
for _, target := range targets {
if host == target {
return pubKey, true
}
}
}
return nil, false
}
// knownHostsLineMatches reports whether a known-hosts line's address list
// contains the normalized target. Comment and blank lines never match.
func knownHostsLineMatches(line, target string) bool {
trimmed := strings.TrimSpace(line)
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
return false
}
fields := strings.Fields(trimmed)
if len(fields) == 0 {
return false
}
for _, addr := range strings.Split(fields[0], ",") {
if addr == target {
return true
}
}
return false
}

View File

@@ -1,104 +0,0 @@
//go:build android
package android
const (
sshSessionsNamespace = "ssh-sessions"
maxStoredSSHSessions = 50
)
type sshSessionRecord struct {
ID string `json:"id"`
Host string `json:"host"`
Port int `json:"port"`
User string `json:"user"`
}
type sshSessionsSection struct {
Sessions []sshSessionRecord `json:"sessions"`
}
// SSHSessionEntry is one stored SSH session, without any credential.
type SSHSessionEntry struct {
ID string
Host string
Port int
User string
}
// SSHSessionArray wraps stored SSH sessions for gomobile compatibility.
type SSHSessionArray struct {
items []*SSHSessionEntry
}
// NewSSHSessionArray creates an empty session array to fill via Add.
func NewSSHSessionArray() *SSHSessionArray {
return &SSHSessionArray{}
}
// Add appends a session entry, oldest first.
func (a *SSHSessionArray) Add(id, host string, port int, user string) {
a.items = append(a.items, &SSHSessionEntry{ID: id, Host: host, Port: port, User: user})
}
// Length returns the number of entries.
func (a *SSHSessionArray) Length() int {
return len(a.items)
}
// Get returns the entry at index i, or nil when out of range.
func (a *SSHSessionArray) Get(i int) *SSHSessionEntry {
if i < 0 || i >= len(a.items) {
return nil
}
return a.items[i]
}
// SSHSessionStore reads and writes a profile's stored SSH sessions.
type SSHSessionStore struct {
prefs prefsStore
}
// NewSSHSessionStore opens the session store of the given profile.
func NewSSHSessionStore(configDir, profileID string) (*SSHSessionStore, error) {
prefs, err := newProfilePrefs(configDir, profileID)
if err != nil {
return nil, err
}
return &SSHSessionStore{prefs: prefs}, nil
}
// Load returns the stored sessions, oldest first.
func (s *SSHSessionStore) Load() (*SSHSessionArray, error) {
var section sshSessionsSection
if _, err := s.prefs.Get(sshSessionsNamespace, &section); err != nil {
return nil, err
}
out := NewSSHSessionArray()
for _, record := range section.Sessions {
if record.ID == "" || record.Host == "" {
continue
}
out.Add(record.ID, record.Host, record.Port, record.User)
}
return out, nil
}
// Save replaces the stored sessions, keeping only the newest entries when the
// list exceeds the storage cap.
func (s *SSHSessionStore) Save(sessions *SSHSessionArray) error {
var items []*SSHSessionEntry
if sessions != nil {
items = sessions.items
}
if len(items) > maxStoredSSHSessions {
items = items[len(items)-maxStoredSSHSessions:]
}
records := make([]sshSessionRecord, 0, len(items))
for _, item := range items {
records = append(records, sshSessionRecord{ID: item.ID, Host: item.Host, Port: item.Port, User: item.User})
}
return s.prefs.Put(sshSessionsNamespace, sshSessionsSection{Sessions: records})
}

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, nil)
networkMapController := controller.NewController(ctx, store, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.cloud", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersmanager), config)
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,7 +21,7 @@ import (
"github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
nbssh "github.com/netbirdio/netbird/client/ssh"
sshcommon "github.com/netbirdio/netbird/client/ssh"
"github.com/netbirdio/netbird/client/system"
"github.com/netbirdio/netbird/shared/management/domain"
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
@@ -521,7 +521,12 @@ func (c *Client) VerifySSHHostKey(peerAddress string, key []byte) error {
return err
}
return nbssh.PeerKeyLookup(engine.GetPeerSSHKey).VerifySSHHostKey(peerAddress, key)
storedKey, found := engine.GetPeerSSHKey(peerAddress)
if !found {
return sshcommon.ErrPeerNotFound
}
return sshcommon.VerifyHostKey(storedKey, key, peerAddress)
}
// SetPerformance retunes a running Client. Only PreallocatedBuffersPerPool

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, nil)
networkMapController := controller.NewController(context.Background(), testStore, metrics, updateManager, requestBuffer, mgmt.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(testStore, peersManager), cfg)
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

@@ -16,47 +16,28 @@ import (
"google.golang.org/grpc"
nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/netsweep"
)
func WithCustomDialer(_ bool, _ string) grpc.DialOption {
return grpc.WithContextDialer(dialContext)
}
// WithSweeper dials like WithCustomDialer but registers connections and
// dials with the sweeper. Append it after WithCustomDialer: gRPC applies
// dial options in order, so the later context dialer wins.
func WithSweeper(sweeper *netsweep.Sweeper) grpc.DialOption {
return grpc.WithContextDialer(func(ctx context.Context, addr string) (net.Conn, error) {
dial := sweeper.StartDial(ctx)
defer dial.Release()
if runtime.GOOS == "linux" {
currentUser, err := user.Current()
if err != nil {
return nil, status.Errorf(codes.FailedPrecondition, "failed to get current user: %v", err)
}
conn, err := dialContext(dial.Ctx(), addr)
if err != nil {
return nil, err
// the custom dialer requires root permissions which are not required for use cases run as non-root
if currentUser.Uid != "0" {
log.Debug("Not running as root, using standard dialer")
dialer := &net.Dialer{}
return dialer.DialContext(ctx, "tcp", addr)
}
}
return dial.WrapConn(conn)
conn, err := nbnet.NewDialer().DialContext(ctx, "tcp", addr)
if err != nil {
return nil, fmt.Errorf("nbnet.NewDialer().DialContext: %w", err)
}
return conn, nil
})
}
func dialContext(ctx context.Context, addr string) (net.Conn, error) {
if runtime.GOOS == "linux" {
currentUser, err := user.Current()
if err != nil {
return nil, status.Errorf(codes.FailedPrecondition, "failed to get current user: %v", err)
}
// the custom dialer requires root permissions which are not required for use cases run as non-root
if currentUser.Uid != "0" {
log.Debug("Not running as root, using standard dialer")
dialer := &net.Dialer{}
return dialer.DialContext(ctx, "tcp", addr)
}
}
conn, err := nbnet.NewDialer().DialContext(ctx, "tcp", addr)
if err != nil {
return nil, fmt.Errorf("nbnet.NewDialer().DialContext: %w", err)
}
return conn, nil
}

View File

@@ -3,7 +3,6 @@ package grpc
import (
"google.golang.org/grpc"
"github.com/netbirdio/netbird/client/netsweep"
"github.com/netbirdio/netbird/util/wsproxy/client"
)
@@ -12,8 +11,3 @@ import (
func WithCustomDialer(tlsEnabled bool, component string) grpc.DialOption {
return client.WithWebSocketDialer(tlsEnabled, component)
}
// WithSweeper is a no-op on WASM/JS: there is no network change signal.
func WithSweeper(_ *netsweep.Sweeper) grpc.DialOption {
return grpc.EmptyDialOption{}
}

View File

@@ -1,49 +0,0 @@
package grpc
import (
"context"
"errors"
"time"
"github.com/cenkalti/backoff/v4"
"github.com/netbirdio/netbird/client/netstate"
)
// Retry mirrors backoff.Retry, but the sleep between attempts also wakes on
// OS network availability transitions: an operation cut down by a network
// change retries the moment the network settles instead of sleeping through
// the recovery. A nil netState never fires, leaving plain backoff.Retry
// behavior.
func Retry(ctx context.Context, operation backoff.Operation, bo backoff.BackOff, netState *netstate.State) error {
bo.Reset()
for {
err := operation()
if err == nil {
return nil
}
var permanent *backoff.PermanentError
if errors.As(err, &permanent) {
return permanent.Err
}
next := bo.NextBackOff()
if next == backoff.Stop {
if cerr := ctx.Err(); cerr != nil {
return cerr
}
return err
}
timer := time.NewTimer(next)
select {
case <-timer.C:
case <-netState.Changed():
timer.Stop()
case <-ctx.Done():
timer.Stop()
return ctx.Err()
}
}
}

View File

@@ -1,91 +0,0 @@
package grpc
import (
"context"
"errors"
"testing"
"time"
"github.com/cenkalti/backoff/v4"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/netstate"
)
func TestRetryWakesOnNetworkChange(t *testing.T) {
ns := netstate.New()
attempts := 0
operation := func() error {
attempts++
if attempts == 1 {
return errors.New("cut by network change")
}
return nil
}
go func() {
time.Sleep(20 * time.Millisecond)
ns.Set(false)
}()
start := time.Now()
err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Minute), ns)
require.NoError(t, err)
assert.Equal(t, 2, attempts, "network change must cause one immediate retry")
assert.Less(t, time.Since(start), time.Second, "the transition must cut the minute-long sleep short")
}
func TestRetryPermanentError(t *testing.T) {
sentinel := errors.New("permission denied")
operation := func() error {
return backoff.Permanent(sentinel)
}
err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Millisecond), nil)
assert.ErrorIs(t, err, sentinel, "permanent errors must stop retries")
}
func TestRetryNilNetState(t *testing.T) {
attempts := 0
operation := func() error {
attempts++
if attempts < 3 {
return errors.New("transient")
}
return nil
}
err := Retry(context.Background(), operation, backoff.NewConstantBackOff(time.Millisecond), nil)
require.NoError(t, err)
assert.Equal(t, 3, attempts, "nil network state must preserve timed retries")
}
func TestRetryStops(t *testing.T) {
failure := errors.New("still failing")
operation := func() error {
return failure
}
err := Retry(context.Background(), operation, &backoff.StopBackOff{}, nil)
assert.ErrorIs(t, err, failure, "stop backoff must return the operation error")
}
func TestRetryCtxCancelDuringSleep(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
operation := func() error {
return errors.New("failing")
}
go func() {
time.Sleep(20 * time.Millisecond)
cancel()
}()
start := time.Now()
err := Retry(ctx, operation, backoff.NewConstantBackOff(time.Minute), netstate.New())
assert.ErrorIs(t, err, context.Canceled, "context cancellation must stop the retry loop")
assert.Less(t, time.Since(start), time.Second, "context cancellation must interrupt backoff sleep")
}

View File

@@ -138,37 +138,26 @@ func (a *Auth) IsSSOSupported(ctx context.Context) (bool, error) {
// GetOAuthFlow returns an OAuth flow (PKCE or Device) using the existing management connection
// This avoids creating a new connection to the management server
func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool, hint string) (OAuthFlow, error) {
func (a *Auth) GetOAuthFlow(ctx context.Context, forceDeviceAuth bool) (OAuthFlow, error) {
var flow OAuthFlow
var err error
err := a.withRetry(ctx, func(client *mgm.GrpcClient) error {
err = a.withRetry(ctx, func(client *mgm.GrpcClient) error {
if forceDeviceAuth {
deviceFlow, err := a.getDeviceFlow(client)
if err != nil {
return err
}
deviceFlow.SetLoginHint(hint)
flow = deviceFlow
return nil
flow, err = a.getDeviceFlow(client)
return err
}
// Try PKCE flow first
pkceFlow, err := a.getPKCEFlow(client)
flow, err = a.getPKCEFlow(client)
if err != nil {
// If PKCE not supported, try Device flow
if s, ok := status.FromError(err); ok && (s.Code() == codes.NotFound || s.Code() == codes.Unimplemented) {
deviceFlow, err := a.getDeviceFlow(client)
if err != nil {
return err
}
deviceFlow.SetLoginHint(hint)
flow = deviceFlow
return nil
flow, err = a.getDeviceFlow(client)
return err
}
return err
}
pkceFlow.SetLoginHint(hint)
flow = pkceFlow
return nil
})

View File

@@ -97,7 +97,9 @@ func authenticateWithPKCEFlow(ctx context.Context, config *profilemanager.Config
return nil, fmt.Errorf("getting pkce authorization flow info failed with error: %v", err)
}
pkceFlowInfo.SetLoginHint(hint)
if hint != "" {
pkceFlowInfo.SetLoginHint(hint)
}
return pkceFlowInfo, nil
}
@@ -125,7 +127,9 @@ func authenticateWithDeviceCodeFlow(ctx context.Context, config *profilemanager.
}
}
deviceFlowInfo.SetLoginHint(hint)
if hint != "" {
deviceFlowInfo.SetLoginHint(hint)
}
return deviceFlowInfo, nil
}

View File

@@ -38,8 +38,6 @@ import (
"github.com/netbirdio/netbird/client/internal/updater"
"github.com/netbirdio/netbird/client/internal/updater/installer"
nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/netstate"
"github.com/netbirdio/netbird/client/netsweep"
cProto "github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/ssh"
sshconfig "github.com/netbirdio/netbird/client/ssh/config"
@@ -72,42 +70,18 @@ type ConnectClient struct {
updateManager *updater.Manager
persistSyncResponse bool
// netState gates every reconnection loop on OS-reported network
// availability. Nil (the default) disables gating; mobile platforms
// inject it via WithNetworkState.
netState *netstate.State
// sweeper cuts the management, signal and relay connections on network
// change; nil disables it.
sweeper *netsweep.Sweeper
}
// ConnectClientOption configures optional ConnectClient behavior.
type ConnectClientOption func(*ConnectClient)
// WithNetworkState injects the OS network availability state that gates every
// reconnection loop; without it gating is disabled.
func WithNetworkState(netState *netstate.State) ConnectClientOption {
return func(c *ConnectClient) { c.netState = netState }
}
// WithSweeper injects the network change sweeper.
func WithSweeper(sweeper *netsweep.Sweeper) ConnectClientOption {
return func(c *ConnectClient) { c.sweeper = sweeper }
}
func NewConnectClient(
ctx context.Context,
config *profilemanager.Config,
statusRecorder *peer.Status,
opts ...ConnectClientOption,
) *ConnectClient {
// Derive the run context here so Stop owns the cancel that unblocks the run
// loop. runCancel is set once at construction, so Stop can call it without
// racing the run loop's startup. Callers therefore need not cancel before Stop.
runCtx, runCancel := context.WithCancel(ctx)
c := &ConnectClient{
return &ConnectClient{
ctx: runCtx,
runCancel: runCancel,
runExited: make(chan struct{}),
@@ -115,10 +89,6 @@ func NewConnectClient(
statusRecorder: statusRecorder,
engineMutex: sync.Mutex{},
}
for _, opt := range opts {
opt(c)
}
return c
}
func (c *ConnectClient) SetUpdateManager(um *updater.Manager) {
@@ -304,13 +274,6 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
return nil
}
// suspend connection attempts while the OS reports no usable network
if waited, err := c.netState.Wait(c.ctx); err != nil {
return nil
} else if waited {
backOff.Reset()
}
state.Set(StatusConnecting)
engineCtx, cancel := context.WithCancel(c.ctx)
@@ -322,8 +285,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
}()
log.Debugf("connecting to the Management service %s", c.config.ManagementURL.Host)
mgmClient, err := mgm.NewClient(engineCtx, c.config.ManagementURL.Host, myPrivateKey, mgmTlsEnabled,
mgm.WithNetworkState(c.netState), mgm.WithSweeper(c.sweeper))
mgmClient, err := mgm.NewClient(engineCtx, c.config.ManagementURL.Host, myPrivateKey, mgmTlsEnabled)
if err != nil {
// On daemon shutdown / Down() the parent context is cancelled
// and the dial fails with "context canceled". Wrapping that
@@ -398,7 +360,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
}()
// with the global Netbird config in hand connect (just a connection, no stream yet) Signal
signalClient, err := connectToSignal(engineCtx, loginResp.GetNetbirdConfig(), myPrivateKey, c.netState, c.sweeper)
signalClient, err := connectToSignal(engineCtx, loginResp.GetNetbirdConfig(), myPrivateKey)
if err != nil {
log.Error(err)
return wrapErr(err)
@@ -434,8 +396,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
engineConfig.StateDir = filepath.Dir(path)
}
relayManager := relayClient.NewManager(engineCtx, relayURLs, myPrivateKey.PublicKey().String(), engineConfig.MTU,
relayClient.WithNetworkState(c.netState), relayClient.WithSweeper(c.sweeper))
relayManager := relayClient.NewManager(engineCtx, relayURLs, myPrivateKey.PublicKey().String(), engineConfig.MTU)
c.statusRecorder.SetRelayMgr(relayManager)
if len(relayURLs) > 0 {
if token != nil {
@@ -463,7 +424,6 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
UpdateManager: c.updateManager,
ClientMetrics: c.clientMetrics,
MetricsCtx: c.ctx,
NetState: c.netState,
}, mobileDependency)
engine.SetSyncResponsePersistence(c.persistSyncResponse)
c.engine = engine
@@ -520,16 +480,6 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
// status stream stuck at Connecting.
err = backoff.Retry(operation, backoff.WithContext(backOff, c.ctx))
if err != nil {
// Once the client context is cancelled backoff.WithContext surfaces the
// bare context error, and any attempt torn down mid-flight reports the
// same. That cancellation is the caller asking us to stop (Stop, Down or
// an engine restart), so exit cleanly instead of handing back a failure
// the caller would have to distinguish from a real one.
if c.ctx.Err() != nil && errors.Is(err, context.Canceled) {
log.Info("exiting client retry loop, context cancelled")
return nil
}
log.Debugf("exiting client retry loop due to unrecoverable error: %s", err)
if s, ok := gstatus.FromError(err); ok && (s.Code() == codes.PermissionDenied) {
state.Set(StatusNeedsLogin)
@@ -723,7 +673,7 @@ func selectMTU(localMTU uint16, peerMTU int32) uint16 {
}
// connectToSignal creates Signal Service client and established a connection
func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourPrivateKey wgtypes.Key, netState *netstate.State, sweeper *netsweep.Sweeper) (*signal.GrpcClient, error) {
func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourPrivateKey wgtypes.Key) (*signal.GrpcClient, error) {
var sigTLSEnabled bool
if wtConfig.Signal.Protocol == mgmProto.HostConfig_HTTPS {
sigTLSEnabled = true
@@ -731,8 +681,7 @@ func connectToSignal(ctx context.Context, wtConfig *mgmProto.NetbirdConfig, ourP
sigTLSEnabled = false
}
signalClient, err := signal.NewClient(ctx, wtConfig.Signal.Uri, ourPrivateKey, sigTLSEnabled,
signal.WithNetworkState(netState), signal.WithSweeper(sweeper))
signalClient, err := signal.NewClient(ctx, wtConfig.Signal.Uri, ourPrivateKey, sigTLSEnabled)
if err != nil {
log.Errorf("error while connecting to the Signal Exchange Service %s: %s", wtConfig.Signal.Uri, err)
return nil, gstatus.Errorf(codes.FailedPrecondition, "failed connecting to Signal Service : %s", err)

View File

@@ -59,7 +59,6 @@ import (
"github.com/netbirdio/netbird/client/internal/syncstore"
"github.com/netbirdio/netbird/client/internal/updater"
"github.com/netbirdio/netbird/client/jobexec"
"github.com/netbirdio/netbird/client/netstate"
cProto "github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/system"
nbdns "github.com/netbirdio/netbird/dns"
@@ -182,9 +181,6 @@ type EngineServices struct {
UpdateManager *updater.Manager
ClientMetrics *metrics.ClientMetrics
MetricsCtx context.Context
// NetState gates the reconnection loops on OS-reported network
// availability; nil disables gating.
NetState *netstate.State
}
// Engine is a mechanism responsible for reacting on Signal and Management stream events and managing connections to the remote peers.
@@ -208,10 +204,6 @@ type Engine struct {
config *EngineConfig
mobileDep MobileDependency
// netState gates the peer reconnection guards on OS-reported network
// availability; nil disables gating.
netState *netstate.State
// STUNs is a list of STUN servers used by ICE
STUNs []*stun.URI
// TURNs is a list of STUN servers used by ICE
@@ -345,7 +337,6 @@ func NewEngine(
syncMsgMux: &sync.Mutex{},
config: config,
mobileDep: mobileDep,
netState: services.NetState,
STUNs: []*stun.URI{},
TURNs: []*stun.URI{},
networkSerial: 0,
@@ -1902,8 +1893,7 @@ func (e *Engine) createPeerConn(pubKey string, allowedIPs []netip.Prefix, agentV
Addr: e.getRosenpassAddr(),
PermissiveMode: e.config.RosenpassPermissive,
},
ICEConfig: e.createICEConfig(),
NetworkState: e.netState,
ICEConfig: e.createICEConfig(),
}
serviceDependencies := peer.ServiceDependencies{

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, nil)
networkMapController := controller.NewController(context.Background(), store, metrics, updateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config)
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

@@ -26,7 +26,6 @@ import (
"github.com/netbirdio/netbird/client/internal/portforward"
"github.com/netbirdio/netbird/client/internal/rosenpass"
"github.com/netbirdio/netbird/client/internal/stdnet"
"github.com/netbirdio/netbird/client/netstate"
"github.com/netbirdio/netbird/route"
relayClient "github.com/netbirdio/netbird/shared/relay/client"
)
@@ -94,10 +93,6 @@ type ConnConfig struct {
// ICEConfig ICE protocol configuration
ICEConfig icemaker.Config
// NetworkState gates the reconnection guard on OS-reported network
// availability; nil disables gating.
NetworkState *netstate.State
}
type Conn struct {
@@ -259,7 +254,7 @@ func (conn *Conn) open(engineCtx context.Context, firstPacket []byte) error {
conn.handshaker.AddICEListener(conn.workerICE.OnNewOffer)
}
conn.guard = guard.NewGuard(conn.Log, conn.isConnectedOnAllWay, conn.config.Timeout, conn.srWatcher, conn.config.NetworkState)
conn.guard = guard.NewGuard(conn.Log, conn.isConnectedOnAllWay, conn.config.Timeout, conn.srWatcher)
conn.wg.Add(1)
go func() {

View File

@@ -6,8 +6,6 @@ import (
"github.com/cenkalti/backoff/v4"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/netstate"
)
// ConnStatus represents the connection state as seen by the guard.
@@ -33,26 +31,20 @@ type connStatusFunc func() ConnStatus
// - Relayed connection disconnected
// - ICE candidate changes
type Guard struct {
log *log.Entry
isConnectedOnAllWay connStatusFunc
timeout time.Duration
srWatcher *SRWatcher
// netState gates reconnect attempts on OS-reported network availability;
// nil disables gating.
netState *netstate.State
log *log.Entry
isConnectedOnAllWay connStatusFunc
timeout time.Duration
srWatcher *SRWatcher
relayedConnDisconnected chan struct{}
iCEConnDisconnected chan struct{}
}
// NewGuard creates a reconnection guard for a peer connection. A nil netState
// disables network availability gating.
func NewGuard(log *log.Entry, isConnectedFn connStatusFunc, timeout time.Duration, srWatcher *SRWatcher, netState *netstate.State) *Guard {
func NewGuard(log *log.Entry, isConnectedFn connStatusFunc, timeout time.Duration, srWatcher *SRWatcher) *Guard {
return &Guard{
log: log,
isConnectedOnAllWay: isConnectedFn,
timeout: timeout,
srWatcher: srWatcher,
netState: netState,
relayedConnDisconnected: make(chan struct{}, 1),
iCEConnDisconnected: make(chan struct{}, 1),
}
@@ -104,16 +96,9 @@ func (g *Guard) reconnectLoopWithRetry(ctx context.Context, callback func()) {
iceState := &iceRetryState{log: g.log}
defer iceState.reset()
netChanged := g.netState.Changed()
for {
select {
case <-tickerChannel:
// skip attempts while the OS reports no usable network; the
// netChanged case below resumes the loop once it returns
if !g.netState.IsOnline() {
continue
}
switch g.isConnectedOnAllWay() {
case ConnStatusConnected:
// all good, nothing to do
@@ -150,23 +135,6 @@ func (g *Guard) reconnectLoopWithRetry(ctx context.Context, callback func()) {
tickerChannel = ticker.C
iceState.reset()
case <-netChanged:
// Re-arm for the next transition before acting on this one.
netChanged = g.netState.Changed()
if !g.netState.IsOnline() {
continue
}
// Ticks skipped while offline drove the backoff towards its
// maximum without ever attempting, and left the ICE budget
// frozen — possibly in hourly mode. Recover on our own so the
// peer does not depend on a signal or relay event that never
// comes when both stayed up across the outage.
g.log.Debugf("network is back, reset reconnection ticker")
ticker.Stop()
ticker = g.newReconnectTicker(ctx)
tickerChannel = ticker.C
iceState.reset()
case <-ctx.Done():
g.log.Debugf("context is done, stop reconnect loop")
return

View File

@@ -15,7 +15,7 @@ import (
func newTestGuard(status connStatusFunc) *Guard {
srw := NewSRWatcher(nil, nil, nil, ice.Config{})
return NewGuard(log.WithField("test", "guard"), status, 50*time.Millisecond, srw, nil)
return NewGuard(log.WithField("test", "guard"), status, 50*time.Millisecond, srw)
}
// countBackoffTickerGoroutines returns how many goroutines are currently sitting

View File

@@ -1,107 +0,0 @@
package guard
import (
"context"
"sync/atomic"
"testing"
"time"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/peer/ice"
"github.com/netbirdio/netbird/client/netstate"
)
// newTestGuardWithNetState builds a guard with a realistic MaxInterval: the
// backoff must be able to grow well past the outage, as it does in production
// where the timeout is seconds to minutes.
func newTestGuardWithNetState(status connStatusFunc, netState *netstate.State) *Guard {
srw := NewSRWatcher(nil, nil, nil, ice.Config{})
return NewGuard(log.WithField("test", "guard"), status, 30*time.Second, srw, netState)
}
// TestGuard_RecoversAfterOfflineToOnline covers a peer that stays disconnected
// across a network outage while neither signal nor relay reports an event —
// both stayed up, as on a short airplane mode toggle over Wi-Fi.
//
// Every tick taken while offline is skipped, but it still advances the
// exponential backoff, so by the time the network returns the next tick can be
// tens of seconds away. Without an explicit reaction to the transition the
// peer waits out that interval for a recovery that could start immediately.
func TestGuard_RecoversAfterOfflineToOnline(t *testing.T) {
netState := netstate.New()
var attempts atomic.Int32
g := newTestGuardWithNetState(func() ConnStatus { return ConnStatusDisconnected }, netState)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
// Start from the reconnect ticker (800ms initial interval), the state a
// peer is in after it loses its connection.
go g.Start(ctx, func() { attempts.Add(1) })
g.SetRelayedConnDisconnected()
// Let the backoff climb: 0.8s, 1.6s, 3.2s, 6.4s ... every tick is skipped
// while offline, but each one doubles the wait for the next.
netState.Set(false)
time.Sleep(8 * time.Second)
offlineAttempts := attempts.Load()
if offlineAttempts != 0 {
t.Fatalf("callback ran %d times while offline, want 0", offlineAttempts)
}
netState.Set(true)
// The next organic tick is now several seconds out, so anything within
// this window can only come from reacting to the transition itself.
pollCtx, stopPolling := context.WithTimeout(ctx, 2*time.Second)
defer stopPolling()
select {
case <-pollCtx.Done():
t.Fatal("peer was not retried within 2s of the network coming back, " +
"with neither a signal nor a relay event to fall back on")
case <-pollUntil(pollCtx, func() bool { return attempts.Load() > 0 }):
}
}
// TestGuard_OfflineTransitionDoesNotRetry checks the other direction: going
// offline must not itself trigger an attempt.
func TestGuard_OfflineTransitionDoesNotRetry(t *testing.T) {
netState := netstate.New()
var attempts atomic.Int32
g := newTestGuardWithNetState(func() ConnStatus { return ConnStatusDisconnected }, netState)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go g.Start(ctx, func() { attempts.Add(1) })
netState.Set(false)
time.Sleep(5 * time.Second)
if got := attempts.Load(); got != 0 {
t.Fatalf("callback ran %d times after going offline, want 0", got)
}
}
// pollUntil closes the returned channel once cond holds. It gives up when ctx
// is done, so the polling goroutine never outlives the test that started it.
func pollUntil(ctx context.Context, cond func() bool) <-chan struct{} {
done := make(chan struct{})
go func() {
for {
if cond() {
close(done)
return
}
select {
case <-ctx.Done():
return
case <-time.After(10 * time.Millisecond):
}
}
}()
return done
}

View File

@@ -1,40 +1,11 @@
package peer
// ClientState identifies the client connection state delivered via
// Listener.OnStateChanged.
type ClientState int
// Client states. The numeric values cross the gomobile boundary (the mobile
// bindings re-export them as integer constants), so they are a wire format:
// append new states at the end, never reorder or insert.
const (
ClientStateDisconnected ClientState = iota
ClientStateConnected
ClientStateConnecting
ClientStateDisconnecting
// ClientStateNoNetwork is an overlay state: it is never stored as the
// last notification, only derived from ClientStateConnecting while the
// OS reports no usable network (see notifier.effectiveState).
ClientStateNoNetwork
)
// Listener is a callback type about the NetBird network connection state
type Listener interface {
// OnStateChanged reports every client state transition. New states are
// delivered only through this callback; the per-state callbacks below
// are kept for compatibility and will be removed once all consumers
// have migrated.
OnStateChanged(state ClientState)
// Deprecated: consume OnStateChanged instead.
OnConnected()
// Deprecated: consume OnStateChanged instead.
OnDisconnected()
// Deprecated: consume OnStateChanged instead.
OnConnecting()
// Deprecated: consume OnStateChanged instead.
OnDisconnecting()
OnAddressChanged(string, string)
OnPeersListChanged(int)
}

View File

@@ -4,64 +4,31 @@ import (
"sync"
)
const (
stateDisconnected = iota
stateConnected
stateConnecting
stateDisconnecting
)
type notifier struct {
// publishLock orders state publication: it is held across computing the
// effective state and handing it to the listener, so a transition cannot
// overtake a newer one and leave the listener on a stale state.
publishLock sync.Mutex
serverStateLock sync.Mutex
listenersLock sync.Mutex
listener Listener
currentClientState bool
lastNotification ClientState
lastNotification int
lastNumberOfPeers int
lastFqdnAddress string
lastIPAddress string
networkAvailable bool
}
func newNotifier() *notifier {
return &notifier{
networkAvailable: true,
}
}
// effectiveState maps the computed state to what listeners should see:
// while the OS reports no usable network, "Connecting" would be a lie —
// connection attempts are suspended — so it is reported as NoNetwork.
// Caller must hold serverStateLock.
func (n *notifier) effectiveState(state ClientState) ClientState {
if !n.networkAvailable && state == ClientStateConnecting {
return ClientStateNoNetwork
}
return state
}
// setNetworkAvailable records the OS network availability and re-notifies
// the listener when the flag flips the effective state (Connecting <->
// NoNetwork).
func (n *notifier) setNetworkAvailable(available bool) {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock()
if n.networkAvailable == available {
n.serverStateLock.Unlock()
return
}
previous := n.effectiveState(n.lastNotification)
n.networkAvailable = available
current := n.effectiveState(n.lastNotification)
n.serverStateLock.Unlock()
if previous != current {
n.notify(current)
}
return &notifier{}
}
func (n *notifier) setListener(listener Listener) {
n.serverStateLock.Lock()
lastNotification := n.effectiveState(n.lastNotification)
lastNotification := n.lastNotification
numOfPeers := n.lastNumberOfPeers
fqdnAddress := n.lastFqdnAddress
address := n.lastIPAddress
@@ -85,9 +52,6 @@ func (n *notifier) removeListener() {
}
func (n *notifier) updateServerStates(mgmState bool, signalState bool) {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock()
calculatedState := n.calculateState(mgmState, signalState)
@@ -97,54 +61,43 @@ func (n *notifier) updateServerStates(mgmState bool, signalState bool) {
}
n.lastNotification = calculatedState
effective := n.effectiveState(calculatedState)
n.serverStateLock.Unlock()
n.notify(effective)
n.notify(calculatedState)
}
func (n *notifier) clientStart() {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock()
n.currentClientState = true
n.lastNotification = ClientStateConnecting
effective := n.effectiveState(ClientStateConnecting)
n.lastNotification = stateConnecting
n.serverStateLock.Unlock()
n.notify(effective)
n.notify(stateConnecting)
}
func (n *notifier) clientStop() {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock()
n.currentClientState = false
n.lastNotification = ClientStateDisconnected
n.lastNotification = stateDisconnected
n.serverStateLock.Unlock()
n.notify(ClientStateDisconnected)
n.notify(stateDisconnected)
}
func (n *notifier) clientTearDown() {
n.publishLock.Lock()
defer n.publishLock.Unlock()
n.serverStateLock.Lock()
n.currentClientState = false
n.lastNotification = ClientStateDisconnecting
n.lastNotification = stateDisconnecting
n.serverStateLock.Unlock()
n.notify(ClientStateDisconnecting)
n.notify(stateDisconnecting)
}
func (n *notifier) isServerStateChanged(newState ClientState) bool {
func (n *notifier) isServerStateChanged(newState int) bool {
return n.lastNotification != newState
}
func (n *notifier) notify(state ClientState) {
func (n *notifier) notify(state int) {
n.listenersLock.Lock()
listener := n.listener
n.listenersLock.Unlock()
@@ -156,20 +109,20 @@ func (n *notifier) notify(state ClientState) {
notifyListener(listener, state)
}
func (n *notifier) calculateState(managementConn, signalConn bool) ClientState {
func (n *notifier) calculateState(managementConn, signalConn bool) int {
if managementConn && signalConn {
return ClientStateConnected
return stateConnected
}
if !managementConn && !signalConn && !n.currentClientState {
return ClientStateDisconnected
return stateDisconnected
}
if n.lastNotification == ClientStateDisconnecting {
return ClientStateDisconnecting
if n.lastNotification == stateDisconnecting {
return stateDisconnecting
}
return ClientStateConnecting
return stateConnecting
}
func (n *notifier) peerListChanged(numOfPeers int) {
@@ -206,19 +159,15 @@ func (n *notifier) localAddressChanged(fqdn, address string) {
listener.OnAddressChanged(fqdn, address)
}
func notifyListener(l Listener, state ClientState) {
// legacy per-state callbacks; NoNetwork is delivered only via
// OnStateChanged below
func notifyListener(l Listener, state int) {
switch state {
case ClientStateDisconnected:
case stateDisconnected:
l.OnDisconnected()
case ClientStateConnected:
case stateConnected:
l.OnConnected()
case ClientStateConnecting:
case stateConnecting:
l.OnConnecting()
case ClientStateDisconnecting:
case stateDisconnecting:
l.OnDisconnecting()
}
l.OnStateChanged(state)
}

View File

@@ -1,108 +0,0 @@
package peer
import (
"sync"
"testing"
"time"
)
type recordingListener struct {
mu sync.Mutex
states []ClientState
onState func(ClientState)
}
func (l *recordingListener) OnStateChanged(state ClientState) {
l.mu.Lock()
l.states = append(l.states, state)
hook := l.onState
l.mu.Unlock()
if hook != nil {
hook(state)
}
}
func (l *recordingListener) last() (ClientState, bool) {
l.mu.Lock()
defer l.mu.Unlock()
if len(l.states) == 0 {
return 0, false
}
return l.states[len(l.states)-1], true
}
func (l *recordingListener) snapshot() []ClientState {
l.mu.Lock()
defer l.mu.Unlock()
return append([]ClientState(nil), l.states...)
}
func (l *recordingListener) OnConnected() {}
func (l *recordingListener) OnDisconnected() {}
func (l *recordingListener) OnConnecting() {}
func (l *recordingListener) OnDisconnecting() {}
func (l *recordingListener) OnAddressChanged(string, string) {}
func (l *recordingListener) OnPeersListChanged(int) {}
// TestNotifier_ConcurrentAvailabilityFlipOrdersPublication holds the first
// transition inside the listener callback and flips availability again from
// another goroutine while it is parked. The second flip must not publish
// ahead of the one in flight, otherwise the listener ends up on a state the
// notifier already superseded.
func TestNotifier_ConcurrentAvailabilityFlipOrdersPublication(t *testing.T) {
n := newNotifier()
n.currentClientState = true
n.lastNotification = ClientStateConnecting
entered := make(chan struct{})
release := make(chan struct{})
l := &recordingListener{}
l.onState = func(state ClientState) {
if state != ClientStateNoNetwork {
return
}
l.mu.Lock()
l.onState = nil
l.mu.Unlock()
close(entered)
<-release
}
n.listener = l
var wg sync.WaitGroup
wg.Add(1)
go func() {
defer wg.Done()
n.setNetworkAvailable(false)
}()
<-entered
flipped := make(chan struct{})
go func() {
defer close(flipped)
n.setNetworkAvailable(true)
}()
select {
case <-flipped:
t.Fatal("the online transition published while the offline one was " +
"still in flight; publication is not serialized")
case <-time.After(200 * time.Millisecond):
}
close(release)
<-flipped
wg.Wait()
got, ok := l.last()
if !ok {
t.Fatal("listener never observed a state")
}
if got != ClientStateConnecting {
t.Fatalf("listener holds %v after the network came back, want Connecting; sequence: %v",
got, l.snapshot())
}
}

View File

@@ -6,32 +6,29 @@ import (
)
type mocListener struct {
lastState ClientState
lastState int
wg sync.WaitGroup
peersWg sync.WaitGroup
peers int
}
func (l *mocListener) OnConnected() {
l.lastState = ClientStateConnected
l.lastState = stateConnected
l.wg.Done()
}
func (l *mocListener) OnDisconnected() {
l.lastState = ClientStateDisconnected
l.lastState = stateDisconnected
l.wg.Done()
}
func (l *mocListener) OnConnecting() {
l.lastState = ClientStateConnecting
l.lastState = stateConnecting
l.wg.Done()
}
func (l *mocListener) OnDisconnecting() {
l.lastState = ClientStateDisconnecting
l.lastState = stateDisconnecting
l.wg.Done()
}
func (l *mocListener) OnStateChanged(state ClientState) {
}
func (l *mocListener) OnAddressChanged(host, addr string) {
}
@@ -60,15 +57,15 @@ func Test_notifier_serverState(t *testing.T) {
type scenario struct {
name string
expected ClientState
expected int
mgmState bool
signalState bool
}
scenarios := []scenario{
{"connected", ClientStateConnected, true, true},
{"mgm down", ClientStateConnecting, false, true},
{"signal down", ClientStateConnecting, true, false},
{"disconnected", ClientStateDisconnected, false, false},
{"connected", stateConnected, true, true},
{"mgm down", stateConnecting, false, true},
{"signal down", stateConnecting, true, false},
{"disconnected", stateDisconnected, false, false},
}
for _, tt := range scenarios {
@@ -88,7 +85,7 @@ func Test_notifier_SetListener(t *testing.T) {
listener.setPeersWaiter()
n := newNotifier()
n.lastNotification = ClientStateConnecting
n.lastNotification = stateConnecting
n.setListener(listener)
listener.wait()
listener.waitPeers()
@@ -102,7 +99,7 @@ func Test_notifier_RemoveListener(t *testing.T) {
listener.setWaiter()
listener.setPeersWaiter()
n := newNotifier()
n.lastNotification = ClientStateConnecting
n.lastNotification = stateConnecting
n.setListener(listener)
// setListener replays cached state on a goroutine; wait for both the state
// and peers callbacks to finish so we don't race on listener.peers.

View File

@@ -1211,12 +1211,6 @@ func (d *Status) ClientTeardown() {
d.notifyStateChange()
}
// SetNetworkAvailable records the OS-reported network availability; while
// unavailable, listeners see NoNetwork instead of Connecting.
func (d *Status) SetNetworkAvailable(available bool) {
d.notifier.setNetworkAvailable(available)
}
// SetConnectionListener set a listener to the notifier
func (d *Status) SetConnectionListener(listener Listener) {
d.notifier.setListener(listener)

View File

@@ -1,130 +0,0 @@
package profilemanager
import (
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"sync"
"github.com/netbirdio/netbird/util"
)
const prefsFileSuffix = ".prefs.json"
var prefsMu sync.Mutex
// Prefs is a namespaced per-profile preference store backed by a single JSON
// file next to the profile config; it is deleted together with the profile.
type Prefs struct {
path string
}
// ProfilePrefs returns the preference store of the profile identified by id.
func (s *ServiceManager) ProfilePrefs(id ID, username string) (*Prefs, error) {
if !IsValidProfileFilenameStem(id) {
return nil, fmt.Errorf("invalid profile ID: %q", id)
}
if id == defaultProfileName {
return &Prefs{path: filepath.Join(filepath.Dir(DefaultConfigPath), id.String()+prefsFileSuffix)}, nil
}
configDir, err := s.getConfigDir(username)
if err != nil {
return nil, fmt.Errorf("get config directory for user %s: %w", username, err)
}
return &Prefs{path: filepath.Join(configDir, id.String()+prefsFileSuffix)}, nil
}
// Get unmarshals the namespace section into v and reports whether it exists.
func (p *Prefs) Get(namespace string, v any) (bool, error) {
if namespace == "" {
return false, fmt.Errorf("empty prefs namespace")
}
prefsMu.Lock()
defer prefsMu.Unlock()
sections, err := readPrefsFile(p.path)
if err != nil {
return false, err
}
raw, ok := sections[namespace]
if !ok {
return false, nil
}
if err := json.Unmarshal(raw, v); err != nil {
return false, fmt.Errorf("decode prefs namespace %q: %w", namespace, err)
}
return true, nil
}
// Put stores v as the namespace section, replacing any previous value.
func (p *Prefs) Put(namespace string, v any) error {
if namespace == "" {
return fmt.Errorf("empty prefs namespace")
}
raw, err := json.Marshal(v)
if err != nil {
return fmt.Errorf("encode prefs namespace %q: %w", namespace, err)
}
prefsMu.Lock()
defer prefsMu.Unlock()
sections, err := readPrefsFile(p.path)
if err != nil {
return err
}
sections[namespace] = raw
return writePrefsFile(p.path, sections)
}
// Remove deletes the namespace section; a missing one is not an error.
func (p *Prefs) Remove(namespace string) error {
if namespace == "" {
return fmt.Errorf("empty prefs namespace")
}
prefsMu.Lock()
defer prefsMu.Unlock()
sections, err := readPrefsFile(p.path)
if err != nil {
return err
}
if _, ok := sections[namespace]; !ok {
return nil
}
delete(sections, namespace)
return writePrefsFile(p.path, sections)
}
func removePrefsFile(path string) error {
prefsMu.Lock()
defer prefsMu.Unlock()
return os.Remove(path)
}
func readPrefsFile(path string) (map[string]json.RawMessage, error) {
data, err := os.ReadFile(path)
if os.IsNotExist(err) {
return map[string]json.RawMessage{}, nil
}
if err != nil {
return nil, fmt.Errorf("read prefs: %w", err)
}
sections := map[string]json.RawMessage{}
if err := json.Unmarshal(data, &sections); err != nil {
return nil, fmt.Errorf("decode prefs: %w", err)
}
return sections, nil
}
func writePrefsFile(path string, sections map[string]json.RawMessage) error {
if err := util.WriteJsonWithRestrictedPermission(context.Background(), path, sections); err != nil {
return fmt.Errorf("write prefs: %w", err)
}
return nil
}

View File

@@ -1,138 +0,0 @@
package profilemanager
import (
"errors"
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type testPrefsSection struct {
Mode uint8 `json:"mode"`
Dest string `json:"dest"`
}
func TestProfilePrefs_RoundTrip(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
created, err := sm.AddProfile("work", username)
require.NoError(t, err)
prefs, err := sm.ProfilePrefs(created.ID, username)
require.NoError(t, err)
require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 2, Dest: "/tmp/x"}))
require.NoError(t, prefs.Put("other", map[string]int{"n": 1}))
var got testPrefsSection
found, err := prefs.Get("filedrop", &got)
require.NoError(t, err)
assert.True(t, found)
assert.Equal(t, testPrefsSection{Mode: 2, Dest: "/tmp/x"}, got)
var other map[string]int
found, err = prefs.Get("other", &other)
require.NoError(t, err)
assert.True(t, found)
assert.Equal(t, map[string]int{"n": 1}, other)
})
}
func TestProfilePrefs_GetMissingNamespace(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
created, err := sm.AddProfile("work", username)
require.NoError(t, err)
prefs, err := sm.ProfilePrefs(created.ID, username)
require.NoError(t, err)
var got testPrefsSection
found, err := prefs.Get("filedrop", &got)
require.NoError(t, err)
assert.False(t, found)
})
}
func TestProfilePrefs_RemoveNamespace(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
created, err := sm.AddProfile("work", username)
require.NoError(t, err)
prefs, err := sm.ProfilePrefs(created.ID, username)
require.NoError(t, err)
require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 1}))
require.NoError(t, prefs.Put("other", map[string]int{"n": 1}))
require.NoError(t, prefs.Remove("filedrop"))
require.NoError(t, prefs.Remove("missing"))
var got testPrefsSection
found, err := prefs.Get("filedrop", &got)
require.NoError(t, err)
assert.False(t, found)
var other map[string]int
found, err = prefs.Get("other", &other)
require.NoError(t, err)
assert.True(t, found)
assert.Equal(t, map[string]int{"n": 1}, other)
})
}
func TestProfilePrefs_RejectsInvalidID(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
_, err := sm.ProfilePrefs("../escape", username)
assert.Error(t, err)
})
}
func TestProfilePrefs_RejectsEmptyNamespace(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
created, err := sm.AddProfile("work", username)
require.NoError(t, err)
prefs, err := sm.ProfilePrefs(created.ID, username)
require.NoError(t, err)
_, err = prefs.Get("", &testPrefsSection{})
assert.Error(t, err)
assert.Error(t, prefs.Put("", testPrefsSection{}))
assert.Error(t, prefs.Remove(""))
})
}
func TestProfilePrefs_DefaultProfile(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
prefs, err := sm.ProfilePrefs(defaultProfileName, username)
require.NoError(t, err)
require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 1}))
expected := filepath.Join(filepath.Dir(DefaultConfigPath), "default"+prefsFileSuffix)
_, err = os.Stat(expected)
require.NoError(t, err)
})
}
func TestRemoveProfile_DeletesPrefsFile(t *testing.T) {
withTestSM(t, func(sm *ServiceManager, username string) {
created, err := sm.AddProfile("work", username)
require.NoError(t, err)
prefs, err := sm.ProfilePrefs(created.ID, username)
require.NoError(t, err)
require.NoError(t, prefs.Put("filedrop", testPrefsSection{Mode: 2}))
configDir, err := sm.getConfigDir(username)
require.NoError(t, err)
prefsPath := filepath.Join(configDir, created.ID.String()+prefsFileSuffix)
_, err = os.Stat(prefsPath)
require.NoError(t, err)
require.NoError(t, sm.RemoveProfile(created.ID, username))
_, err = os.Stat(prefsPath)
assert.True(t, errors.Is(err, os.ErrNotExist), "prefs file should be removed")
})
}

View File

@@ -420,11 +420,6 @@ func (s *ServiceManager) RemoveProfile(id ID, username string) error {
log.Warnf("failed to remove profile state file %s: %v", stateFile, err)
}
prefsFile := filepath.Join(filepath.Dir(target.Path), id.String()+prefsFileSuffix)
if err := removePrefsFile(prefsFile); err != nil && !os.IsNotExist(err) {
log.Warnf("failed to remove profile prefs file %s: %v", prefsFile, err)
}
return nil
}

View File

@@ -37,32 +37,23 @@
// Updater Process (Setup):
//
// 1. Receives parameters from service via command-line arguments
// 2. Terminates the UI so the installer does not have to replace a locked image
// file, which would otherwise leave the install needing a reboot
// 3. Runs installer with appropriate silent/quiet flags:
// 2. Runs installer with appropriate silent/quiet flags:
// - Windows EXE: installer.exe /S
// - Windows MSI: msiexec.exe /i installer.msi /qn /norestart REBOOT=ReallySuppress /l*v msi.log
// - Windows MSI: msiexec.exe /i installer.msi /quiet /qn /l*v msi.log
// - macOS PKG: installer -pkg installer.pkg -target /
// - macOS Homebrew: brew upgrade netbirdio/tap/netbird
// 4. Installer terminates the daemon
// 5. Installer replaces binaries with new version
// 6. Updater waits for installer to complete. On Windows, MSI exit codes 3010
// (ERROR_SUCCESS_REBOOT_REQUIRED) and 1641 (ERROR_SUCCESS_REBOOT_INITIATED)
// are a pending-reboot outcome, not a failure: the install succeeded, but
// some files are only replaced on the next restart (the reboot itself is
// suppressed via /norestart and REBOOT=ReallySuppress), and the flow
// continues as on success
// 7. Updater restarts daemon:
// 3. Installer terminates daemon and UI processes
// 4. Installer replaces binaries with new version
// 5. Updater waits for installer to complete
// 6. Updater restarts daemon:
// - Windows: netbird.exe service start
// - macOS/Linux: netbird service start
// 8. Updater restarts UI:
// - Windows: Launches netbird-ui.exe using CreateProcessAsUser in every
// session it was terminated in, falling back to the active console session
// 7. Updater restarts UI:
// - Windows: Launches netbird-ui.exe as active console user using CreateProcessAsUser
// - macOS: Uses launchctl asuser to launch NetBird.app for console user
// - Linux: Not implemented (UI typically auto-starts)
// 9. Updater writes result.json with success/error status (a pending reboot is
// recorded as success)
// 10. Updater process exits
// 8. Updater writes result.json with success/error status
// 9. Updater process exits
//
// # Result Communication
//

View File

@@ -42,9 +42,6 @@ func NewWithDir(tempDir string) *Installer {
// This will run by the original service process
func (u *Installer) RunInstallation(ctx context.Context, targetVersion string) (err error) {
resultHandler := NewResultHandler(u.tempDir)
if err := resultHandler.ClearStaleResult(); err != nil {
log.Warnf("clear stale installer result: %v", err)
}
defer func() {
if err != nil {

View File

@@ -2,7 +2,6 @@ package installer
import (
"context"
"errors"
"fmt"
"os"
"os/exec"
@@ -23,12 +22,6 @@ const (
msiLogFile = "msi.log"
// ERROR_SUCCESS_REBOOT_REQUIRED and ERROR_SUCCESS_REBOOT_INITIATED
msiRebootRequired = 3010
msiRebootInitiated = 1641
processExitWait = 10 * time.Second
msiDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.msi"
exeDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.exe"
)
@@ -45,8 +38,6 @@ var (
func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string, daemonFolder string) (resultErr error) {
resultHandler := NewResultHandler(u.tempDir)
var uiSessions []uint32
// Always ensure daemon and UI are restarted after setup
defer func() {
log.Infof("starting daemon back")
@@ -55,7 +46,7 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
}
log.Infof("starting UI back")
if err := u.startUI(daemonFolder, uiSessions); err != nil {
if err := u.startUIAsUser(daemonFolder); err != nil {
log.Errorf("failed to start UI: %v", err)
}
@@ -84,14 +75,6 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
return
}
// The UI holds an open handle on its own image. Left running, Restart Manager
// cannot shut it down (msiexec runs as LocalSystem here, the UI as the
// interactive user), so the MSI falls back to replacing the file on reboot and
// marks the install as restart-required. The deferred close-application action
// in the package runs too late to prevent that, it happens after
// InstallValidate has already registered the file as in use.
uiSessions = killUI()
var cmd *exec.Cmd
switch installerType {
case TypeExe:
@@ -101,9 +84,7 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
installerDir := filepath.Dir(installerFile)
logPath := filepath.Join(installerDir, msiLogFile)
log.Infof("run msi installer: %s", installerFile)
// REBOOT=ReallySuppress: a silent install has no way to ask, so without it
// msiexec reboots the machine on its own if it decides one is needed.
cmd = exec.CommandContext(ctx, "msiexec.exe", "/i", filepath.Base(installerFile), "/qn", "/norestart", "REBOOT=ReallySuppress", "/l*v", logPath)
cmd = exec.CommandContext(ctx, "msiexec.exe", "/i", filepath.Base(installerFile), "/quiet", "/qn", "/l*v", logPath)
}
cmd.Dir = filepath.Dir(installerFile)
@@ -114,13 +95,9 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
}
log.Infof("installer started with PID %d", cmd.Process.Pid)
if err := cmd.Wait(); err != nil {
if !isRebootPending(err) {
resultErr = err
log.Errorf("installer process finished with error: %v", err)
return
}
log.Warnf("installer completed but reported a pending reboot, some files will be replaced on the next restart")
if resultErr = cmd.Wait(); resultErr != nil {
log.Errorf("installer process finished with error: %v", resultErr)
return
}
return nil
@@ -140,142 +117,16 @@ func (u *Installer) startDaemon(daemonFolder string) error {
return nil
}
func (u *Installer) startUI(daemonFolder string, sessionIDs []uint32) error {
func (u *Installer) startUIAsUser(daemonFolder string) error {
uiPath := filepath.Join(daemonFolder, uiName)
log.Infof("starting netbird-ui: %s", uiPath)
if len(sessionIDs) == 0 {
sessionID := windows.WTSGetActiveConsoleSessionId()
if sessionID == 0xFFFFFFFF {
return fmt.Errorf("no active user session found")
}
sessionIDs = []uint32{sessionID}
// Get the active console session ID
sessionID := windows.WTSGetActiveConsoleSessionId()
if sessionID == 0xFFFFFFFF {
return fmt.Errorf("no active user session found")
}
var errs []error
for _, sessionID := range sessionIDs {
if err := startUIInSession(uiPath, sessionID); err != nil {
errs = append(errs, fmt.Errorf("session %d: %w", sessionID, err))
continue
}
log.Infof("netbird-ui started successfully in session %d", sessionID)
}
return errors.Join(errs...)
}
// isRebootPending reports whether the installer exit code means it succeeded but
// left work for the next restart. The reboot itself is suppressed, so this is not
// a failure.
func isRebootPending(err error) bool {
var exitErr *exec.ExitError
if !errors.As(err, &exitErr) {
return false
}
switch exitErr.ExitCode() {
case msiRebootRequired, msiRebootInitiated:
return true
default:
return false
}
}
// killUI terminates any running netbird-ui process and returns the IDs of the
// interactive sessions the terminated processes belonged to. Setup starts the
// UI again in those sessions once the installer is done.
func killUI() []uint32 {
pids, err := processIDsByName(uiName)
if err != nil {
log.Warnf("failed to look up %s processes: %v", uiName, err)
return nil
}
sessions := make(map[uint32]struct{})
for _, pid := range pids {
var sessionID uint32
if err := windows.ProcessIdToSessionId(pid, &sessionID); err != nil {
log.Warnf("failed to look up session of %s (PID %d): %v", uiName, pid, err)
}
if err := terminateProcess(pid); err != nil {
log.Warnf("failed to terminate %s (PID %d): %v", uiName, pid, err)
continue
}
log.Infof("terminated %s (PID %d) in session %d", uiName, pid, sessionID)
if sessionID != 0 {
sessions[sessionID] = struct{}{}
}
}
sessionIDs := make([]uint32, 0, len(sessions))
for sessionID := range sessions {
sessionIDs = append(sessionIDs, sessionID)
}
return sessionIDs
}
func processIDsByName(name string) ([]uint32, error) {
snapshot, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPPROCESS, 0)
if err != nil {
return nil, fmt.Errorf("create process snapshot: %w", err)
}
defer func() {
if err := windows.CloseHandle(snapshot); err != nil {
log.Warnf("failed to close process snapshot: %v", err)
}
}()
var entry windows.ProcessEntry32
entry.Size = uint32(unsafe.Sizeof(entry))
var pids []uint32
for err = windows.Process32First(snapshot, &entry); err == nil; err = windows.Process32Next(snapshot, &entry) {
if strings.EqualFold(windows.UTF16ToString(entry.ExeFile[:]), name) {
pids = append(pids, entry.ProcessID)
}
}
if !errors.Is(err, windows.ERROR_NO_MORE_FILES) {
return nil, fmt.Errorf("enumerate processes: %w", err)
}
return pids, nil
}
func terminateProcess(pid uint32) error {
handle, err := windows.OpenProcess(windows.PROCESS_TERMINATE|windows.SYNCHRONIZE, false, pid)
if err != nil {
// The process may have exited between enumeration and now.
if errors.Is(err, windows.ERROR_INVALID_PARAMETER) {
return nil
}
return fmt.Errorf("open process: %w", err)
}
defer func() {
if err := windows.CloseHandle(handle); err != nil {
log.Warnf("failed to close process handle: %v", err)
}
}()
if err := windows.TerminateProcess(handle, 0); err != nil {
return fmt.Errorf("terminate process: %w", err)
}
// Wait for the handle to signal so the image file is released before the
// installer tries to overwrite it. A timeout is reported through the returned
// event, not through err, which stays nil unless the wait itself failed.
event, err := windows.WaitForSingleObject(handle, uint32(processExitWait.Milliseconds()))
if err != nil {
return fmt.Errorf("wait for process exit: %w", err)
}
if event != windows.WAIT_OBJECT_0 {
return fmt.Errorf("wait for process exit: unexpected wait result %#x", event)
}
return nil
}
func startUIInSession(uiPath string, sessionID uint32) error {
// Get the user token for that session
var userToken windows.Token
err := windows.WTSQueryUserToken(sessionID, &userToken)
@@ -307,16 +158,6 @@ func startUIInSession(uiPath string, sessionID uint32) error {
}
}()
var env *uint16
if err := windows.CreateEnvironmentBlock(&env, primaryToken, false); err != nil {
return fmt.Errorf("create environment block: %w", err)
}
defer func() {
if err := windows.DestroyEnvironmentBlock(env); err != nil {
log.Warnf("failed to destroy environment block: %v", err)
}
}()
// Prepare startup info
var si windows.StartupInfo
si.Cb = uint32(unsafe.Sizeof(si))
@@ -339,7 +180,7 @@ func startUIInSession(uiPath string, sessionID uint32) error {
nil,
false,
creationFlags,
env,
nil,
nil,
&si,
&pi,
@@ -356,6 +197,7 @@ func startUIInSession(uiPath string, sessionID uint32) error {
log.Warnf("failed to close thread handle: %v", err)
}
log.Infof("netbird-ui started successfully in session %d", sessionID)
return nil
}

View File

@@ -1,108 +0,0 @@
package installer
import (
"errors"
"os/exec"
"slices"
"strconv"
"testing"
)
// exitErrorWithCode returns a real *exec.ExitError carrying the given exit code.
func exitErrorWithCode(t *testing.T, code int) error {
t.Helper()
err := exec.Command("cmd.exe", "/c", "exit "+strconv.Itoa(code)).Run()
if err == nil {
t.Fatalf("expected a non-zero exit for code %d", code)
}
return err
}
func TestIsRebootPending(t *testing.T) {
tests := []struct {
name string
code int
want bool
}{
{name: "reboot required", code: msiRebootRequired, want: true},
{name: "reboot initiated", code: msiRebootInitiated, want: true},
{name: "generic failure", code: 1603, want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := isRebootPending(exitErrorWithCode(t, tt.code)); got != tt.want {
t.Errorf("isRebootPending(exit %d) = %v, want %v", tt.code, got, tt.want)
}
})
}
}
// TestProcessIDsByNameAndTerminate spawns a long-running system process, finds it
// by name and terminates it, covering the path the updater uses to release the UI
// image file before the installer replaces it.
func TestProcessIDsByNameAndTerminate(t *testing.T) {
cmd := exec.Command("ping.exe", "-n", "60", "127.0.0.1")
if err := cmd.Start(); err != nil {
t.Fatalf("start ping: %v", err)
}
pid := uint32(cmd.Process.Pid)
killed := false
t.Cleanup(func() {
if !killed {
_ = cmd.Process.Kill()
}
_ = cmd.Wait()
})
// Name matching must be case-insensitive: the snapshot reports PING.EXE.
pids, err := processIDsByName("ping.exe")
if err != nil {
t.Fatalf("processIDsByName: %v", err)
}
if !slices.Contains(pids, pid) {
t.Fatalf("PID %d not among the ping.exe processes found: %v", pid, pids)
}
if err := terminateProcess(pid); err != nil {
t.Fatalf("terminateProcess: %v", err)
}
killed = true
// terminateProcess only returns once the handle has signalled, so the process
// is already gone and Wait must not block. It exits with the code passed to
// TerminateProcess, which is 0, so Wait reports no error.
if err := cmd.Wait(); err != nil {
t.Fatalf("wait for terminated ping: %v", err)
}
if !cmd.ProcessState.Exited() {
t.Error("process did not exit after terminateProcess")
}
remaining, err := processIDsByName("ping.exe")
if err != nil {
t.Fatalf("processIDsByName after terminate: %v", err)
}
if slices.Contains(remaining, pid) {
t.Errorf("PID %d still listed after terminateProcess", pid)
}
}
func TestProcessIDsByNameNoMatch(t *testing.T) {
pids, err := processIDsByName("netbird-nonexistent-process.exe")
if err != nil {
t.Fatalf("processIDsByName: %v", err)
}
if len(pids) != 0 {
t.Errorf("expected no matches, got %v", pids)
}
}
func TestIsRebootPendingNonExitError(t *testing.T) {
if isRebootPending(errors.New("start installer: file not found")) {
t.Error("a non-exit error must not be treated as a pending reboot")
}
}

View File

@@ -54,12 +54,6 @@ func (rh *ResultHandler) GetErrorResultReason() string {
return ""
}
// ClearStaleResult removes a result file left over from a previous installation
// attempt so result watchers cannot read an outdated outcome for the current attempt.
func (rh *ResultHandler) ClearStaleResult() error {
return rh.cleanup()
}
func (rh *ResultHandler) WriteSuccess() error {
result := Result{
Success: true,

View File

@@ -22,8 +22,6 @@ import (
"github.com/netbirdio/netbird/client/internal/listener"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/netstate"
"github.com/netbirdio/netbird/client/netsweep"
"github.com/netbirdio/netbird/client/system"
"github.com/netbirdio/netbird/formatter"
"github.com/netbirdio/netbird/route"
@@ -38,6 +36,11 @@ const (
AnonymizeLevelStrict = nbAnonymize.LevelStrictString
)
// ConnectionListener export internal Listener for mobile
type ConnectionListener interface {
peer.Listener
}
// RouteListener export internal RouteListener for mobile
type NetworkChangeListener interface {
listener.NetworkChangeListener
@@ -84,12 +87,6 @@ type Client struct {
onHostDnsFn func([]string)
dnsManager dns.IosDnsManager
loginComplete bool
// netState outlives engine restarts: it mirrors the OS connectivity, not
// the engine lifecycle. Run injects it into each new ConnectClient, which
// distributes it to every reconnection loop.
netState *netstate.State
// sweeper also outlives engine restarts; NotifyNetworkChange sweeps it.
sweeper *netsweep.Sweeper
// preloadedConfig holds config loaded from JSON (used on tvOS where file writes are blocked)
preloadedConfig *profilemanager.Config
@@ -112,8 +109,6 @@ func NewClient(cfgFile, stateFile, cacheDir, logFilePath, deviceName string, osV
ctxCancelLock: &sync.Mutex{},
networkChangeListener: networkChangeListener,
dnsManager: dnsManager,
netState: netstate.New(),
sweeper: netsweep.New(),
}
}
@@ -189,8 +184,7 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error {
c.onHostDnsFn = func([]string) {}
cfg.WgIface = interfaceName
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder,
internal.WithNetworkState(c.netState), internal.WithSweeper(c.sweeper))
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder)
c.setState(cfg, connectClient)
// Persist the latest sync response so DebugBundle can include the network
// map. On iOS this is backed by disk to keep it out of the constrained
@@ -199,25 +193,6 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error {
return connectClient.RunOniOS(fd, c.networkChangeListener, c.dnsManager, c.stateFile, c.cacheDir, c.logFilePath)
}
// SetNetworkAvailable feeds OS-reported network availability into the client
// (e.g. from NWPathMonitor). While unavailable, the internal reconnect loops
// suspend their attempts and the connection listener reports NoNetwork
// instead of Connecting; when availability returns, the loops resume
// immediately with a fresh backoff.
func (c *Client) SetNetworkAvailable(available bool) {
c.netState.Set(available)
c.recorder.SetNetworkAvailable(available)
}
// NotifyNetworkChange marks the management, signal and relay connections
// stale after the OS switched networks and schedules a sweep that cuts
// whatever has not redialed on the new network by then. The engine and the
// TUN device stay untouched.
func (c *Client) NotifyNetworkChange() {
c.sweeper.MarkNetworkChange()
log.Infof("network change: connections marked stale")
}
// Stop the internal client and free the resources
func (c *Client) Stop() {
c.ctxCancelLock.Lock()
@@ -356,11 +331,7 @@ func (c *Client) GetStatusDetails() *StatusDetails {
// SetConnectionListener set the network connection listener
func (c *Client) SetConnectionListener(listener ConnectionListener) {
if listener == nil {
c.recorder.RemoveConnectionListener()
return
}
c.recorder.SetConnectionListener(connectionListenerAdapter{listener})
c.recorder.SetConnectionListener(listener)
}
// RemoveConnectionListener remove connection listener

View File

@@ -1,43 +0,0 @@
//go:build ios
package NetBirdSDK
import (
"github.com/netbirdio/netbird/client/internal/peer"
)
// Client state values, re-exported as basic constants so gomobile emits them
// into the generated bindings. They mirror peer.ClientState*: append-only,
// never reorder.
const (
ClientStateDisconnected = int(peer.ClientStateDisconnected)
ClientStateConnected = int(peer.ClientStateConnected)
ClientStateConnecting = int(peer.ClientStateConnecting)
ClientStateDisconnecting = int(peer.ClientStateDisconnecting)
ClientStateNoNetwork = int(peer.ClientStateNoNetwork)
)
// ConnectionListener export internal Listener for mobile.
//
// It intentionally lacks OnStateChanged for now: adding a method to a gomobile
// interface breaks every Swift implementation, so the iOS app keeps building
// against the legacy per-state callbacks. A follow-up will extend it together
// with the app.
type ConnectionListener interface {
OnConnected()
OnDisconnected()
OnConnecting()
OnDisconnecting()
OnAddressChanged(string, string)
OnPeersListChanged(int)
}
// connectionListenerAdapter adapts the gomobile-facing ConnectionListener to
// peer.Listener.
type connectionListenerAdapter struct {
ConnectionListener
}
// OnStateChanged is dropped on iOS until the app adopts the state callback;
// the legacy per-state callbacks continue to fire.
func (a connectionListenerAdapter) OnStateChanged(peer.ClientState) {}

View File

@@ -323,7 +323,7 @@ func (a *Auth) login(urlOpener URLOpener, forceDeviceAuth bool, deviceName strin
const authInfoRequestTimeout = 30 * time.Second
func (a *Auth) foregroundGetTokenInfo(authClient *auth.Auth, urlOpener URLOpener, forceDeviceAuth bool) (*auth.TokenInfo, error) {
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth, "")
oAuthFlow, err := authClient.GetOAuthFlow(a.ctx, forceDeviceAuth)
if err != nil {
return nil, fmt.Errorf("failed to get OAuth flow: %v", err)
}

View File

@@ -1,110 +0,0 @@
// Package netstate tracks OS-reported network availability for the client.
//
// A State instance is owned by the platform integration (e.g. the Android or
// iOS bindings, fed from ConnectivityManager callbacks or NWPathMonitor) and
// is injected into the connection retry loops (management, signal, relay,
// peer guards and the top-level connect loop), which consult it to avoid
// burning CPU and battery on reconnect attempts while the device has no
// network at all (e.g. airplane mode), and to reset their backoff as soon as
// the network returns.
//
// Consumers hold a *State that may be nil — every non-mobile platform leaves
// it unset. The read methods are safe on a nil receiver: they report online
// and never block, so consumers behave as if this package did not exist.
package netstate
import (
"context"
"sync"
log "github.com/sirupsen/logrus"
)
// State holds the OS-reported network availability. The zero value is not
// usable; create instances with New.
type State struct {
mu sync.Mutex
online bool
changed chan struct{}
}
// New creates a State that starts online. Platforms without network tracking
// pass a nil *State instead: the read methods treat nil as always online and
// never block, so consumers need no nil guards.
func New() *State {
return &State{
online: true,
changed: make(chan struct{}),
}
}
// Set records whether the OS reports any usable network. Transitions wake up
// all Wait callers immediately. Unlike the read methods, Set is not nil-safe:
// it is only for the platform owner that created the State with New.
func (s *State) Set(online bool) {
s.mu.Lock()
defer s.mu.Unlock()
if s.online == online {
return
}
s.online = online
close(s.changed)
s.changed = make(chan struct{})
log.Infof("OS network availability changed: online=%t", online)
}
// IsOnline reports whether the OS reports at least one usable network. On a
// nil receiver — no State injected — it reports online.
func (s *State) IsOnline() bool {
if s == nil {
return true
}
s.mu.Lock()
defer s.mu.Unlock()
return s.online
}
// Changed returns a channel closed on the next availability transition, for
// callers that already own a select loop and cannot block in Wait. Re-read it
// after every fire: each transition installs a fresh channel. On a nil
// receiver — no State injected — it returns nil, which blocks forever in a
// select, so the caller simply never observes a transition.
func (s *State) Changed() <-chan struct{} {
if s == nil {
return nil
}
s.mu.Lock()
defer s.mu.Unlock()
return s.changed
}
// Wait blocks while the network is offline. It reports whether it had to
// wait, so callers can reset their backoff after an outage. It returns early
// with the context error when ctx is done. On a nil receiver — no State
// injected — it returns immediately.
func (s *State) Wait(ctx context.Context) (bool, error) {
if s == nil {
return false, nil
}
waited := false
for {
s.mu.Lock()
if s.online {
s.mu.Unlock()
return waited, nil
}
ch := s.changed
s.mu.Unlock()
if !waited {
waited = true
log.Debugf("network is offline, pausing connection attempts")
}
select {
case <-ctx.Done():
return waited, ctx.Err()
case <-ch:
}
}
}

View File

@@ -1,170 +0,0 @@
package netstate
import (
"context"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestNewStateIsOnline(t *testing.T) {
assert.True(t, New().IsOnline(), "a fresh State should start online")
}
func TestSetTogglesOnlineState(t *testing.T) {
s := New()
s.Set(false)
assert.False(t, s.IsOnline(), "state should be offline after Set(false)")
s.Set(true)
assert.True(t, s.IsOnline(), "state should be online after Set(true)")
}
func TestWaitReturnsImmediatelyWhenOnline(t *testing.T) {
s := New()
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
waited, err := s.Wait(ctx)
require.NoError(t, err)
assert.False(t, waited, "Wait should not block when the network is online")
}
func TestWaitBlocksUntilOnline(t *testing.T) {
s := New()
s.Set(false)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
result := make(chan bool, 1)
go func() {
waited, err := s.Wait(ctx)
if err != nil {
result <- false
return
}
result <- waited
}()
// Verify Wait is actually blocking while offline
select {
case <-result:
t.Fatal("Wait should block while the network is offline")
case <-time.After(100 * time.Millisecond):
}
s.Set(true)
select {
case waited := <-result:
assert.True(t, waited, "Wait should report that it had to wait for the network")
case <-time.After(2 * time.Second):
t.Fatal("Wait should return promptly after the network becomes available")
}
}
func TestWaitReturnsOnContextCancel(t *testing.T) {
s := New()
s.Set(false)
ctx, cancel := context.WithCancel(context.Background())
result := make(chan error, 1)
go func() {
_, err := s.Wait(ctx)
result <- err
}()
cancel()
select {
case err := <-result:
assert.ErrorIs(t, err, context.Canceled)
case <-time.After(2 * time.Second):
t.Fatal("Wait should return promptly after context cancellation")
}
}
func TestWaitWakesAllWaiters(t *testing.T) {
s := New()
s.Set(false)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
const waiters = 10
var wg sync.WaitGroup
results := make(chan bool, waiters)
for i := 0; i < waiters; i++ {
wg.Add(1)
go func() {
defer wg.Done()
waited, err := s.Wait(ctx)
if err != nil {
results <- false
return
}
results <- waited
}()
}
time.Sleep(100 * time.Millisecond)
s.Set(true)
wg.Wait()
close(results)
count := 0
for waited := range results {
assert.True(t, waited, "every waiter should report that it waited")
count++
}
assert.Equal(t, waiters, count, "all waiters should have returned")
}
func TestNilStateReadsAreNoops(t *testing.T) {
var s *State
assert.True(t, s.IsOnline(), "nil State should report online")
waited, err := s.Wait(context.Background())
require.NoError(t, err)
assert.False(t, waited, "nil State's Wait should not block")
}
func TestConcurrentSetAndWait(t *testing.T) {
s := New()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
var wg sync.WaitGroup
for i := 0; i < 4; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for j := 0; j < 100; j++ {
s.Set(j%2 == 0)
s.IsOnline()
}
}()
}
for i := 0; i < 4; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for j := 0; j < 100; j++ {
if _, err := s.Wait(ctx); err != nil {
return
}
}
}()
}
wg.Wait()
}

View File

@@ -1,267 +0,0 @@
// Package netsweep cuts network-bound activity when the OS switches networks:
// a sweep closes the registered connections and aborts the in-flight dials, so
// their owners redial immediately instead of waiting for the old sockets to
// time out.
//
// A nil *Sweeper disables everything: all methods are nil-safe no-ops.
package netsweep
import (
"context"
"errors"
"net"
"sync"
"time"
"github.com/cenkalti/backoff/v4"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/netstate"
)
// DefaultSweepDelay absorbs network flapping while the OS settles on a
// default network before the stale registrations are cut.
const DefaultSweepDelay = 500 * time.Millisecond
const recentMarkWindow = 3 * time.Second
// Config customizes a Sweeper. The zero value applies the defaults.
type Config struct {
// SweepDelay overrides DefaultSweepDelay when positive.
SweepDelay time.Duration
}
// ErrSwept reports that a dial finished after a network change swept its
// registration. The connection is already closed; the caller must treat it
// as a failed dial and redial on the new network.
var ErrSwept = errors.New("netsweep: connection swept by network change")
// sweepID identifies one registration in a sweeper. Connections and dials
// draw from the same counter, so an id is unique across both registries.
type sweepID uint64
type connEntry struct {
conn net.Conn
gen uint64
}
// Dial tracks one dial from start to connection registration. It hands the
// dialed connection to the sweeper atomically, so a sweep can never fall
// between the dial finishing and the connection being registered.
type Dial struct {
sweeper *Sweeper
ctx context.Context
cancel context.CancelFunc
id sweepID
done bool // set by a sweep, WrapConn or Release; guarded by sweeper.mu
gen uint64
}
// Ctx returns the dial's context. A sweep cancels it, so a dial started on the
// old network aborts instead of waiting out its handshake timeout.
func (d *Dial) Ctx() context.Context {
return d.ctx
}
// Release ends the dial's registration and cancels its context. It is
// idempotent and safe after WrapConn, so callers can defer it.
func (d *Dial) Release() {
s := d.sweeper
if s == nil {
return
}
s.mu.Lock()
d.done = true
delete(s.dials, d.id)
s.mu.Unlock()
d.cancel()
}
// sweptConn deregisters itself from the sweeper when closed.
type sweptConn struct {
net.Conn
sweeper *Sweeper
id sweepID
}
func (c *sweptConn) Close() error {
c.sweeper.deregister(c.id)
return c.Conn.Close()
}
// Sweeper registers live connections and in-flight dials so the
// network-change sweep can cut everything registered before the change.
type Sweeper struct {
mu sync.Mutex
conns map[sweepID]connEntry
dials map[sweepID]*Dial
nextID sweepID
gen uint64
timer *time.Timer
sweepDelay time.Duration
lastMark time.Time
}
// New creates an empty sweeper with the default configuration.
func New() *Sweeper {
return NewWithConfig(Config{})
}
// NewWithConfig creates an empty sweeper customized by cfg.
func NewWithConfig(cfg Config) *Sweeper {
delay := cfg.SweepDelay
if delay <= 0 {
delay = DefaultSweepDelay
}
return &Sweeper{
conns: make(map[sweepID]connEntry),
dials: make(map[sweepID]*Dial),
sweepDelay: delay,
}
}
// StartDial registers an in-flight dial. Dial with Ctx, hand the result to
// WrapConn, and Release the dial when the attempt is over, typically deferred.
func (s *Sweeper) StartDial(ctx context.Context) *Dial {
if s == nil {
return &Dial{ctx: ctx}
}
ctx, cancel := context.WithCancel(ctx)
d := &Dial{sweeper: s, ctx: ctx, cancel: cancel}
s.mu.Lock()
d.id = s.nextID
s.nextID++
d.gen = s.gen
s.dials[d.id] = d
s.mu.Unlock()
return d
}
// WrapConn hands conn over to the sweeper. If a sweep ran since StartDial,
// the connection belongs to the old network: it is closed and ErrSwept is
// returned. Otherwise conn is registered against the next sweep and returned
// wrapped, deregistering itself on Close. Call it once, before Release.
func (d *Dial) WrapConn(conn net.Conn) (net.Conn, error) {
s := d.sweeper
if s == nil {
return conn, nil
}
s.mu.Lock()
if d.done {
s.mu.Unlock()
if err := conn.Close(); err != nil {
log.Debugf("swept dial close error: %v", err)
}
return nil, ErrSwept
}
d.done = true
delete(s.dials, d.id)
id := s.nextID
s.nextID++
// The conn inherits the dial's generation: the socket was bound to the
// network that was default when the dial started, not when it finished.
s.conns[id] = connEntry{conn: conn, gen: d.gen}
s.mu.Unlock()
return &sweptConn{Conn: conn, sweeper: s, id: id}, nil
}
// MarkNetworkChange records that the OS switched networks: everything
// registered so far becomes stale, and a sweep is (re)scheduled after the
// configured delay to cut whatever is still stale by then. Owners that
// redialed in the meantime hold fresh-generation registrations and survive,
// so no cancellation is needed around the sweep.
func (s *Sweeper) MarkNetworkChange() {
if s == nil {
return
}
s.mu.Lock()
s.gen++
cutoff := s.gen
s.lastMark = time.Now()
if s.timer != nil {
s.timer.Stop()
}
s.timer = time.AfterFunc(s.sweepDelay, func() {
n := s.sweep(cutoff)
log.Infof("network change sweep: closed %d stale connections", n)
})
s.mu.Unlock()
}
// QuickRetryBackoff wraps bo so that after each Reset the first retry comes
// quickly when the disconnect followed a recent network change and the
// network is online. Any other failure keeps bo's spread, so the clients of
// a restarted server still scatter their reconnects. A nil sweeper returns
// bo unchanged.
func (s *Sweeper) QuickRetryBackoff(ctx context.Context, bo backoff.BackOff, netState *netstate.State) backoff.BackOff {
if s == nil {
return bo
}
return backoff.WithContext(newQuickRetryBackoff(bo, s, netState), ctx)
}
func (s *Sweeper) markedRecently() bool {
if s == nil {
return false
}
s.mu.Lock()
defer s.mu.Unlock()
return !s.lastMark.IsZero() && time.Since(s.lastMark) < recentMarkWindow
}
// sweep closes the registered connections and aborts the in-flight dials
// older than cutoff, and returns how many connections it closed. A dial
// whose connection was not yet handed to WrapConn is marked, so the late
// WrapConn closes it instead of registering it.
func (s *Sweeper) sweep(cutoff uint64) int {
if s == nil {
return 0
}
s.mu.Lock()
var conns []net.Conn
for id, e := range s.conns {
if e.gen < cutoff {
delete(s.conns, id)
conns = append(conns, e.conn)
}
}
var dials []*Dial
for id, d := range s.dials {
if d.gen < cutoff {
d.done = true
delete(s.dials, id)
dials = append(dials, d)
}
}
s.mu.Unlock()
if len(dials) > 0 {
log.Debugf("aborting %d in-flight dials", len(dials))
for _, d := range dials {
d.cancel()
}
}
for _, conn := range conns {
log.Debugf("sweeping connection %s -> %s", conn.LocalAddr(), conn.RemoteAddr())
if err := conn.Close(); err != nil {
log.Debugf("swept connection close error: %v", err)
}
}
return len(conns)
}
func (s *Sweeper) deregister(id sweepID) {
s.mu.Lock()
delete(s.conns, id)
s.mu.Unlock()
}

View File

@@ -1,241 +0,0 @@
package netsweep
import (
"context"
"math"
"net"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestSweepClosesRegisteredConns(t *testing.T) {
sweeper := New()
c1 := wrap(t, sweeper, connPair(t))
c2 := wrap(t, sweeper, connPair(t))
assert.Equal(t, 2, sweeper.sweepAll(), "both live connections should be closed")
// The wrappers must report closed now.
buf := make([]byte, 1)
_, err := c1.Read(buf)
assert.Error(t, err, "first connection should be unusable after the sweep")
_, err = c2.Read(buf)
assert.Error(t, err, "second connection should be unusable after the sweep")
assert.Equal(t, 0, sweeper.sweepAll(), "second sweep should find nothing")
}
func TestCloseDeregisters(t *testing.T) {
sweeper := New()
conn := wrap(t, sweeper, connPair(t))
require.NoError(t, conn.Close())
assert.Equal(t, 0, sweeper.sweepAll(), "closed connection must leave the registry")
}
func TestCloseIsIdempotent(t *testing.T) {
sweeper := New()
conn := wrap(t, sweeper, connPair(t))
require.NoError(t, conn.Close())
assert.Error(t, conn.Close(), "double close surfaces the underlying error but must not panic")
}
func TestSweepOnlyAffectsOlderConns(t *testing.T) {
sweeper := New()
_ = wrap(t, sweeper, connPair(t))
assert.Equal(t, 1, sweeper.sweepAll())
// A connection dialed after the sweep must survive until the next one.
_ = wrap(t, sweeper, connPair(t))
assert.Equal(t, 1, sweeper.sweepAll(), "post-sweep connection belongs to the next sweep")
}
func TestSweepAbortsInFlightDials(t *testing.T) {
sweeper := New()
dial := sweeper.StartDial(context.Background())
defer dial.Release()
sweeper.sweepAll()
assert.ErrorIs(t, dial.Ctx().Err(), context.Canceled, "sweep must cancel the in-flight dial context")
}
func TestReleasedDialIsNotAborted(t *testing.T) {
sweeper := New()
// Simulate a dial that finished before the sweep.
released := sweeper.StartDial(context.Background())
released.Release()
// A dial still in flight during the sweep.
pending := sweeper.StartDial(context.Background())
defer pending.Release()
sweeper.sweepAll()
assert.ErrorIs(t, pending.Ctx().Err(), context.Canceled, "pending dial must be aborted")
}
func TestSweepBetweenDialAndHandoffClosesConn(t *testing.T) {
sweeper := New()
dial := sweeper.StartDial(context.Background())
defer dial.Release()
// The dial succeeds on the old network, then the sweep lands before the
// connection is handed over.
conn := connPair(t)
assert.Equal(t, 0, sweeper.sweepAll(), "the connection is not registered yet")
wrapped, err := dial.WrapConn(conn)
require.ErrorIs(t, err, ErrSwept)
require.Nil(t, wrapped)
buf := make([]byte, 1)
_, err = conn.Read(buf)
assert.Error(t, err, "the old-network connection must be closed, not leaked")
assert.Equal(t, 0, sweeper.sweepAll(), "nothing may leak into the next sweep")
}
func TestMarkNetworkChangeSparesFreshConns(t *testing.T) {
sweeper := NewWithConfig(Config{SweepDelay: 10 * time.Millisecond})
stale := wrap(t, sweeper, connPair(t))
sweeper.MarkNetworkChange()
_ = wrap(t, sweeper, connPair(t))
_ = stale.SetReadDeadline(time.Now().Add(time.Second))
buf := make([]byte, 1)
_, err := stale.Read(buf)
require.ErrorIs(t, err, net.ErrClosed, "stale connection must be closed by the delayed sweep")
assert.Equal(t, 1, sweeper.sweepAll(), "the fresh connection must survive the stale sweep")
}
func TestMarkNetworkChangeAbortsStaleDials(t *testing.T) {
sweeper := NewWithConfig(Config{SweepDelay: 10 * time.Millisecond})
stale := sweeper.StartDial(context.Background())
defer stale.Release()
sweeper.MarkNetworkChange()
fresh := sweeper.StartDial(context.Background())
defer fresh.Release()
assert.Eventually(t, func() bool {
return stale.Ctx().Err() != nil
}, time.Second, 5*time.Millisecond, "stale dial must be aborted by the delayed sweep")
assert.NoError(t, fresh.Ctx().Err(), "post-mark dial must not be aborted")
}
func TestConnInheritsDialGeneration(t *testing.T) {
sweeper := NewWithConfig(Config{SweepDelay: 20 * time.Millisecond})
// The dial starts before the network change but completes after it: the
// socket is bound to the old network, so the sweep must still cut it.
dial := sweeper.StartDial(context.Background())
defer dial.Release()
sweeper.MarkNetworkChange()
wrapped, err := dial.WrapConn(connPair(t))
require.NoError(t, err)
_ = wrapped.SetReadDeadline(time.Now().Add(time.Second))
buf := make([]byte, 1)
_, err = wrapped.Read(buf)
require.ErrorIs(t, err, net.ErrClosed, "old-generation connection must be swept")
}
func TestRepeatedMarksCoalesce(t *testing.T) {
sweeper := NewWithConfig(Config{SweepDelay: 20 * time.Millisecond})
first := wrap(t, sweeper, connPair(t))
sweeper.MarkNetworkChange()
second := wrap(t, sweeper, connPair(t))
sweeper.MarkNetworkChange()
_ = wrap(t, sweeper, connPair(t))
buf := make([]byte, 1)
for _, conn := range []net.Conn{first, second} {
_ = conn.SetReadDeadline(time.Now().Add(time.Second))
_, err := conn.Read(buf)
require.ErrorIs(t, err, net.ErrClosed, "every pre-mark connection must be swept by the rescheduled sweep")
}
assert.Equal(t, 1, sweeper.sweepAll(), "only the newest-generation connection may remain")
}
func TestNilSweeperIsNoop(t *testing.T) {
var sweeper *Sweeper
conn := connPair(t)
dial := sweeper.StartDial(context.Background())
defer dial.Release()
wrapped, err := dial.WrapConn(conn)
require.NoError(t, err)
assert.Equal(t, conn, wrapped, "nil sweeper must return the conn unchanged")
assert.NoError(t, dial.Ctx().Err(), "nil sweeper must not cancel the dial context")
assert.Equal(t, 0, sweeper.sweepAll(), "nil sweeper closes nothing")
}
// wrap registers conn with the sweeper through a completed dial.
func wrap(t *testing.T, sweeper *Sweeper, conn net.Conn) net.Conn {
t.Helper()
dial := sweeper.StartDial(context.Background())
defer dial.Release()
wrapped, err := dial.WrapConn(conn)
require.NoError(t, err)
return wrapped
}
// connPair dials a loopback TCP connection and keeps the accepted peer open
// until the test ends: a peer that closed early would make the connection
// unreadable on its own, so a read error after the sweep would prove nothing.
func connPair(t *testing.T) net.Conn {
t.Helper()
l, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
t.Cleanup(func() {
if err := l.Close(); err != nil {
t.Logf("listener close error: %v", err)
}
})
accepted := make(chan net.Conn, 1)
go func() {
conn, err := l.Accept()
if err != nil {
close(accepted)
return
}
accepted <- conn
}()
conn, err := net.Dial("tcp", l.Addr().String())
require.NoError(t, err)
peer, ok := <-accepted
require.True(t, ok, "listener must accept the dialed connection")
t.Cleanup(func() {
if err := peer.Close(); err != nil {
t.Logf("peer close error: %v", err)
}
})
return conn
}
// sweepAll cuts every registration regardless of generation.
func (s *Sweeper) sweepAll() int {
return s.sweep(math.MaxUint64)
}

View File

@@ -1,39 +0,0 @@
package netsweep
import (
"time"
"github.com/cenkalti/backoff/v4"
"github.com/netbirdio/netbird/client/netstate"
)
const quickRetryDelay = 200 * time.Millisecond
type quickRetryBackoff struct {
backoff.BackOff
sweeper *Sweeper
netState *netstate.State
used bool
}
func newQuickRetryBackoff(bo backoff.BackOff, sweeper *Sweeper, netState *netstate.State) *quickRetryBackoff {
return &quickRetryBackoff{
BackOff: bo,
sweeper: sweeper,
netState: netState,
}
}
func (b *quickRetryBackoff) NextBackOff() time.Duration {
if !b.used && b.sweeper.markedRecently() && b.netState.IsOnline() {
b.used = true
return quickRetryDelay
}
return b.BackOff.NextBackOff()
}
func (b *quickRetryBackoff) Reset() {
b.used = false
b.BackOff.Reset()
}

View File

@@ -1,58 +0,0 @@
package netsweep
import (
"context"
"testing"
"time"
"github.com/cenkalti/backoff/v4"
"github.com/stretchr/testify/assert"
)
func TestQuickRetryAfterRecentMark(t *testing.T) {
sweeper := New()
sweeper.MarkNetworkChange()
slow := backoff.NewConstantBackOff(5 * time.Second)
bo := sweeper.QuickRetryBackoff(context.Background(), slow, nil)
assert.Equal(t, quickRetryDelay, bo.NextBackOff(), "first retry after a mark must be quick")
assert.Equal(t, 5*time.Second, bo.NextBackOff(), "second retry must fall back to the wrapped backoff")
bo.Reset()
assert.Equal(t, quickRetryDelay, bo.NextBackOff(), "reset must re-arm the quick retry")
}
func TestQuickRetryWithoutMarkKeepsSpread(t *testing.T) {
sweeper := New()
slow := backoff.NewConstantBackOff(5 * time.Second)
bo := sweeper.QuickRetryBackoff(context.Background(), slow, nil)
assert.Equal(t, 5*time.Second, bo.NextBackOff(), "without a mark the wrapped backoff decides")
sweeper.mu.Lock()
sweeper.lastMark = time.Now().Add(-recentMarkWindow)
sweeper.mu.Unlock()
assert.Equal(t, 5*time.Second, bo.NextBackOff(), "a stale mark must not trigger the quick retry")
}
func TestQuickRetryNilSweeperPassthrough(t *testing.T) {
var sweeper *Sweeper
slow := backoff.NewConstantBackOff(5 * time.Second)
bo := sweeper.QuickRetryBackoff(context.Background(), slow, nil)
assert.Equal(t, backoff.BackOff(slow), bo, "nil sweeper must return the backoff unchanged")
}
func TestQuickRetryHonorsContext(t *testing.T) {
sweeper := New()
sweeper.MarkNetworkChange()
ctx, cancel := context.WithCancel(context.Background())
cancel()
bo := sweeper.QuickRetryBackoff(ctx, backoff.NewConstantBackOff(time.Millisecond), nil)
assert.Equal(t, backoff.Stop, bo.NextBackOff(), "cancelled context must stop the retry loop")
}

View File

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

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, nil)
networkMapController := controller.NewController(context.Background(), store, metrics, peersUpdateManager, requestBuffer, server.MockIntegratedValidator{}, settingsMockManager, "netbird.selfhosted", port_forwarding.NewControllerMock(), manager.NewEphemeralManager(store, peersManager), config)
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

@@ -313,23 +313,21 @@ func Dial(ctx context.Context, addr, user string, opts DialOptions) (*Client, er
// dialSSH establishes an SSH connection without JWT authentication
func dialSSH(ctx context.Context, network, addr string, config *ssh.ClientConfig) (*Client, error) {
if config.Timeout > 0 {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, config.Timeout)
defer cancel()
}
dialer := &net.Dialer{}
conn, err := dialer.DialContext(ctx, network, addr)
if err != nil {
return nil, fmt.Errorf("dial %s: %w", addr, err)
}
client, err := nbssh.Handshake(ctx, conn, addr, config)
clientConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config)
if err != nil {
return nil, err
if closeErr := conn.Close(); closeErr != nil {
log.Debugf("connection close after handshake failure: %v", closeErr)
}
return nil, fmt.Errorf("ssh handshake: %w", err)
}
client := ssh.NewClient(clientConn, chans, reqs)
return &Client{
client: client,
}, nil

View File

@@ -12,8 +12,6 @@ import (
log "github.com/sirupsen/logrus"
"golang.org/x/crypto/ssh"
"golang.org/x/term"
nbssh "github.com/netbirdio/netbird/client/ssh"
)
func (c *Client) setupTerminalMode(ctx context.Context, session *ssh.Session) error {
@@ -84,7 +82,37 @@ func (c *Client) setupTerminal(session *ssh.Session, fd int) error {
return fmt.Errorf("get terminal size: %w", err)
}
modes := nbssh.DefaultTerminalModes
modes := ssh.TerminalModes{
ssh.ECHO: 1,
ssh.TTY_OP_ISPEED: 14400,
ssh.TTY_OP_OSPEED: 14400,
// Ctrl+C
ssh.VINTR: 3,
// Ctrl+\
ssh.VQUIT: 28,
// Backspace
ssh.VERASE: 127,
// Ctrl+U
ssh.VKILL: 21,
// Ctrl+D
ssh.VEOF: 4,
ssh.VEOL: 0,
ssh.VEOL2: 0,
// Ctrl+Q
ssh.VSTART: 17,
// Ctrl+S
ssh.VSTOP: 19,
// Ctrl+Z
ssh.VSUSP: 26,
// Ctrl+O
ssh.VDISCARD: 15,
// Ctrl+R
ssh.VREPRINT: 18,
// Ctrl+W
ssh.VWERASE: 23,
// Ctrl+V
ssh.VLNEXT: 22,
}
terminal := os.Getenv("TERM")
if terminal == "" {

View File

@@ -10,8 +10,6 @@ import (
log "github.com/sirupsen/logrus"
"golang.org/x/crypto/ssh"
nbssh "github.com/netbirdio/netbird/client/ssh"
)
const (
@@ -82,14 +80,28 @@ func (c *Client) setupTerminalMode(_ context.Context, session *ssh.Session) erro
w, h := c.getWindowsConsoleSize()
modes := ssh.TerminalModes{
ssh.ICRNL: 1,
ssh.OPOST: 1,
ssh.ONLCR: 1,
ssh.ISIG: 1,
ssh.ICANON: 1,
}
for mode, value := range nbssh.DefaultTerminalModes {
modes[mode] = value
ssh.ECHO: 1,
ssh.TTY_OP_ISPEED: 14400,
ssh.TTY_OP_OSPEED: 14400,
ssh.ICRNL: 1,
ssh.OPOST: 1,
ssh.ONLCR: 1,
ssh.ISIG: 1,
ssh.ICANON: 1,
ssh.VINTR: 3, // Ctrl+C
ssh.VQUIT: 28, // Ctrl+\
ssh.VERASE: 127, // Backspace
ssh.VKILL: 21, // Ctrl+U
ssh.VEOF: 4, // Ctrl+D
ssh.VEOL: 0,
ssh.VEOL2: 0,
ssh.VSTART: 17, // Ctrl+Q
ssh.VSTOP: 19, // Ctrl+S
ssh.VSUSP: 26, // Ctrl+Z
ssh.VDISCARD: 15, // Ctrl+O
ssh.VWERASE: 23, // Ctrl+W
ssh.VLNEXT: 22, // Ctrl+V
ssh.VREPRINT: 18, // Ctrl+R
}
if err := session.RequestPty("xterm-256color", h, w, modes); err != nil {

View File

@@ -35,19 +35,6 @@ type HostKeyVerifier interface {
VerifySSHHostKey(peerAddress string, key []byte) error
}
// PeerKeyLookup returns the stored SSH host key for a peer address.
type PeerKeyLookup func(peerAddress string) ([]byte, bool)
// VerifySSHHostKey implements HostKeyVerifier by looking up the stored key
// and comparing it against the presented key.
func (l PeerKeyLookup) VerifySSHHostKey(peerAddress string, presentedKey []byte) error {
storedKey, found := l(peerAddress)
if !found {
return ErrPeerNotFound
}
return VerifyHostKey(storedKey, presentedKey, peerAddress)
}
// DaemonHostKeyVerifier implements HostKeyVerifier using the NetBird daemon
type DaemonHostKeyVerifier struct {
client proto.DaemonServiceClient

View File

@@ -1,45 +0,0 @@
package ssh
import (
"context"
"fmt"
"io"
"net"
"time"
log "github.com/sirupsen/logrus"
"golang.org/x/crypto/ssh"
)
// Handshake runs the SSH client handshake on an already dialed conn and
// returns the resulting client. Dialing bounds only the TCP establishment;
// without a deadline on the socket a peer that accepts and then goes silent
// blocks the handshake forever, so the context deadline is applied to conn
// for the duration of the handshake. conn is closed on any error.
func Handshake(ctx context.Context, conn net.Conn, addr string, config *ssh.ClientConfig) (*ssh.Client, error) {
if deadline, ok := ctx.Deadline(); ok {
if err := conn.SetDeadline(deadline); err != nil {
closeHandshake(conn, "conn after deadline error")
return nil, fmt.Errorf("set handshake deadline: %w", err)
}
}
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config)
if err != nil {
closeHandshake(conn, "conn after handshake error")
return nil, fmt.Errorf("ssh handshake: %w", err)
}
if err := conn.SetDeadline(time.Time{}); err != nil {
closeHandshake(sshConn, "ssh conn after deadline clear error")
return nil, fmt.Errorf("clear handshake deadline: %w", err)
}
return ssh.NewClient(sshConn, chans, reqs), nil
}
func closeHandshake(c io.Closer, label string) {
if err := c.Close(); err != nil {
log.Debugf("ssh: close %s: %v", label, err)
}
}

View File

@@ -610,10 +610,13 @@ func (p *SSHProxy) dialBackend(ctx context.Context, addr, user, jwtToken string)
return nil, fmt.Errorf("connect to server: %w", err)
}
handshakeCtx, cancel := context.WithTimeout(ctx, sshHandshakeTimeout)
defer cancel()
clientConn, chans, reqs, err := cryptossh.NewClientConn(conn, addr, config)
if err != nil {
_ = conn.Close()
return nil, fmt.Errorf("SSH handshake: %w", err)
}
return nbssh.Handshake(handshakeCtx, conn, addr, config)
return cryptossh.NewClient(clientConn, chans, reqs), nil
}
func (p *SSHProxy) verifyHostKey(hostname string, remote net.Addr, key cryptossh.PublicKey) error {

View File

@@ -1,84 +0,0 @@
package ssh
import (
"fmt"
"io"
log "github.com/sirupsen/logrus"
"golang.org/x/crypto/ssh"
)
// DefaultTerminalModes are the PTY modes used by the interactive terminal clients.
var DefaultTerminalModes = ssh.TerminalModes{
ssh.ECHO: 1,
ssh.TTY_OP_ISPEED: 14400,
ssh.TTY_OP_OSPEED: 14400,
ssh.VINTR: 3, // Ctrl+C
ssh.VQUIT: 28, // Ctrl+\
ssh.VERASE: 127, // Backspace
ssh.VKILL: 21, // Ctrl+U
ssh.VEOF: 4, // Ctrl+D
ssh.VEOL: 0,
ssh.VEOL2: 0,
ssh.VSTART: 17, // Ctrl+Q
ssh.VSTOP: 19, // Ctrl+S
ssh.VSUSP: 26, // Ctrl+Z
ssh.VDISCARD: 15, // Ctrl+O
ssh.VREPRINT: 18, // Ctrl+R
ssh.VWERASE: 23, // Ctrl+W
ssh.VLNEXT: 22, // Ctrl+V
}
// PTYSession is an interactive shell session with a PTY and its I/O pipes.
type PTYSession struct {
Session *ssh.Session
Stdin io.WriteCloser
Stdout io.Reader
Stderr io.Reader
}
// StartPTYSession opens a session on the client, requests an xterm-256color PTY
// with the default terminal modes, wires up the I/O pipes and starts a shell.
// The session is closed on any error.
func StartPTYSession(client *ssh.Client, cols, rows int) (*PTYSession, error) {
session, err := client.NewSession()
if err != nil {
return nil, fmt.Errorf("new session: %w", err)
}
pty, err := setupPTYSession(session, cols, rows)
if err != nil {
if closeErr := session.Close(); closeErr != nil {
log.Debugf("ssh: session close after setup error: %v", closeErr)
}
return nil, err
}
return pty, nil
}
// setupPTYSession requests the PTY, opens the pipes and starts the shell on an
// already created session.
func setupPTYSession(session *ssh.Session, cols, rows int) (*PTYSession, error) {
if err := session.RequestPty("xterm-256color", rows, cols, DefaultTerminalModes); err != nil {
return nil, fmt.Errorf("request pty: %w", err)
}
stdin, err := session.StdinPipe()
if err != nil {
return nil, fmt.Errorf("stdin pipe: %w", err)
}
stdout, err := session.StdoutPipe()
if err != nil {
return nil, fmt.Errorf("stdout pipe: %w", err)
}
stderr, err := session.StderrPipe()
if err != nil {
return nil, fmt.Errorf("stderr pipe: %w", err)
}
if err := session.Shell(); err != nil {
return nil, fmt.Errorf("start shell: %w", err)
}
return &PTYSession{Session: session, Stdin: stdin, Stdout: stdout, Stderr: stderr}, nil
}

View File

@@ -30,11 +30,6 @@ func GetInfo(ctx context.Context) *Info {
kernelVersion = osInfo[2]
}
addrs, err := networkAddresses()
if err != nil {
log.Warnf("discover network addresses: %s", err)
}
gio := &Info{
GoOS: runtime.GOOS,
Kernel: kernel,
@@ -46,7 +41,6 @@ func GetInfo(ctx context.Context) *Info {
NetbirdVersion: version.NetbirdVersion(),
UIVersion: extractUIVersion(ctx),
KernelVersion: kernelVersion,
NetworkAddresses: addrs,
SystemSerialNumber: serial(),
SystemProductName: productModel(),
SystemManufacturer: productManufacturer(),

View File

@@ -1,4 +1,4 @@
//go:build !ios && !android
//go:build !ios
package system

View File

@@ -1,89 +0,0 @@
package system
import (
"net/netip"
"strings"
)
var iFaceDiscover IFaceDiscover
type IFaceDiscover interface {
IFaces() (string, error)
}
// SetIFaceDiscover configures the Android interface discovery provider.
func SetIFaceDiscover(discover IFaceDiscover) {
iFaceDiscover = discover
}
func networkAddresses() ([]NetworkAddress, error) {
if iFaceDiscover == nil {
return nil, nil
}
ifaces, err := iFaceDiscover.IFaces()
if err != nil {
return nil, err
}
var netAddresses []NetworkAddress
for _, line := range strings.Split(ifaces, "\n") {
addresses, ok := interfaceAddresses(line)
if !ok {
continue
}
for _, address := range addresses {
netAddr, ok := toNetworkAddress(address)
if !ok {
continue
}
if isDuplicated(netAddresses, netAddr) {
continue
}
netAddresses = append(netAddresses, netAddr)
}
}
return netAddresses, nil
}
func interfaceAddresses(line string) ([]string, bool) {
parts := strings.Split(line, "|")
if len(parts) != 2 {
return nil, false
}
flags := strings.Fields(parts[0])
if len(flags) != 8 {
return nil, false
}
up, loopback := flags[3], flags[5]
if up != "true" || loopback == "true" {
return nil, false
}
return strings.Fields(parts[1]), true
}
func toNetworkAddress(address string) (NetworkAddress, bool) {
prefix, err := netip.ParsePrefix(address)
if err != nil {
return NetworkAddress{}, false
}
if prefix.Addr().Is4In6() {
if prefix.Bits() < 96 {
return NetworkAddress{}, false
}
prefix = netip.PrefixFrom(prefix.Addr().Unmap(), prefix.Bits()-96)
}
ip := prefix.Addr()
if ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsMulticast() {
return NetworkAddress{}, false
}
return NetworkAddress{NetIP: prefix}, true
}
func isDuplicated(addresses []NetworkAddress, addr NetworkAddress) bool {
for _, duplicated := range addresses {
if duplicated.NetIP == addr.NetIP {
return true
}
}
return false
}

View File

@@ -1,4 +1,4 @@
//go:build !ios && !android
//go:build !ios
package system

View File

@@ -80,12 +80,13 @@ func (c *Client) Connect(host string, port int, username, jwtToken string, ipVer
return fmt.Errorf("dial %s: %w", addr, err)
}
sshClient, err := nbssh.Handshake(ctx, conn, addr, config)
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, config)
if err != nil {
return err
closeWithLog(conn, "connection after handshake error")
return fmt.Errorf("SSH handshake: %w", err)
}
c.sshClient = sshClient
c.sshClient = ssh.NewClient(sshConn, chans, reqs)
logrus.Infof("SSH: Connected to %s", addr)
return nil
@@ -118,26 +119,57 @@ func (c *Client) getAuthMethods(jwtToken string) ([]ssh.AuthMethod, error) {
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
}
// StartSession starts an SSH session with PTY. It holds the client lock for
// the whole startup so Close cannot tear the client down mid-setup and the
// new session cannot be installed into an already closed client.
// StartSession starts an SSH session with PTY
func (c *Client) StartSession(cols, rows int) error {
c.mu.Lock()
defer c.mu.Unlock()
if c.sshClient == nil {
return fmt.Errorf("SSH client not connected")
}
pty, err := nbssh.StartPTYSession(c.sshClient, cols, rows)
session, err := c.sshClient.NewSession()
if err != nil {
return err
return fmt.Errorf("create session: %w", err)
}
c.session = pty.Session
c.stdin = pty.Stdin
c.stdout = pty.Stdout
c.stderr = pty.Stderr
c.mu.Lock()
defer c.mu.Unlock()
c.session = session
modes := ssh.TerminalModes{
ssh.ECHO: 1,
ssh.TTY_OP_ISPEED: 14400,
ssh.TTY_OP_OSPEED: 14400,
ssh.VINTR: 3,
ssh.VQUIT: 28,
ssh.VERASE: 127,
}
if err := session.RequestPty("xterm-256color", rows, cols, modes); err != nil {
closeWithLog(session, "session after PTY error")
return fmt.Errorf("PTY request: %w", err)
}
c.stdin, err = session.StdinPipe()
if err != nil {
closeWithLog(session, "session after stdin error")
return fmt.Errorf("get stdin: %w", err)
}
c.stdout, err = session.StdoutPipe()
if err != nil {
closeWithLog(session, "session after stdout error")
return fmt.Errorf("get stdout: %w", err)
}
c.stderr, err = session.StderrPipe()
if err != nil {
closeWithLog(session, "session after stderr error")
return fmt.Errorf("get stderr: %w", err)
}
if err := session.Shell(); err != nil {
closeWithLog(session, "session after shell error")
return fmt.Errorf("start shell: %w", err)
}
logrus.Info("SSH: Session started with PTY")
return nil

9
go.mod
View File

@@ -72,18 +72,17 @@ require (
github.com/hashicorp/go-multierror v1.1.1
github.com/hashicorp/go-secure-stdlib/base62 v0.1.2
github.com/hashicorp/go-version v1.7.0
github.com/jackc/pgx/v5 v5.10.0
github.com/jackc/pgx/v5 v5.5.5
github.com/libdns/route53 v1.5.0
github.com/libp2p/go-nat v0.2.0
github.com/libp2p/go-netroute v0.4.0
github.com/lrh3321/ipset-go v0.0.0-20250619021614-54a0a98ace81
github.com/magefile/mage v1.17.2
github.com/mdlayher/socket v0.5.1
github.com/mdp/qrterminal/v3 v3.2.1
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-20260803100840-78e79ba20f87
github.com/netbirdio/management-integrations/integrations v0.0.0-20260416123949-2355d972be42
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
@@ -237,8 +236,8 @@ require (
github.com/huin/goupnp v1.2.0 // indirect
github.com/inconshreveable/mousetrap v1.1.0 // indirect
github.com/jackc/pgpassfile v1.0.0 // indirect
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
github.com/jackc/puddle/v2 v2.2.2 // indirect
github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a // indirect
github.com/jackc/puddle/v2 v2.2.1 // indirect
github.com/jackpal/go-nat-pmp v1.0.2 // indirect
github.com/jinzhu/inflection v1.0.0 // indirect
github.com/jinzhu/now v1.1.5 // indirect

18
go.sum
View File

@@ -341,12 +341,12 @@ github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0=
github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4=
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a h1:bbPeKD0xmW/Y25WS6cokEszi5g+S0QxI/d45PkRi7Nk=
github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
github.com/jackc/pgx/v5 v5.5.5 h1:amBjrZVmksIdNjxGW/IiIMzxMKZFelXbUoPNb+8sjQw=
github.com/jackc/pgx/v5 v5.5.5/go.mod h1:ez9gk+OAat140fv9ErkZDYFWmXLfV+++K0uAOiwgm1A=
github.com/jackc/puddle/v2 v2.2.1 h1:RhxXJtFG022u4ibrCSMSiu5aOq1i77R3OHKNJj77OAk=
github.com/jackc/puddle/v2 v2.2.1/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/jackpal/go-nat-pmp v1.0.2 h1:KzKSgb7qkJvOUTqYl9/Hg/me3pWgBmERKrTGD7BdWus=
github.com/jackpal/go-nat-pmp v1.0.2/go.mod h1:QPH045xvCAeXUZOxsnwmrtiCoxIr9eob+4orBN1SBKc=
github.com/jcmturner/aescts/v2 v2.0.0 h1:9YKLH6ey7H4eDBXW8khjYslgyqG2xZikXP0EQFKrle8=
@@ -415,8 +415,6 @@ github.com/lrh3321/ipset-go v0.0.0-20250619021614-54a0a98ace81 h1:J56rFEfUTFT9j9
github.com/lrh3321/ipset-go v0.0.0-20250619021614-54a0a98ace81/go.mod h1:RD8ML/YdXctQ7qbcizZkw5mZ6l8Ogrl1dodBzVJduwI=
github.com/lufia/plan9stats v0.0.0-20240513124658-fba389f38bae h1:dIZY4ULFcto4tAFlj1FYZl8ztUZ13bdq+PLY+NOfbyI=
github.com/lufia/plan9stats v0.0.0-20240513124658-fba389f38bae/go.mod h1:ilwx/Dta8jXAgpFYFvSWEMwxmbWXyiUHkd5FwyKhb5k=
github.com/magefile/mage v1.17.2 h1:fyXVu1eadI8Ap1HCCNgEhJ5McIWiYhLR8uol64ZZc40=
github.com/magefile/mage v1.17.2/go.mod h1:Yj51kqllmsgFpvvSzgrZPK9WtluG3kUhFaBUVLo4feA=
github.com/magiconair/properties v1.8.10 h1:s31yESBquKXCV9a/ScB3ESkOjUYYv+X0rg8SYxI99mE=
github.com/magiconair/properties v1.8.10/go.mod h1:Dhd985XPs7jluiymwWYZ0G4Z61jb3vdS329zhj2hYo0=
github.com/matryer/is v1.4.1 h1:55ehd8zaGABKLXQUe2awZ99BD/PTc2ls+KV/dXphgEQ=
@@ -484,8 +482,8 @@ 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-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/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/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=

View File

@@ -1,58 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"testing"
"time"
"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 := conn(t, ctx).GetAccountSettings(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 = conn(t, ctx).GetAccountSettings(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 = conn(t, ctx).GetAccountSettings(ctx, "account-3")
assert.NoError(t, err)
assert.Equal(t, accountSettings, nmdata.AccountSettingsInfo{})
}

View File

@@ -1,53 +0,0 @@
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','"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','"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','"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

@@ -1,25 +0,0 @@
//go:build integration
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 := conn(t, ctx).GetDnsSettings(ctx, "account-1")
assert.NoError(t, err)
assert.Equal(t, settings, nmdata.DNSSettings{
DisabledManagementGroups: []string{"disabled-group-1", "disabled-group-2"},
})
settings, err = conn(t, ctx).GetDnsSettings(ctx, "account-2")
assert.NoError(t, err)
assert.Equal(t, settings, nmdata.DNSSettings{})
}

View File

@@ -1,62 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"testing"
"github.com/miekg/dns"
"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 := conn(t, ctx).GetAppliedZoneCandidates(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

@@ -1,39 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"database/sql"
"testing"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"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 := conn(t, ctx).GetDomains(ctx, "account-1")
assert.NoError(t, err)
assert.Len(t, domains, 2)
assert.Contains(t, domains, networkmapdb.Domain{
Domain: sql.NullString{String: "test-1.com", Valid: true},
TargetCluster: sql.NullString{String: "target-1.cluster.local", Valid: true},
})
assert.Contains(t, domains, networkmapdb.Domain{
Domain: sql.NullString{String: "test-2.com", Valid: true},
TargetCluster: sql.NullString{String: "target-2.cluster.local", Valid: true},
})
}

View File

@@ -1,54 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"testing"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestGetGroups(t *testing.T) {
ctx := context.TODO()
groups, resourceToGroupIdx, err := conn(t, ctx).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()
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')")
groups, _, err := conn(t, ctx).GetGroups(ctx, "random-id")
assert.NoError(t, err)
require.Len(t, groups, 1)
assert.NotEmpty(t, groups[0].PublicID)
}

View File

@@ -1,99 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
_ "embed"
"os"
"testing"
"time"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql"
networkmap_sqlite "github.com/netbirdio/netbird/management/internals/network_map_db/sqlite"
"github.com/netbirdio/netbird/management/server/types"
)
//go:embed base_data.sql
var baseData string
var (
pgstore *networkmap_pgsql.PgStore
sqlitestore *networkmap_sqlite.SqliteStore
engine string
)
func TestMain(m *testing.M) {
var cleanup func()
kind, _ := os.LookupEnv("NETBIRD_STORE_ENGINE")
switch kind {
case string(types.PostgresStoreEngine):
engine = string(types.PostgresStoreEngine)
pgstore, cleanup = createPGTestStore(baseData)
pgstore.UsingTimeZone(time.UTC)
case "", string(types.SqliteStoreEngine):
engine = string(types.SqliteStoreEngine)
sqlitestore, cleanup = createSqliteTestStore(baseData)
default:
log.Fatalf("unsupported db '%s' in NETBIRD_STORE_ENGINE env var", kind)
}
code := m.Run()
cleanup()
os.Exit(code)
}
func conn(t *testing.T, ctx context.Context) networkmapdb.NetworkMapDBStoreConn {
t.Helper()
switch engine {
case string(types.PostgresStoreEngine):
c, err := pgstore.Pool.Acquire(ctx)
assert.NoError(t, err)
return pgstore.UsingConnection(c.Conn())
case string(types.SqliteStoreEngine):
return sqlitestore.UsingConn()
}
log.Fatalf("unknown db engine kind %s", engine)
return nil
}
func store(t *testing.T) networkmapdb.NetworkMapDBStore {
t.Helper()
switch engine {
case string(types.PostgresStoreEngine):
return pgstore
case string(types.SqliteStoreEngine):
return sqlitestore
}
log.Fatalf("unknown db engine kind %s", engine)
return nil
}
func execQuery(t *testing.T, ctx context.Context, q string) {
t.Helper()
switch engine {
case string(types.PostgresStoreEngine):
_, err := pgstore.Pool.Exec(ctx, q)
assert.NoError(t, err)
case string(types.SqliteStoreEngine):
_, err := sqlitestore.Db.ExecContext(ctx, q)
assert.NoError(t, err)
}
}
// use to parse time in time.RFC3339Nano format
// returns the time in the UTC time zone
func mustParseTime(t string) *time.Time {
tt, err := time.Parse(time.RFC3339Nano, t)
if err != nil {
panic(err)
}
utc := tt.UTC()
return &utc
}

View File

@@ -1,61 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"net/netip"
"testing"
"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 := conn(t, ctx).GetNameServerGroups(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

@@ -1,98 +0,0 @@
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-33','network-331','{"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, '["33-group-one-resource-id"]', false,
false, 'disabled', false, false);
insert into groups (id, account_id, name, resources, public_id) VALUES('33-group-one-resource-id','account-33','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('33-group-two-resources-id','account-33','group-2-name', '[{"ID":"subnet-id-1","Type":"subnet"}, {"ID":"host-id-2","Type":"host"}]','33-group-two-resources-id-public');
insert into groups (id, account_id, name, resources, public_id) VALUES('33-group-no-resources-id','account-33','group-3-name', null,'33-group-no-resources-id-public');
insert into group_peers (account_id, peer_id, group_id) VALUES('account-33','peer-id-331','33-group-one-resource-id');
insert into group_peers (account_id, peer_id, group_id) VALUES('account-33','peer-id-332','33-group-two-resources-id');
insert into group_peers (account_id, peer_id, group_id) VALUES('account-33','peer-id-333','33-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-331','account-33','key-331','ssh-key-1','peer-1','["extra-peer-1"]','user-id-1',true,true,'2026-08-06 13:25:59.12999','"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-332','account-33','key-332','ssh-key-2','peer-2','["extra-peer-2"]','user-id-2',true,true,'2026-08-06 14:25:59.12999','"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-333','account-33','key-333','ssh-key-3','peer-3','["extra-peer-3"]','user-id-3',true,true,'2026-08-06 12:25:59.12999','"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"');
insert into zones (id, account_id, domain, enable_search_domain, distribution_groups)
VALUES('zone-331','account-33','test-331.com',true,'["33-group-one-resource-id"]');
insert into records (id, account_id, zone_id, name, type, ttl, content)
VALUES('record-331','account-33','zone-331','test.test-331.com','A',1800,'1.1.1.1');
insert into records (id, account_id, zone_id, name, type, ttl, content)
VALUES('record-332','account-33','zone-331','test2.test-331.com','A',1800,'1.1.1.2');
insert into domains (id, account_id, domain, target_cluster)
VALUES('domain-331','account-33','test-331.com','target-1.cluster.local');
insert into name_server_groups (id, public_id, name, description, name_servers, groups, domains, enabled, search_domains_enabled, "primary", account_id)
VALUES('nsgroup-331','nsgroup-1-public','nsgroup-1','nsgroup-1','[{"IP":"192.168.31.2","NSType":1,"Port":53}]','["33-group-one-resource-id"]','["test-1.com"]',TRUE,FALSE,TRUE,'account-33');
insert into name_server_groups (id, public_id, name, description, name_servers, groups, domains, enabled, search_domains_enabled,"primary",account_id)
VALUES('nsgroup-332','nsgroup-2-public','nsgroup-2','nsgroup-2','[{"IP":"192.168.32.3","NSType":1,"Port":53}]','["33-group-one-resource-id","33-group-no-resources-id"]','["test-1.com","test-2.com"]',TRUE,FALSE,TRUE,'account-33');
insert into network_resources (id, account_id, network_id, public_id, name, description, type, domain, prefix, enabled)
VALUES('net-resource-331','account-33','network-331','net-resource-public-1','network-resource-1','network-resource-1','subnet','','"10.0.0.0/16"',TRUE);
insert into network_resources (id, account_id, network_id, public_id, name, description, type, domain, prefix, enabled)
VALUES('net-resource-332','account-33','network-332','net-resource-public-2','network-resource-2','network-resource-2','domain','test.com','',TRUE);
insert into network_routers (id, account_id, public_id, peer, network_id, masquerade, metric, enabled, peer_groups)
VALUES('test-nr-id-331','account-33','public-id-1','peer-id-331','network-id-1',TRUE,999,TRUE,'["33-group-one-resource-id"]');
insert into network_routers (id, account_id, public_id, peer, network_id, masquerade, metric, enabled, peer_groups)
VALUES('test-nr-id-332','account-33','public-id-2','','network-id-2',TRUE,333,TRUE,'["33-group-two-resources-id","33-group-no-resources-id"]');
insert into networks (id, account_id, public_id) VALUES('network-331','account-33','network-1-public');
insert into networks (id, account_id, public_id) VALUES('network-332','account-33','network-2-public');
insert into policies (id, public_id, account_id, enabled, source_posture_checks)
values('policy-331','policy-1-public','account-33',true,'["posture-checks-1","posture-checks-2"]');
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-331-rule-1','policy-331',true,'accept','tcp',true,'["33-group-one-resource-id","33-group-two-resources-id"]','["33-group-one-resource-id","33-group-two-resources-id"]',
'{"ID":"host-id-1","Type":"host"}','{"ID":"domain-331","Type":"domain"}','["8080","8443"]', '[{"Start":8080,"End":8090}]',
'{"33-group-one-resource-id":["user-1", "user-2"]}','user-3');
insert into posture_checks (id, account_id, public_id, checks)
VALUES('posturecheck-331','account-33','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"]}}');
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-331','account-33','route-1-public','"172.0.0.0/16"','["test-1.com"]',true,'route-331-net-id','route-1',
'peer-id-331','["33-group-one-resource-id"]',1,true,9999,true,
'["33-group-one-resource-id"]','["33-group-one-resource-id"]',false);
insert into services (id, account_id, enabled, private, access_groups, proxy_cluster, domain)
values('service-331','account-33',true,true,'["33-group-one-resource-id"]','test-1.com','test-332.com');

View File

@@ -1,516 +0,0 @@
{
"Peers": {
"peer-id-331": {
"ID": "peer-id-331",
"Key": "key-331",
"SSHKey": "ssh-key-1",
"DNSLabel": "peer-1",
"UserID": "user-id-1",
"SSHEnabled": true,
"LoginExpirationEnabled": true,
"LastLogin": "2026-08-06T13:25:59.12999Z",
"IP": "10.10.10.1",
"IPv6": "fdf4:ba80:6aa5:89f1:44d7:8701:8699:4940",
"RequiresApproval": false,
"ExtraDNSLabels": [
"extra-peer-1"
],
"Meta": {
"WtVersion": "0.76.0",
"GoOS": "linux",
"OSVersion": "26.4.1",
"KernelVersion": "6.8.0-134-generic",
"NetworkAddresses": [
{
"NetIP": "fe80::8b4c:973f:a76b:3771/64"
},
{
"NetIP": "192.168.16.1/20"
}
],
"Files": [
{
"Path": "/usr/bin/netbird",
"ProcessIsRunning": false
}
],
"Capabilities": [
1,
2
],
"Flags": {
"ServerSSHAllowed": true,
"DisableIPv6": false
},
"SyncMessageVersion": 1
},
"ProxyMeta": {
"Embedded": true
},
"Location": {
"CountryCode": "DE",
"CityName": "Berlin",
"ConnectionIP": "46.201.148.187"
}
},
"peer-id-332": {
"ID": "peer-id-332",
"Key": "key-332",
"SSHKey": "ssh-key-2",
"DNSLabel": "peer-2",
"UserID": "user-id-2",
"SSHEnabled": true,
"LoginExpirationEnabled": true,
"LastLogin": "2026-08-06T14:25:59.12999Z",
"IP": "10.10.100.1",
"IPv6": "fdf5:ba80:6aa5:89f1:44d7:8701:8699:4940",
"RequiresApproval": false,
"ExtraDNSLabels": [
"extra-peer-2"
],
"Meta": {
"WtVersion": "0.76.1",
"GoOS": "linux",
"OSVersion": "26.4.2",
"KernelVersion": "6.8.0-135-generic",
"NetworkAddresses": [
{
"NetIP": "fe81::8b4c:973f:a76b:3771/64"
},
{
"NetIP": "192.168.17.1/20"
}
],
"Files": [
{
"Path": "/usr/bin/netbird",
"ProcessIsRunning": false
}
],
"Capabilities": [
1,
2
],
"Flags": {
"ServerSSHAllowed": true,
"DisableIPv6": false
},
"SyncMessageVersion": 0
},
"ProxyMeta": {
"Embedded": true
},
"Location": {
"CountryCode": "DE",
"CityName": "Berlin",
"ConnectionIP": "46.201.149.187"
}
},
"peer-id-333": {
"ID": "peer-id-333",
"Key": "key-333",
"SSHKey": "ssh-key-3",
"DNSLabel": "peer-3",
"UserID": "user-id-3",
"SSHEnabled": true,
"LoginExpirationEnabled": true,
"LastLogin": "2026-08-06T12:25:59.12999Z",
"IP": "10.10.200.1",
"IPv6": "fdf6:ba80:6aa5:89f1:44d7:8701:8699:4940",
"RequiresApproval": false,
"ExtraDNSLabels": [
"extra-peer-3"
],
"Meta": {
"WtVersion": "0.76.2",
"GoOS": "linux",
"OSVersion": "26.4.3",
"KernelVersion": "6.8.0-136-generic",
"NetworkAddresses": [
{
"NetIP": "fe82::8b4c:973f:a76b:3771/64"
},
{
"NetIP": "192.168.18.1/20"
}
],
"Files": [
{
"Path": "/usr/bin/netbird",
"ProcessIsRunning": false
}
],
"Capabilities": [
1,
2
],
"Flags": {
"ServerSSHAllowed": true,
"DisableIPv6": false
},
"SyncMessageVersion": 1
},
"ProxyMeta": {
"Embedded": true
},
"Location": {
"CountryCode": "DE",
"CityName": "Berlin",
"ConnectionIP": "46.201.150.187"
}
}
},
"Groups": {
"33-group-no-resources-id": {
"ID": "33-group-no-resources-id",
"Name": "group-3-name",
"PublicID": "33-group-no-resources-id-public",
"Peers": null,
"Resources": null
},
"33-group-one-resource-id": {
"ID": "33-group-one-resource-id",
"Name": "group-1-name",
"PublicID": "group-one-resource-id-public",
"Peers": [
"peer-id-331"
],
"Resources": [
{
"ID": "host-id-1",
"Type": "host"
}
]
},
"33-group-two-resources-id": {
"ID": "33-group-two-resources-id",
"Name": "group-2-name",
"PublicID": "33-group-two-resources-id-public",
"Peers": [
"peer-id-332",
"peer-id-333"
],
"Resources": [
{
"ID": "subnet-id-1",
"Type": "subnet"
},
{
"ID": "host-id-2",
"Type": "host"
}
]
}
},
"Policies": [
{
"ID": "policy-331",
"PublicID": "policy-1-public",
"Enabled": true,
"SourcePostureChecks": [
"posture-checks-1",
"posture-checks-2"
],
"Rules": [
{
"ID": "policy-331",
"PolicyID": "policy-331",
"Enabled": true,
"Action": "accept",
"Protocol": "tcp",
"Bidirectional": true,
"Sources": [
"33-group-one-resource-id",
"33-group-two-resources-id"
],
"Destinations": [
"33-group-one-resource-id",
"33-group-two-resources-id"
],
"SourceResource": {
"ID": "host-id-1",
"Type": "host"
},
"DestinationResource": {
"ID": "domain-331",
"Type": "domain"
},
"Ports": [
"8080",
"8443"
],
"PortRanges": [
{
"Start": 8080,
"End": 8090
}
],
"AuthorizedGroups": {
"33-group-one-resource-id": [
"user-1",
"user-2"
]
},
"AuthorizedUser": "user-3"
}
]
}
],
"Routes": [
{
"ID": "route-331",
"AccountID": "account-33",
"PublicID": "route-1-public",
"Network": "172.0.0.0/16",
"Domains": [
"test-1.com"
],
"KeepRoute": true,
"NetID": "route-331-net-id",
"Description": "route-1",
"Peer": "peer-id-331",
"PeerID": "peer-id-331",
"PeerGroups": [
"33-group-one-resource-id"
],
"NetworkType": 1,
"Masquerade": true,
"Metric": 9999,
"Enabled": true,
"Groups": [
"33-group-one-resource-id"
],
"AccessControlGroups": [
"33-group-one-resource-id"
],
"SkipAutoApply": false
}
],
"NameServerGroups": [
{
"ID": "nsgroup-331",
"PublicID": "nsgroup-1-public",
"Name": "nsgroup-1",
"Description": "nsgroup-1",
"NameServers": [
{
"IP": "192.168.31.2",
"NSType": 1,
"Port": 53
}
],
"Groups": [
"33-group-one-resource-id"
],
"Primary": true,
"Domains": [
"test-1.com"
],
"Enabled": true,
"SearchDomainsEnabled": false
},
{
"ID": "nsgroup-332",
"PublicID": "nsgroup-2-public",
"Name": "nsgroup-2",
"Description": "nsgroup-2",
"NameServers": [
{
"IP": "192.168.32.3",
"NSType": 1,
"Port": 53
}
],
"Groups": [
"33-group-one-resource-id",
"33-group-no-resources-id"
],
"Primary": true,
"Domains": [
"test-1.com",
"test-2.com"
],
"Enabled": true,
"SearchDomainsEnabled": false
}
],
"NetworkResources": [
{
"ID": "net-resource-331",
"NetworkID": "network-331",
"AccountID": "account-33",
"PublicID": "net-resource-public-1",
"Name": "network-resource-1",
"Description": "network-resource-1",
"Type": "subnet",
"Address": "",
"Domain": "",
"Prefix": "10.0.0.0/16",
"Enabled": true
},
{
"ID": "net-resource-332",
"NetworkID": "network-332",
"AccountID": "account-33",
"PublicID": "net-resource-public-2",
"Name": "network-resource-2",
"Description": "network-resource-2",
"Type": "domain",
"Address": "",
"Domain": "test.com",
"Prefix": "",
"Enabled": true
}
],
"Network": {
"Identifier": "network-331",
"Net": {
"IP": "100.103.0.0",
"Mask": "//8AAA=="
},
"NetV6": {
"IP": "fdde:e995:fd38:a465::",
"Mask": "//////////8AAAAAAAAAAA=="
},
"Dns": "",
"Serial": 1
},
"DNSSettings": {
"DisabledManagementGroups": [
"disabled-group-1",
"disabled-group-2"
]
},
"AccountSettings": {
"PeerLoginExpirationEnabled": true,
"PeerLoginExpiration": 86400000000000,
"PeerInactivityExpirationEnabled": false,
"PeerInactivityExpiration": 86400000000000,
"DNSDomain": "",
"IPv6EnabledGroups": [
"33-group-one-resource-id"
],
"RoutingPeerDNSResolutionEnabled": false,
"LazyConnectionEnabled": false,
"AutoUpdateVersion": "disabled",
"AutoUpdateAlways": false,
"MetricsPushEnabled": false
},
"PostureChecks": {
"posturecheck-331": {
"ID": "posturecheck-331",
"Checks": {
"NBVersionCheck": {
"MinVersion": "0.25.0"
},
"OSVersionCheck": {
"Android": null,
"Darwin": {
"MinVersion": "12.0"
},
"Ios": null,
"Linux": null,
"Windows": null
},
"GeoLocationCheck": {
"Locations": [
{
"CountryCode": "FI",
"CityName": ""
}
],
"Action": "allow"
},
"PeerNetworkRangeCheck": {
"Action": "deny",
"Ranges": [
"192.168.0.1/24"
]
},
"ProcessCheck": null
}
}
},
"PostureValidation": null,
"AllowedUserIDs": {},
"NetworkXIDToPublicID": {
"network-331": "network-1-public",
"network-332": "network-2-public"
},
"PostureCheckXIDToPublicID": {
"posturecheck-331": "posturecheck-1-public"
},
"ValidatedPeers": {
"peer-id-1": {},
"peer-id-2": {},
"peer-id-3": {}
},
"ResourcePolicies": {},
"Routers": {
"network-id-1": {
"peer-id-331": {
"PublicID": "public-id-1",
"PeerGroups": [
"33-group-one-resource-id"
],
"Masquerade": true,
"Metric": 999,
"Enabled": true
}
},
"network-id-2": {
"peer-id-332": {
"PublicID": "public-id-2",
"PeerGroups": [
"33-group-two-resources-id",
"33-group-no-resources-id"
],
"Masquerade": true,
"Metric": 333,
"Enabled": true
},
"peer-id-333": {
"PublicID": "public-id-2",
"PeerGroups": [
"33-group-two-resources-id",
"33-group-no-resources-id"
],
"Masquerade": true,
"Metric": 333,
"Enabled": true
}
}
},
"GroupIDToUserIDs": {},
"DNSDomain": "",
"ProxyTargetedDomainResourceIDs": {},
"AppliedZoneCandidates": [
{
"DistributionGroups": [
"33-group-one-resource-id"
],
"Zone": {
"Domain": "test-331.com",
"Records": [
{
"Name": "test.test-331.com",
"Type": 1,
"Class": "IN",
"TTL": 1800,
"RData": "1.1.1.1"
},
{
"Name": "test2.test-331.com",
"Type": 1,
"Class": "IN",
"TTL": 1800,
"RData": "1.1.1.2"
}
],
"SearchDomainDisabled": false,
"NonAuthoritative": false
}
}
],
"PrivateServiceCandidates": null
}

View File

@@ -1,73 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
_ "embed"
"encoding/json"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"github.com/golang/mock/gomock"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/management/server/types"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
)
//go:embed network_map_data.sql
var nmapData string
//go:embed network_map_data_golden.json
var goldenNMap string
const EnvUpdateGoldenData = "NMAP_UPDATE_GOLDEN_DATA"
func TestGetNetworkMapData(t *testing.T) {
ctx := context.TODO()
ctrl := gomock.NewController(t)
extraSettingsManager := settings.NewMockManager(ctrl)
extraSettingsManager.EXPECT().GetExtraSettings(gomock.Any(), gomock.Any()).Return(&types.ExtraSettings{}, nil)
peerValidators := integrated_validator.NewMockIntegratedValidator(ctrl)
peerValidators.EXPECT().GetValidatedPeers(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Return(
map[string]struct{}{
"peer-id-1": {},
"peer-id-2": {},
"peer-id-3": {},
}, nil)
storeImpl := networkmapdb.NetworkMapDBStoreImpl{
Store: store(t),
ExtraSettingsManager: extraSettingsManager,
IntegratedPeerValidator: peerValidators,
}
for _, query := range strings.Split(nmapData, ";") {
if err := store(t).Exec(ctx, query); err != nil {
log.Fatalf("error initializing nmap test: %s", err.Error())
}
}
nmap, err := storeImpl.GetNetworkMapData(ctx, "account-33")
assert.NoError(t, err)
serializedNMap, err := json.MarshalIndent(nmap, "", " ")
assert.NoError(t, err)
if _, ok := os.LookupEnv(EnvUpdateGoldenData); ok {
_, filename, _, _ := runtime.Caller(0)
tosavepath := filepath.Join(filepath.Dir(filename), "network_map_data_golden.json")
err = os.WriteFile(tosavepath, serializedNMap, 0644)
assert.NoError(t, err)
goldenNMap = string(serializedNMap)
}
assert.Equal(t, goldenNMap, string(serializedNMap))
}

View File

@@ -1,65 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"net/netip"
"testing"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetNetworkResources(t *testing.T) {
ctx := context.TODO()
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 := conn(t, ctx).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

@@ -1,33 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"testing"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetNetworkRouters(t *testing.T) {
ctx := context.TODO()
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 := conn(t, ctx).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

@@ -1,56 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"encoding/json"
"net"
"testing"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetNetwork(t *testing.T) {
ctx := context.TODO()
network, err := conn(t, ctx).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 = conn(t, ctx).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

@@ -1,26 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"testing"
"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 := conn(t, ctx).GetNetworkXIDToPublicIdMap(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

@@ -1,163 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"net"
"net/netip"
"testing"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetPeers(t *testing.T) {
ctx := context.TODO()
peers, clusterToPeersIdx, err := conn(t, ctx).GetPeers(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

@@ -1,121 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"fmt"
"regexp"
"strings"
"time"
log "github.com/sirupsen/logrus"
"github.com/google/uuid"
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"
"gorm.io/driver/postgres"
"gorm.io/gorm"
)
func createPGTestStore(baseData string) (*networkmap_pgsql.PgStore, func()) {
_, 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()
pgstore, 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 := pgstore.Pool.Exec(ctx, query); err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
}
return pgstore, cleanup
}
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}`)
}

View File

@@ -1,146 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"testing"
"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 := conn(t, ctx).GetPolicies(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: nil,
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: nil,
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

@@ -1,61 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"net/netip"
"testing"
"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 := conn(t, ctx).GetPostureChecks(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

@@ -1,87 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"net/netip"
"testing"
"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 := conn(t, ctx).GetRoutes(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

@@ -1,109 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"database/sql"
"testing"
"github.com/stretchr/testify/assert"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
)
func TestGetPrivateServices(t *testing.T) {
ctx := context.TODO()
execQuery(t, 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')`)
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')`)
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)`)
services, err := conn(t, ctx).GetPrivateServices(ctx, "account-1")
assert.NoError(t, err)
assert.Contains(t, services, networkmapdb.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, networkmapdb.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, networkmapdb.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 TestGetProxyTargetedDomainResourceIDs(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 := conn(t, ctx).GetProxyTargetedDomainResourceIDs(ctx, "account-1")
assert.NoError(t, err)
assert.Equal(t, servtargetedDomains, map[string]struct{}{
"target-1": {},
"target-6": {},
"target-7": {},
})
}

View File

@@ -1,48 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"fmt"
"runtime"
"strings"
networkmap_sqlite "github.com/netbirdio/netbird/management/internals/network_map_db/sqlite"
gormstore "github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
log "github.com/sirupsen/logrus"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
func createSqliteTestStore(baseData string) (*networkmap_sqlite.SqliteStore, func()) {
storeSqliteFileName := ":memory:"
storeStr := fmt.Sprintf("%s?cache=shared", storeSqliteFileName)
if runtime.GOOS == "windows" {
// Vo avoid `The process cannot access the file because it is being used by another process` on Windows
storeStr = storeSqliteFileName
}
db, err := gorm.Open(sqlite.Open(storeStr), &gorm.Config{})
if err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
_, err = gormstore.NewSqlStore(context.TODO(), db, types.SqliteStoreEngine, nil, false)
if err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
sqldb, err := db.DB()
if err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
for _, query := range strings.Split(baseData, ";") {
if _, err := sqldb.Exec(query); err != nil {
log.Fatalf("error initializing db: %s", err.Error())
}
}
return &networkmap_sqlite.SqliteStore{Db: sqldb}, func() {}
}

View File

@@ -1,57 +0,0 @@
//go:build integration
package networkmap_pgsql
import (
"context"
"testing"
"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 := conn(t, ctx).GetAllowedUsers(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

@@ -1,10 +0,0 @@
//mage:multiline
// Set the general description you want to have displayed with mage -l here.
package main
// mg contains helpful utility functions, like Deps
// Default target to run when none is specified
// If not set, running mage will list available targets
//var Default = Integrationtest.All

View File

@@ -1,74 +0,0 @@
package main
import (
"errors"
"strings"
"github.com/magefile/mage/mg"
"github.com/magefile/mage/sh"
)
var defaultcli = []string{"test", "-tags=integration", "-timeout=20m"}
type Integrationtest mg.Namespace
func (i Integrationtest) All(gotestflags *string) error {
var errs []error
if err := i.Api(gotestflags); err != nil {
errs = append(errs, err)
}
if err := i.NmapDb(gotestflags); err != nil {
errs = append(errs, err)
}
if len(errs) > 0 {
return errors.Join(errs...)
}
return nil
}
func (Integrationtest) NmapDb(gotestflags *string) error {
cli := defaultcli
if gotestflags != nil {
cli = append(cli, strings.Split(*gotestflags, " ")...)
}
cli = append(cli, "./integration_tests/management/network_map_db/...")
return sh.RunV("go", cli...)
}
func (Integrationtest) NmapDbPostgres(gotestflags *string) error {
cli := defaultcli
if gotestflags != nil {
cli = append(cli, strings.Split(*gotestflags, " ")...)
}
cli = append(cli, "./integration_tests/management/network_map_db/...")
return sh.RunWithV(map[string]string{"NETBIRD_STORE_ENGINE": "postgres"}, "go", cli...)
}
func (Integrationtest) NmapDbSqlite(gotestflags *string) error {
cli := defaultcli
if gotestflags != nil {
cli = append(cli, strings.Split(*gotestflags, " ")...)
}
cli = append(cli, "./integration_tests/management/network_map_db/...")
return sh.RunWithV(map[string]string{"NETBIRD_STORE_ENGINE": "sqlite"}, "go", cli...)
}
func (Integrationtest) RegenerateNmapGoldenData(gotestflags *string) error {
cli := defaultcli
if gotestflags != nil {
cli = append(cli, strings.Split(*gotestflags, " ")...)
}
cli = append(cli, "./integration_tests/management/network_map_db/...")
return sh.RunWithV(map[string]string{"NMAP_UPDATE_GOLDEN_DATA": "true", "NETBIRD_STORE_ENGINE": "sqlite"}, "go", cli...)
}
func (Integrationtest) Api(gotestflags *string) error {
cli := defaultcli
if gotestflags != nil {
cli = append(cli, strings.Split(*gotestflags, " ")...)
}
cli = append(cli, "./management/server/http/...")
return sh.RunV("go", cli...)
}

View File

@@ -18,7 +18,6 @@ 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"
@@ -31,8 +30,6 @@ 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"
@@ -64,8 +61,6 @@ type Controller struct {
serverSupportedSyncMessageVersion sharedgrpc.SyncMessageVersion
perAccountServerSupportedSyncMessageVersions map[string]sharedgrpc.SyncMessageVersion
nmdataStore *networkmapdb.NetworkMapDBStoreImpl
}
type bufferUpdate struct {
@@ -83,7 +78,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, nmdataStore *networkmapdb.NetworkMapDBStoreImpl) *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) *Controller {
nMetrics, err := newMetrics(metrics.UpdateChannelMetrics())
if err != nil {
log.Fatal(fmt.Errorf("error creating metrics: %w", err))
@@ -104,7 +99,6 @@ 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,
}
}
@@ -153,11 +147,6 @@ 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)
@@ -178,7 +167,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
return nil
}
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validate peers: %v", err)
}
@@ -266,7 +255,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, types.TwinPeer(p), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
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)
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -287,7 +276,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
}
start = time.Now()
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)
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)
c.metrics.CountToSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -305,261 +294,6 @@ 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)
}
nmData.PrecomputePostureValidation()
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
@@ -592,10 +326,6 @@ 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)
@@ -611,7 +341,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, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validate peers: %v", err)
}
@@ -698,7 +428,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, types.TwinPeer(p), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSetting, maps.Keys(peerGroups), dnsFwdPort)
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)
c.metrics.CountToComponentSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -719,7 +449,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
}
start = time.Now()
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)
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)
c.metrics.CountToSyncResponseDuration(time.Since(start))
c.peersUpdateManager.SendUpdate(ctx, p.ID, &network_map.UpdateMessage{
@@ -776,7 +506,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, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return fmt.Errorf("failed to get validated peers: %v", err)
}
@@ -836,7 +566,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, types.TwinPeer(peer), nil, nil, components, proxyNetworkMap, dnsDomain, postureChecks, types.TwinAccountSettings(account.Settings), extraSettings, maps.Keys(peerGroups), dnsFwdPort)
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)
c.peersUpdateManager.SendUpdate(ctx, peer.ID, &network_map.UpdateMessage{
Update: update,
@@ -853,7 +583,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, types.TwinPeer(peer), nil, nil, nmap, dnsDomain, postureChecks, dnsCache, types.TwinAccountSettings(account.Settings), extraSettings, maps.Keys(peerGroups), dnsFwdPort)
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)
c.peersUpdateManager.SendUpdate(ctx, peer.ID, &network_map.UpdateMessage{
Update: update,
@@ -913,11 +643,7 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
if err != nil {
return nil, nil, nil, nil, 0, err
}
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)
return peer, &types.NetworkMapComponents{Network: network.Copy()}, nil, nil, 0, nil
}
account, err := c.requestBuffer.GetAccountWithBackpressure(ctx, accountID)
@@ -927,7 +653,7 @@ func (c *Controller) GetValidatedPeerWithComponents(ctx context.Context, isRequi
c.injectAllProxyPolicies(ctx, account)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return nil, nil, nil, nil, 0, err
}
@@ -964,21 +690,6 @@ 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 {
@@ -1085,15 +796,11 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
}
emptyMap := &types.NetworkMap{
Network: types.TwinNetwork(network),
Network: network.Copy(),
}
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
@@ -1101,7 +808,7 @@ func (c *Controller) GetValidatedPeerWithMap(ctx context.Context, isRequiresAppr
c.injectAllProxyPolicies(ctx, account)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
approvedPeersMap, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return nil, nil, 0, err
}
@@ -1141,21 +848,6 @@ 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 {
@@ -1218,36 +910,20 @@ 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 {
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 {
if len(peers) == 0 {
return int64(network_map.OldForwarderPort)
}
reqVer := semver.Canonical(requiredVersion)
// Check if all peers have the required version or newer
for _, wtVersion := range wtVersions {
for _, peer := range peers {
// Development version is always supported
if version.IsDevelopmentVersion(wtVersion) {
if version.IsDevelopmentVersion(peer.Meta.WtVersion) {
continue
}
peerVersion := semver.Canonical("v" + wtVersion)
peerVersion := semver.Canonical("v" + peer.Meta.WtVersion)
if peerVersion == "" {
// If any peer doesn't have version info, return 0
return int64(network_map.OldForwarderPort)
@@ -1381,7 +1057,7 @@ func (c *Controller) GetNetworkMap(ctx context.Context, peerID string) (*types.N
groups[groupID] = group.Peers
}
validatedPeers, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, types.TwinGroups(maps.Values(account.Groups)), types.TwinPeers(maps.Values(account.Peers)), account.Settings.Extra)
validatedPeers, err := c.integratedPeerValidator.GetValidatedPeers(ctx, account.Id, maps.Values(account.Groups), maps.Values(account.Peers), account.Settings.Extra)
if err != nil {
return nil, err
}

View File

@@ -1,380 +0,0 @@
package nmaptest
import (
"bytes"
"cmp"
"fmt"
"slices"
"sort"
"strconv"
"strings"
"github.com/netbirdio/netbird/shared/management/proto"
)
// normalizeIDSpace replaces policy and route identifiers with positional
// placeholders so a comparison can reach everything else.
//
// This exists only because the envelope round-trip currently substitutes each
// internal xid with the object's public id, which is a tracked defect and not a
// licence to differ: those identifiers reach the server again inside flow
// events, which resolve them by internal id, so the substitution silently
// breaks flow attribution for component-format peers. TestIDSpaceMatches
// asserts the equality that must eventually hold; this erasure keeps the other
// 40-odd cases reporting on semantics meanwhile. When the id space is unified,
// delete this and the calls to it — every case should still pass.
//
// Cardinality and cross-references survive the erasure: two rules under one
// policy still share a token and a route firewall rule still points at its
// route, so a path that drops a policy, merges two policies, or misattributes a
// rule to the wrong route still fails.
func normalizeIDSpace(nm *proto.NetworkMap) {
if nm == nil {
return
}
policies := newTokenizer("policy")
routes := newTokenizer("route")
for _, i := range orderBy(nm.Routes, routeKeyWithoutID) {
nm.Routes[i].ID = routes.get(nm.Routes[i].ID)
}
for _, i := range orderBy(nm.FirewallRules, firewallKeyWithoutPolicy) {
r := nm.FirewallRules[i]
if len(r.PolicyID) > 0 {
r.PolicyID = []byte(policies.get(string(r.PolicyID)))
}
}
for _, i := range orderBy(nm.RoutesFirewallRules, routeFirewallKeyWithoutIDs) {
r := nm.RoutesFirewallRules[i]
if len(r.PolicyID) > 0 {
r.PolicyID = []byte(policies.get(string(r.PolicyID)))
}
r.RouteID = routes.get(r.RouteID)
}
}
// tokenizer maps identifiers to positional placeholders in order of first use.
type tokenizer struct {
prefix string
seen map[string]string
}
func newTokenizer(prefix string) *tokenizer {
return &tokenizer{prefix: prefix, seen: make(map[string]string)}
}
func (t *tokenizer) get(id string) string {
if id == "" {
return ""
}
if tok, ok := t.seen[id]; ok {
return tok
}
tok := fmt.Sprintf("%s#%d", t.prefix, len(t.seen))
t.seen[id] = tok
return tok
}
// orderBy returns indices sorted by key, so placeholder numbering does not
// depend on the identifiers being erased.
func orderBy[T any](items []T, key func(T) string) []int {
idx := make([]int, len(items))
for i := range idx {
idx[i] = i
}
sort.SliceStable(idx, func(a, b int) bool { return key(items[idx[a]]) < key(items[idx[b]]) })
return idx
}
func routeKeyWithoutID(r *proto.Route) string {
if r == nil {
return ""
}
return fmt.Sprintf("%s|%s|%s|%d|%d|%t|%t|%v",
r.Network, r.NetID, r.Peer, r.Metric, r.NetworkType, r.Masquerade, r.KeepRoute, r.Domains)
}
func firewallKeyWithoutPolicy(r *proto.FirewallRule) string {
if r == nil {
return ""
}
return fmt.Sprintf("%s|%d|%d|%d|%s|%s|%v",
r.PeerIP, r.Direction, r.Action, r.Protocol, r.Port, portInfoKey(r.PortInfo), r.SourcePrefixes) //nolint:staticcheck
}
func routeFirewallKeyWithoutIDs(r *proto.RouteFirewallRule) string {
if r == nil {
return ""
}
return fmt.Sprintf("%s|%d|%d|%s|%v|%v|%t|%d",
r.Destination, r.Protocol, r.Action, portInfoKey(r.PortInfo), r.Domains, r.SourceRanges, r.IsDynamic, r.CustomProtocol)
}
// canonicalize sorts every repeated field of the NetworkMap by a stable key.
// The producing paths iterate Go maps while building these slices, so order
// can differ between runs even when the content is identical; comparing
// without this reports noise.
func canonicalize(nm *proto.NetworkMap) {
if nm == nil {
return
}
slices.SortFunc(nm.RemotePeers, cmpRemotePeer)
slices.SortFunc(nm.OfflinePeers, cmpRemotePeer)
slices.SortFunc(nm.Routes, cmpRoute)
slices.SortFunc(nm.FirewallRules, cmpFirewallRule)
slices.SortFunc(nm.RoutesFirewallRules, cmpRouteFirewallRule)
slices.SortFunc(nm.ForwardingRules, cmpForwardingRule)
for _, r := range nm.FirewallRules {
slices.SortFunc(r.SourcePrefixes, bytes.Compare)
}
for _, r := range nm.RoutesFirewallRules {
slices.Sort(r.SourceRanges)
}
canonicalizeDNSConfig(nm.DNSConfig)
canonicalizeSSHAuth(nm.SshAuth)
}
func canonicalizeDNSConfig(d *proto.DNSConfig) {
if d == nil {
return
}
for _, g := range d.NameServerGroups {
if g == nil {
continue
}
slices.Sort(g.Domains)
slices.SortFunc(g.NameServers, func(a, b *proto.NameServer) int {
if a == nil || b == nil {
return boolCmp(a == nil, b == nil)
}
if c := cmp.Compare(a.IP, b.IP); c != 0 {
return c
}
if c := cmp.Compare(a.Port, b.Port); c != 0 {
return c
}
return cmp.Compare(a.NSType, b.NSType)
})
}
slices.SortFunc(d.NameServerGroups, func(a, b *proto.NameServerGroup) int {
return cmp.Compare(nsgKey(a), nsgKey(b))
})
for _, z := range d.CustomZones {
if z == nil {
continue
}
slices.SortFunc(z.Records, cmpSimpleRecord)
}
slices.SortFunc(d.CustomZones, func(a, b *proto.CustomZone) int {
if a == nil || b == nil {
return boolCmp(a == nil, b == nil)
}
return cmp.Compare(a.Domain, b.Domain)
})
}
// canonicalizeSSHAuth sorts AuthorizedUsers and re-keys MachineUsers.Indexes
// against the new ordering, preserving which machine user maps to which hashes.
func canonicalizeSSHAuth(s *proto.SSHAuth) {
if s == nil || len(s.AuthorizedUsers) == 0 {
return
}
type hashed struct {
bytes []byte
old uint32
}
entries := make([]hashed, len(s.AuthorizedUsers))
for i, b := range s.AuthorizedUsers {
entries[i] = hashed{bytes: b, old: uint32(i)}
}
slices.SortFunc(entries, func(a, b hashed) int { return bytes.Compare(a.bytes, b.bytes) })
remap := make(map[uint32]uint32, len(entries))
sorted := make([][]byte, len(entries))
for newIdx, e := range entries {
remap[e.old] = uint32(newIdx)
sorted[newIdx] = e.bytes
}
s.AuthorizedUsers = sorted
for _, mu := range s.MachineUsers {
if mu == nil {
continue
}
for i, oldIdx := range mu.Indexes {
if newIdx, ok := remap[oldIdx]; ok {
mu.Indexes[i] = newIdx
}
}
slices.Sort(mu.Indexes)
}
}
func boolCmp(a, b bool) int {
if a == b {
return 0
}
if a {
return 1
}
return -1
}
func nsgKey(g *proto.NameServerGroup) string {
if g == nil {
return ""
}
var parts []string
for _, ns := range g.NameServers {
if ns == nil {
continue
}
parts = append(parts, ns.IP+":"+strconv.FormatInt(ns.Port, 10)+":"+strconv.FormatInt(ns.NSType, 10))
}
slices.Sort(parts)
key := strings.Join(parts, ",")
domains := append([]string(nil), g.Domains...)
slices.Sort(domains)
key += "|" + strings.Join(domains, "|")
if g.Primary {
key += "|P"
}
if g.SearchDomainsEnabled {
key += "|S"
}
return key
}
func cmpSimpleRecord(a, b *proto.SimpleRecord) int {
if a == nil || b == nil {
return boolCmp(a == nil, b == nil)
}
if c := cmp.Compare(a.Name, b.Name); c != 0 {
return c
}
if c := cmp.Compare(a.Type, b.Type); c != 0 {
return c
}
if c := cmp.Compare(a.Class, b.Class); c != 0 {
return c
}
if c := cmp.Compare(a.RData, b.RData); c != 0 {
return c
}
return cmp.Compare(a.TTL, b.TTL)
}
func cmpRemotePeer(a, b *proto.RemotePeerConfig) int {
if a == nil || b == nil {
return boolCmp(a == nil, b == nil)
}
return cmp.Compare(a.WgPubKey, b.WgPubKey)
}
func cmpRoute(a, b *proto.Route) int {
if a == nil || b == nil {
return boolCmp(a == nil, b == nil)
}
if c := cmp.Compare(a.ID, b.ID); c != 0 {
return c
}
if c := cmp.Compare(a.NetID, b.NetID); c != 0 {
return c
}
if c := cmp.Compare(a.Network, b.Network); c != 0 {
return c
}
if c := cmp.Compare(a.Peer, b.Peer); c != 0 {
return c
}
if c := cmp.Compare(a.Metric, b.Metric); c != 0 {
return c
}
return slices.Compare(a.Domains, b.Domains)
}
func cmpFirewallRule(a, b *proto.FirewallRule) int {
if a == nil || b == nil {
return boolCmp(a == nil, b == nil)
}
if c := bytes.Compare(a.PolicyID, b.PolicyID); c != 0 {
return c
}
if c := cmp.Compare(a.PeerIP, b.PeerIP); c != 0 { //nolint:staticcheck
return c
}
if c := cmp.Compare(int32(a.Direction), int32(b.Direction)); c != 0 {
return c
}
if c := cmp.Compare(int32(a.Action), int32(b.Action)); c != 0 {
return c
}
if c := cmp.Compare(int32(a.Protocol), int32(b.Protocol)); c != 0 {
return c
}
if c := cmp.Compare(a.Port, b.Port); c != 0 {
return c
}
return cmp.Compare(portInfoKey(a.PortInfo), portInfoKey(b.PortInfo))
}
func cmpRouteFirewallRule(a, b *proto.RouteFirewallRule) int {
if a == nil || b == nil {
return boolCmp(a == nil, b == nil)
}
if c := bytes.Compare(a.PolicyID, b.PolicyID); c != 0 {
return c
}
if c := cmp.Compare(a.RouteID, b.RouteID); c != 0 {
return c
}
if c := cmp.Compare(a.Destination, b.Destination); c != 0 {
return c
}
if c := cmp.Compare(int32(a.Protocol), int32(b.Protocol)); c != 0 {
return c
}
if c := cmp.Compare(portInfoKey(a.PortInfo), portInfoKey(b.PortInfo)); c != 0 {
return c
}
if c := cmp.Compare(int32(a.Action), int32(b.Action)); c != 0 {
return c
}
if c := slices.Compare(a.Domains, b.Domains); c != 0 {
return c
}
if c := slices.Compare(a.SourceRanges, b.SourceRanges); c != 0 {
return c
}
if c := cmp.Compare(a.CustomProtocol, b.CustomProtocol); c != 0 {
return c
}
return boolCmp(a.IsDynamic, b.IsDynamic)
}
func cmpForwardingRule(a, b *proto.ForwardingRule) int {
if a == nil || b == nil {
return boolCmp(a == nil, b == nil)
}
if c := cmp.Compare(int32(a.Protocol), int32(b.Protocol)); c != 0 {
return c
}
return bytes.Compare(a.TranslatedAddress, b.TranslatedAddress)
}
func portInfoKey(pi *proto.PortInfo) string {
if pi == nil {
return ""
}
switch sel := pi.PortSelection.(type) {
case *proto.PortInfo_Port:
return "P" + strconv.FormatUint(uint64(sel.Port), 10)
case *proto.PortInfo_Range_:
if sel.Range == nil {
return "R"
}
return "R" + strconv.FormatUint(uint64(sel.Range.Start), 10) + "-" + strconv.FormatUint(uint64(sel.Range.End), 10)
}
return ""
}

View File

@@ -1,218 +0,0 @@
package nmaptest
import (
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"net"
"os"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
// LoadNetworkMapData reads a fixture holding the NetworkMapData the store
// would return for one account. Unknown fields are rejected so fixture typos
// fail loudly instead of silently testing a default.
func LoadNetworkMapData(path string) (*networkmap.NetworkMapData, error) {
f, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("open fixture: %w", err)
}
defer f.Close()
dec := json.NewDecoder(f)
dec.DisallowUnknownFields()
var nmData networkmap.NetworkMapData
if err := dec.Decode(&nmData); err != nil {
return nil, fmt.Errorf("decode fixture %s: %w", path, err)
}
return &nmData, nil
}
var defaultNetworkNet = func() net.IPNet {
_, ipnet, err := net.ParseCIDR("100.64.0.0/10")
if err != nil {
panic(err)
}
return *ipnet
}()
// applyFixtureDefaults fills the boilerplate a fixture may omit. Map-keyed
// objects inherit their key as ID, peers get a deterministic WG-shaped key
// and their ID as DNS label, PublicIDs default to the internal ID (the
// envelope encoder puts public IDs on the wire and silently degrades on
// empty ones), and a nil ValidatedPeers validates every peer — production
// fills it through the integrated validator, not the store.
func applyFixtureDefaults(nmData *networkmap.NetworkMapData) {
if nmData.Network == nil {
nmData.Network = &nmdata.Network{}
}
if nmData.Network.Identifier == "" {
nmData.Network.Identifier = "network"
}
if nmData.Network.Net.IP == nil {
nmData.Network.Net = defaultNetworkNet
}
if nmData.AccountSettings == nil {
nmData.AccountSettings = &nmdata.AccountSettingsInfo{}
}
if nmData.DNSSettings == nil {
nmData.DNSSettings = &nmdata.DNSSettings{}
}
for id, p := range nmData.Peers {
if p == nil {
continue
}
if p.ID == "" {
p.ID = id
}
if p.Key == "" {
p.Key = derivedWgKey(p.ID)
}
if p.DNSLabel == "" {
p.DNSLabel = p.ID
}
}
for id, g := range nmData.Groups {
if g == nil {
continue
}
if g.ID == "" {
g.ID = id
}
if g.Name == "" {
g.Name = g.ID
}
if g.PublicID == "" {
g.PublicID = g.ID
}
}
for _, policy := range nmData.Policies {
defaultPolicyIDs(policy)
}
resolveResourcePolicyRefs(nmData)
for _, r := range nmData.Routes {
if r != nil && r.PublicID == "" {
r.PublicID = r.ID
}
}
for _, nsg := range nmData.NameServerGroups {
if nsg != nil && nsg.PublicID == "" {
nsg.PublicID = nsg.ID
}
}
for _, res := range nmData.NetworkResources {
if res == nil {
continue
}
if res.PublicID == "" {
res.PublicID = res.ID
}
defaultXIDMapping(&nmData.NetworkXIDToPublicID, res.NetworkID)
}
for networkID, routers := range nmData.Routers {
defaultXIDMapping(&nmData.NetworkXIDToPublicID, networkID)
for _, router := range routers {
if router != nil && router.PublicID == "" {
router.PublicID = networkID
}
}
}
for id, pc := range nmData.PostureChecks {
if pc == nil {
continue
}
if pc.ID == "" {
pc.ID = id
}
defaultXIDMapping(&nmData.PostureCheckXIDToPublicID, pc.ID)
}
if nmData.ValidatedPeers == nil {
nmData.ValidatedPeers = make(map[string]struct{}, len(nmData.Peers))
for id := range nmData.Peers {
nmData.ValidatedPeers[id] = struct{}{}
}
}
}
// resolveResourcePolicyRefs lets a fixture name an account policy by ID in
// ResourcePolicies — {"ID": "pol-x"} with no rules — instead of repeating it.
// The real store puts the same policy pointer in both places, which is what
// resolving the reference reproduces.
func resolveResourcePolicyRefs(nmData *networkmap.NetworkMapData) {
byID := make(map[string]*nmdata.Policy, len(nmData.Policies))
for _, policy := range nmData.Policies {
if policy != nil && policy.ID != "" {
byID[policy.ID] = policy
}
}
for _, policies := range nmData.ResourcePolicies {
for i, policy := range policies {
if policy == nil {
continue
}
if len(policy.Rules) == 0 {
if full, ok := byID[policy.ID]; ok {
policies[i] = full
continue
}
}
defaultPolicyIDs(policy)
}
}
}
func defaultPolicyIDs(policy *nmdata.Policy) {
if policy == nil {
return
}
if policy.PublicID == "" {
policy.PublicID = policy.ID
}
for i, rule := range policy.Rules {
if rule == nil {
continue
}
if rule.PolicyID == "" {
rule.PolicyID = policy.ID
}
if rule.ID == "" {
// Production gives a rule its policy's id (management/server/policy.go:205,
// "when policy can contain multiple rules, need refactor"), so a
// single-rule policy — the only shape the product can create today —
// must be modelled that way or the wire ids come out unrealistic.
rule.ID = policy.ID
if len(policy.Rules) > 1 {
rule.ID = fmt.Sprintf("%s-rule-%d", policy.ID, i)
}
}
}
}
func defaultXIDMapping(m *map[string]string, id string) {
if id == "" {
return
}
if *m == nil {
*m = make(map[string]string)
}
if _, ok := (*m)[id]; !ok {
(*m)[id] = id
}
}
// derivedWgKey returns a deterministic base64 key of 32 bytes, valid for the
// envelope decoder's WG-key identity.
func derivedWgKey(peerID string) string {
sum := sha256.Sum256([]byte(peerID))
return base64.StdEncoding.EncodeToString(sum[:])
}

View File

@@ -1,12 +0,0 @@
package nmaptest_test
import (
"path/filepath"
"testing"
"github.com/netbirdio/netbird/management/internals/controllers/network_map/nmaptest"
)
func TestNetworkMapGolden(t *testing.T) {
nmaptest.RunGoldenDir(t, filepath.Join("testdata", "cases"))
}

View File

@@ -1,311 +0,0 @@
// Package nmaptest measures network map generation on the dedicated store
// path against committed expectations. A case stands in for the store load
// with a NetworkMapData fixture — the value NetworkMapDBStoreImpl returns for
// one account — then runs the production per-peer pipeline the controller
// uses, PeersCustomZone → GetPeerNetworkMapComponents → proto conversion, in
// both wire shapes: the legacy full map (grpc.ToSyncResponse) and the
// component envelope expanded client-side (grpc.ToComponentSyncResponse →
// networkmap.EnvelopeToNetworkMap).
//
// The expectation files are the point of the framework. They state what the
// output should be, so a failing case means the code disagrees with the
// expectation and the answer is normally to fix the code; an expectation
// changes only through a deliberate reviewed edit. Nothing in this package
// writes to testdata — there is no flag that records current behaviour into an
// expectation, because that is how a defect becomes the baseline. Cases whose
// expectation encodes correct behaviour the code does not yet deliver stay red
// on purpose.
//
// A case lives in testdata/cases/<name>/ as case.json (manifest: description,
// peers, optional accountID, dnsDomain, modes), nmdata.json (the fixture the
// mocked store returns, using Go field names; zero values may be omitted and
// applyFixtureDefaults fills the boilerplate) and golden/<peerID>.json.
//
// There is ONE expectation per peer, shared by every mode. The modes are not
// different computations: CalculateNetworkMapFromComponents is
// components.Calculate, and both sides assemble the proto with the same
// encode helpers, so the only variable is what the envelope round-trip did to
// the components in transit. Any difference between modes is therefore a
// round-trip fidelity defect, and a shared expectation is what exposes it.
// Results are canonicalized before comparison, since repeated proto fields
// come from map iteration.
package nmaptest
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"github.com/google/go-cmp/cmp"
"github.com/stretchr/testify/require"
"golang.org/x/exp/maps"
"google.golang.org/protobuf/encoding/protojson"
"google.golang.org/protobuf/testing/protocmp"
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller"
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache"
mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
)
// Mode selects the wire shape a case is verified through. Both end in a
// *proto.NetworkMap, the one comparison surface shared by every path.
type Mode string
const (
// ModeFull is the legacy wire shape: the server runs Calculate and sends
// the expanded map (grpc.ToSyncResponse).
ModeFull Mode = "full"
// ModeEnvelope is the component wire shape: the server encodes components
// into a NetworkMapEnvelope (grpc.ToComponentSyncResponse) and the map is
// expanded the way the client engine does (networkmap.EnvelopeToNetworkMap).
ModeEnvelope Mode = "envelope"
defaultAccountID = "account"
defaultDNSDomain = "netbird.test"
)
var defaultModes = []Mode{ModeFull, ModeEnvelope}
// Case is one nmap-generation scenario: store data for a single account, the
// peers whose network maps are computed, and the directory holding one expected
// *proto.NetworkMap per peer — shared by every mode.
type Case struct {
Name string
AccountID string
DNSDomain string
Peers []string
Modes []Mode
Data *networkmap.NetworkMapData
GoldenDir string
}
type manifest struct {
Description string
AccountID string
DNSDomain string
Peers []string
Modes []Mode
}
// RunGoldenDir discovers and runs every fixture case under dir. A case is a
// directory containing case.json (manifest), nmdata.json (store fixture) and
// golden/<peerID>.json (expected proto.NetworkMap, protojson).
func RunGoldenDir(t *testing.T, dir string) {
t.Helper()
entries, err := os.ReadDir(dir)
require.NoError(t, err, "read cases dir")
ran := 0
for _, entry := range entries {
if !entry.IsDir() || strings.HasPrefix(entry.Name(), ".") {
continue
}
caseDir := filepath.Join(dir, entry.Name())
c, err := loadCase(caseDir)
require.NoError(t, err, "load case %s", entry.Name())
ran++
t.Run(entry.Name(), func(t *testing.T) {
RunCase(t, c)
})
}
require.NotZero(t, ran, "no cases found under %s", dir)
}
func loadCase(caseDir string) (Case, error) {
raw, err := os.ReadFile(filepath.Join(caseDir, "case.json"))
if err != nil {
return Case{}, fmt.Errorf("read manifest: %w", err)
}
dec := json.NewDecoder(bytes.NewReader(raw))
dec.DisallowUnknownFields()
var m manifest
if err := dec.Decode(&m); err != nil {
return Case{}, fmt.Errorf("decode manifest: %w", err)
}
data, err := LoadNetworkMapData(filepath.Join(caseDir, "nmdata.json"))
if err != nil {
return Case{}, err
}
return Case{
Name: filepath.Base(caseDir),
AccountID: m.AccountID,
DNSDomain: m.DNSDomain,
Peers: m.Peers,
Modes: m.Modes,
Data: data,
GoldenDir: filepath.Join(caseDir, "golden"),
}, nil
}
// RunCase computes each target peer's network map through every enabled mode
// and compares the canonicalized result against the peer's expectation file.
// It mirrors the controller's store path: fill fixture defaults, precompute
// posture validation once, then run the per-peer pipeline.
func RunCase(t *testing.T, c Case) {
t.Helper()
require.NotNil(t, c.Data, "case %s: Data is required", c.Name)
require.NotEmpty(t, c.Peers, "case %s: Peers is required", c.Name)
require.NotEmpty(t, c.GoldenDir, "case %s: GoldenDir is required", c.Name)
if c.AccountID == "" {
c.AccountID = defaultAccountID
}
if c.DNSDomain == "" {
c.DNSDomain = defaultDNSDomain
}
if len(c.Modes) == 0 {
c.Modes = defaultModes
}
ctx := context.Background()
nmData := c.Data
applyFixtureDefaults(nmData)
nmData.PrecomputePostureValidation()
dnsDomain := c.DNSDomain
if nmData.AccountSettings.DNSDomain != "" {
dnsDomain = nmData.AccountSettings.DNSDomain
}
zone := networkmap.PeersCustomZone(ctx, c.AccountID, dnsDomain, nmData.Peers, controller.IPv6AllowedPeersFromData(nmData))
dnsFwdPort := controller.ComputeForwarderPortFromData(nmData.Peers, network_map.DnsForwarderPortMinVersion)
for _, mode := range c.Modes {
if mode == ModeEnvelope {
requireEnvelopeSafeKeys(t, nmData, c.Name)
break
}
}
for _, peerID := range c.Peers {
peer := nmData.Peers[peerID]
require.NotNil(t, peer, "case %s: target peer %q not in fixture", c.Name, peerID)
for _, mode := range c.Modes {
t.Run(peerID+"/"+string(mode), func(t *testing.T) {
got := computeMode(t, ctx, mode, nmData, peerID, zone, dnsDomain, dnsFwdPort)
canonicalize(got)
compareGolden(t, filepath.Join(c.GoldenDir, peerID+".json"), got, mode)
})
}
}
}
// computeMode produces the peer's proto.NetworkMap the way the controller does
// for that wire shape.
func computeMode(t *testing.T, ctx context.Context, mode Mode, nmData *networkmap.NetworkMapData,
peerID string, zone nmdata.CustomZone, dnsDomain string, dnsFwdPort int64) *proto.NetworkMap {
t.Helper()
peer := nmData.Peers[peerID]
require.NotNil(t, peer, "target peer %q not in fixture", peerID)
switch mode {
case ModeFull:
nmap := controller.NetworkMapFromData(ctx, nmData, peerID, zone)
return mgmtgrpc.ToSyncResponse(ctx, nil, nil, nil, peer, nil, nil, nmap, dnsDomain, nil,
&cache.DNSConfigCache{}, nmData.AccountSettings, nil, nil, dnsFwdPort).NetworkMap
case ModeEnvelope:
components := nmData.GetPeerNetworkMapComponents(peerID, zone)
peerGroups := maps.Keys(nmData.GetPeerGroups(peerID))
resp := mgmtgrpc.ToComponentSyncResponse(ctx, nil, nil, nil, peer, nil, nil, components, nil,
dnsDomain, nil, nmData.AccountSettings, nil, peerGroups, dnsFwdPort)
res, err := networkmap.EnvelopeToNetworkMap(ctx, resp.NetworkMapEnvelope, peer.Key, dnsDomain)
require.NoError(t, err, "expand envelope")
return res.NetworkMap
default:
t.Fatalf("unknown mode %q", mode)
return nil
}
}
// requireEnvelopeSafeKeys fails fast on peer keys the envelope decoder would
// silently drop: it re-keys peers by base64 of the raw 32-byte WG public key.
func requireEnvelopeSafeKeys(t *testing.T, nmData *networkmap.NetworkMapData, caseName string) {
t.Helper()
for id, p := range nmData.Peers {
if p == nil {
continue
}
raw, err := base64.StdEncoding.DecodeString(p.Key)
if err != nil || len(raw) != 32 {
t.Fatalf("case %s: peer %q Key must be base64 of 32 bytes for mode %q (the envelope decoder drops it otherwise); use a real WireGuard public key or restrict the case to mode %q",
caseName, id, ModeEnvelope, ModeFull)
}
}
}
// compareGolden measures got against the committed expectation file. One
// expectation serves every mode, because the modes run the same computation and
// must therefore agree. The expectation is the authority: a mismatch means the
// code does not produce what this case says it should, so it is reported as a
// failure and not quietly absorbed.
//
// The full mode is compared verbatim, identifiers included, so the expectation
// pins real ids and stays readable. Other modes have identifiers erased on both
// sides first, because the envelope currently rewrites them — a tracked defect
// that TestIDSpaceMatches asserts against on its own, so it does not have to
// drown out every other case here.
// Nothing here writes to testdata. Expectation files are authored by hand and
// only ever change through a reviewed edit, so there is no mode in which a run
// can create or replace one. When a file is missing the computed map is printed
// for the author to read and, if it is genuinely correct, save deliberately.
func compareGolden(t *testing.T, path string, got *proto.NetworkMap, mode Mode) {
t.Helper()
if mode != ModeFull {
normalizeIDSpace(got)
canonicalize(got)
}
raw, err := os.ReadFile(path)
if err != nil {
rendered, mErr := renderNetworkMap(got)
require.NoError(t, mErr)
t.Fatalf("no expectation file %s: %v\nThis case has nothing to measure against — write the "+
"proto.NetworkMap this peer should receive. Mode %s currently produces:\n%s\nRead it before "+
"saving any of it: if the code is wrong, so is this.", path, err, mode, rendered)
}
want := &proto.NetworkMap{}
require.NoError(t, protojson.Unmarshal(raw, want), "parse expectation %s", path)
canonicalize(want)
if mode != ModeFull {
normalizeIDSpace(want)
canonicalize(want)
}
if diff := cmp.Diff(want, got, protocmp.Transform()); diff != "" {
t.Errorf("mode %s does not produce what %s expects (-want +got):\n%s\n"+
"Both modes run the same computation on the same components, so they must produce the same map. "+
"The expectation file is the committed statement of correct output — fix the code, or change the "+
"expectation deliberately if the intended behaviour really moved.", mode, path, diff)
}
}
// renderNetworkMap renders stable protojson: protojson output whitespace is
// deliberately unstable, so it is reformatted through json.Indent.
func renderNetworkMap(nm *proto.NetworkMap) ([]byte, error) {
raw, err := protojson.Marshal(nm)
if err != nil {
return nil, err
}
var buf bytes.Buffer
if err := json.Indent(&buf, raw, "", " "); err != nil {
return nil, err
}
buf.WriteByte('\n')
return buf.Bytes(), nil
}

View File

@@ -1,7 +0,0 @@
{
"description": "Two groups joined by one allow-all policy; peer-c has SSH enabled so the legacy-SSH path fills SshAuth from AllowedUserIDs.",
"peers": [
"peer-a",
"peer-c"
]
}

View File

@@ -1,65 +0,0 @@
{
"Serial": "5",
"peerConfig": {
"address": "100.64.0.1/10",
"sshConfig": {},
"fqdn": "peer-a.netbird.test",
"RoutingPeerDnsResolutionEnabled": true,
"autoUpdate": {}
},
"remotePeers": [
{
"wgPubKey": "4deEImv8zGvsyBmmfC2G0eQkbyMzyGuz/YK7pcYETwM=",
"allowedIps": [
"100.64.0.3/32"
],
"sshConfig": {
"sshPubKey": "c3NoLXBlZXItYw=="
},
"fqdn": "peer-c.netbird.test",
"agentVersion": "0.60.0"
}
],
"DNSConfig": {
"ServiceEnable": true,
"CustomZones": [
{
"Domain": "netbird.test.",
"Records": [
{
"Name": "peer-a.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.1"
},
{
"Name": "peer-c.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.3"
}
]
}
],
"ForwarderPort": "22054"
},
"FirewallRules": [
{
"PeerIP": "100.64.0.3",
"Protocol": "ALL",
"PolicyID": "cG9sLWFsbA=="
},
{
"PeerIP": "100.64.0.3",
"Direction": "OUT",
"Protocol": "ALL",
"PolicyID": "cG9sLWFsbA=="
}
],
"routesFirewallRulesIsEmpty": true,
"sshAuth": {
"UserIDClaim": "sub"
}
}

View File

@@ -1,102 +0,0 @@
{
"Serial": "5",
"peerConfig": {
"address": "100.64.0.3/10",
"sshConfig": {
"sshEnabled": true
},
"fqdn": "peer-c.netbird.test",
"RoutingPeerDnsResolutionEnabled": true,
"autoUpdate": {}
},
"remotePeers": [
{
"wgPubKey": "AvldyrZ12Pf90jzf3AXmhPwg3UcI+jtJHfbpBlupvko=",
"allowedIps": [
"100.64.0.2/32"
],
"sshConfig": {},
"fqdn": "peer-b.netbird.test",
"agentVersion": "0.60.0"
},
{
"wgPubKey": "vblMc9U8RAI6cVopcKEMTVT6lVC3D9nTTMSwot5d3L4=",
"allowedIps": [
"100.64.0.1/32"
],
"sshConfig": {},
"fqdn": "peer-a.netbird.test",
"agentVersion": "0.60.0"
}
],
"DNSConfig": {
"ServiceEnable": true,
"CustomZones": [
{
"Domain": "netbird.test.",
"Records": [
{
"Name": "peer-a.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.1"
},
{
"Name": "peer-b.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.2"
},
{
"Name": "peer-c.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.3"
}
]
}
],
"ForwarderPort": "22054"
},
"FirewallRules": [
{
"PeerIP": "100.64.0.1",
"Protocol": "ALL",
"PolicyID": "cG9sLWFsbA=="
},
{
"PeerIP": "100.64.0.1",
"Direction": "OUT",
"Protocol": "ALL",
"PolicyID": "cG9sLWFsbA=="
},
{
"PeerIP": "100.64.0.2",
"Protocol": "ALL",
"PolicyID": "cG9sLWFsbA=="
},
{
"PeerIP": "100.64.0.2",
"Direction": "OUT",
"Protocol": "ALL",
"PolicyID": "cG9sLWFsbA=="
}
],
"routesFirewallRulesIsEmpty": true,
"sshAuth": {
"UserIDClaim": "sub",
"AuthorizedUsers": [
"u9dHvAXZJKiXITuwP9jD/A=="
],
"machineUsers": {
"*": {
"indexes": [
0
]
}
}
}
}

View File

@@ -1,31 +0,0 @@
{
"Network": {"Serial": 5},
"AccountSettings": {"RoutingPeerDNSResolutionEnabled": true},
"Peers": {
"peer-a": {"IP": "100.64.0.1", "Meta": {"WtVersion": "0.60.0"}},
"peer-b": {"IP": "100.64.0.2", "Meta": {"WtVersion": "0.60.0"}},
"peer-c": {"IP": "100.64.0.3", "SSHEnabled": true, "SSHKey": "ssh-peer-c", "Meta": {"WtVersion": "0.60.0"}}
},
"Groups": {
"grp-dev": {"Peers": ["peer-a", "peer-b"]},
"grp-ops": {"Peers": ["peer-c"]}
},
"Policies": [
{
"ID": "pol-all",
"PublicID": "pol-all-pub",
"Enabled": true,
"Rules": [
{
"Enabled": true,
"Action": "accept",
"Protocol": "all",
"Bidirectional": true,
"Sources": ["grp-dev"],
"Destinations": ["grp-ops"]
}
]
}
],
"AllowedUserIDs": {"user-ops": {}}
}

View File

@@ -1,7 +0,0 @@
{
"description": "Nameserver group and applied custom zone distributed to grp-dev; peer-a (with an extra DNS label) receives them, peer-c outside the group receives neither.",
"peers": [
"peer-a",
"peer-c"
]
}

View File

@@ -1,93 +0,0 @@
{
"Serial": "8",
"peerConfig": {
"address": "100.64.0.1/10",
"sshConfig": {},
"fqdn": "peer-a.netbird.test",
"RoutingPeerDnsResolutionEnabled": true,
"autoUpdate": {}
},
"remotePeers": [
{
"wgPubKey": "AvldyrZ12Pf90jzf3AXmhPwg3UcI+jtJHfbpBlupvko=",
"allowedIps": [
"100.64.0.2/32"
],
"sshConfig": {},
"fqdn": "peer-b.netbird.test",
"agentVersion": "0.60.0"
}
],
"DNSConfig": {
"ServiceEnable": true,
"NameServerGroups": [
{
"NameServers": [
{
"IP": "8.8.8.8",
"Port": "53"
}
],
"Primary": true
}
],
"CustomZones": [
{
"Domain": "corp.internal.",
"Records": [
{
"Name": "db.corp.internal.",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "10.10.0.5"
}
]
},
{
"Domain": "netbird.test.",
"Records": [
{
"Name": "peer-a.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.1"
},
{
"Name": "peer-b.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.2"
},
{
"Name": "www.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.1"
}
]
}
],
"ForwarderPort": "22054"
},
"FirewallRules": [
{
"PeerIP": "100.64.0.2",
"Protocol": "ALL",
"PolicyID": "cG9sLW1lc2g="
},
{
"PeerIP": "100.64.0.2",
"Direction": "OUT",
"Protocol": "ALL",
"PolicyID": "cG9sLW1lc2g="
}
],
"routesFirewallRulesIsEmpty": true,
"sshAuth": {
"UserIDClaim": "sub"
}
}

View File

@@ -1,34 +0,0 @@
{
"Serial": "8",
"peerConfig": {
"address": "100.64.0.3/10",
"sshConfig": {},
"fqdn": "peer-c.netbird.test",
"RoutingPeerDnsResolutionEnabled": true,
"autoUpdate": {}
},
"remotePeersIsEmpty": true,
"DNSConfig": {
"ServiceEnable": true,
"CustomZones": [
{
"Domain": "netbird.test.",
"Records": [
{
"Name": "peer-c.netbird.test",
"Type": "1",
"Class": "IN",
"TTL": "300",
"RData": "100.64.0.3"
}
]
}
],
"ForwarderPort": "22054"
},
"firewallRulesIsEmpty": true,
"routesFirewallRulesIsEmpty": true,
"sshAuth": {
"UserIDClaim": "sub"
}
}

View File

@@ -1,51 +0,0 @@
{
"Network": {"Serial": 8},
"AccountSettings": {"RoutingPeerDNSResolutionEnabled": true},
"Peers": {
"peer-a": {"IP": "100.64.0.1", "ExtraDNSLabels": ["www"], "Meta": {"WtVersion": "0.60.0"}},
"peer-b": {"IP": "100.64.0.2", "Meta": {"WtVersion": "0.60.0"}},
"peer-c": {"IP": "100.64.0.3", "Meta": {"WtVersion": "0.60.0"}}
},
"Groups": {
"grp-dev": {"Peers": ["peer-a", "peer-b"]},
"grp-ops": {"Peers": ["peer-c"]}
},
"Policies": [
{
"ID": "pol-mesh",
"PublicID": "pol-mesh-pub",
"Enabled": true,
"Rules": [
{
"Enabled": true,
"Action": "accept",
"Protocol": "all",
"Bidirectional": true,
"Sources": ["grp-dev"],
"Destinations": ["grp-dev"]
}
]
}
],
"NameServerGroups": [
{
"ID": "nsg-1",
"Name": "dns-primary",
"NameServers": [{"IP": "8.8.8.8", "Port": 53}],
"Groups": ["grp-dev"],
"Primary": true,
"Enabled": true
}
],
"AppliedZoneCandidates": [
{
"DistributionGroups": ["grp-dev"],
"Zone": {
"Domain": "corp.internal.",
"Records": [
{"Name": "db.corp.internal.", "Type": 1, "Class": "IN", "TTL": 300, "RData": "10.10.0.5"}
]
}
}
]
}

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