Merge branch 'main' into fix/pkce-flow-session-extend

# Conflicts:
#	client/ios/NetBirdSDK/login.go
#	client/server/server.go
#	shared/management/proto/management.pb.go
This commit is contained in:
Zoltán Papp
2026-10-05 15:37:54 +02:00
847 changed files with 67577 additions and 25637 deletions
+59 -28
View File
@@ -3,7 +3,6 @@ package cmd
import (
"context"
"fmt"
"os/user"
"strings"
"time"
@@ -24,7 +23,10 @@ import (
"github.com/netbirdio/netbird/version"
)
const errCloseConnection = "Failed to close connection: %v"
const (
errCloseConnection = "Failed to close connection: %v"
noUpDownFlag = "no-updown"
)
var (
logFileCount uint32
@@ -114,7 +116,7 @@ func debugConfigDump(cmd *cobra.Command, _ []string) error {
if err != nil {
return fmt.Errorf("get active profile: %v", err)
}
currUser, err := user.Current()
currUser, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %v", err)
}
@@ -258,13 +260,14 @@ func runForDuration(cmd *cobra.Command, args []string) error {
}
stateWasDown := stat.Status != string(internal.StatusConnected) && stat.Status != string(internal.StatusConnecting)
noUpDown, _ := cmd.Flags().GetBool(noUpDownFlag)
initialLogLevel, err := client.GetLogLevel(cmd.Context(), &proto.GetLogLevelRequest{})
if err != nil {
return fmt.Errorf("failed to get log level: %v", status.Convert(err).Message())
}
if stateWasDown {
if stateWasDown && !noUpDown {
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
} else {
@@ -285,34 +288,20 @@ func runForDuration(cmd *cobra.Command, args []string) error {
}
needsRestoreUp := false
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message())
if noUpDown {
enableSyncResponsePersistence(cmd, client)
} else {
needsRestoreUp = !stateWasDown
cmd.Println("netbird down")
needsRestoreUp = restartDaemon(cmd, client, stateWasDown)
}
time.Sleep(1 * time.Second)
// Enable sync response persistence before bringing the service up
if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{
Enabled: true,
}); err != nil {
cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message())
}
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
} else {
needsRestoreUp = false
cmd.Println("netbird up")
}
time.Sleep(3 * time.Second)
cpuProfilingStarted := false
if _, err := client.StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil {
cmd.PrintErrf("Failed to start CPU profiling: %v\n", err)
if msg := status.Convert(err).Message(); strings.Contains(msg, "already in progress") {
cmd.PrintErrln("CPU profiling is already running (started with `netbird debug cpu start`). " +
"It is left running and is included in a bundle created after `netbird debug cpu stop`.")
} else {
cmd.PrintErrf("Failed to start CPU profiling: %v\n", msg)
}
} else {
cpuProfilingStarted = true
defer func() {
@@ -402,7 +391,7 @@ func runForDuration(cmd *cobra.Command, args []string) error {
}
}
if stateWasDown {
if stateWasDown && !noUpDown {
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
cmd.PrintErrf("Failed to restore service down state: %v\n", status.Convert(err).Message())
} else {
@@ -459,6 +448,47 @@ func setSyncResponsePersistence(cmd *cobra.Command, args []string) error {
return nil
}
// enableSyncResponsePersistence asks the daemon to keep the latest sync
// response so the bundle carries the network map. With a running daemon only
// syncs received after the call are kept.
func enableSyncResponsePersistence(cmd *cobra.Command, client proto.DaemonServiceClient) {
if _, err := client.SetSyncResponsePersistence(cmd.Context(), &proto.SetSyncResponsePersistenceRequest{
Enabled: true,
}); err != nil {
cmd.PrintErrf("Failed to enable sync response persistence: %v\n", status.Convert(err).Message())
}
}
// restartDaemon cycles the daemon down and up with sync response persistence
// enabled so the bundle carries the network map. It reports whether the
// daemon was left down although it was running before, so the caller can
// bring it back up.
func restartDaemon(cmd *cobra.Command, client proto.DaemonServiceClient, stateWasDown bool) bool {
needsRestoreUp := false
if _, err := client.Down(cmd.Context(), &proto.DownRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service down: %v\n", status.Convert(err).Message())
} else {
needsRestoreUp = !stateWasDown
cmd.Println("netbird down")
}
time.Sleep(1 * time.Second)
// Enable sync response persistence before bringing the service up
enableSyncResponsePersistence(cmd, client)
if _, err := client.Up(cmd.Context(), &proto.UpRequest{}); err != nil {
cmd.PrintErrf("Failed to bring service up: %v\n", status.Convert(err).Message())
} else {
needsRestoreUp = false
cmd.Println("netbird up")
}
time.Sleep(3 * time.Second)
return needsRestoreUp
}
func waitForDurationOrCancel(ctx context.Context, duration time.Duration, cmd *cobra.Command) error {
ticker := time.NewTicker(1 * time.Second)
defer ticker.Stop()
@@ -547,4 +577,5 @@ func init() {
forCmd.Flags().StringVar(&uploadBundleURLFlag, "upload-bundle-url", types.DefaultBundleURL, "Service URL to get an URL to upload the debug bundle")
forCmd.Flags().BoolVar(&uploadBundleInsecureFlag, "upload-bundle-insecure", false, "Allow uploading to an http or untrusted-TLS upload server (self-hosted); requires root")
forCmd.Flags().Bool("capture", false, "Capture packets during the debug duration and include in bundle")
forCmd.Flags().Bool(noUpDownFlag, false, "Keep the daemon running instead of bringing it down and up before collecting. The bundle only includes the network map if a sync arrives during the run")
}
+83
View File
@@ -0,0 +1,83 @@
package cmd
import (
"fmt"
log "github.com/sirupsen/logrus"
"github.com/spf13/cobra"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/proto"
)
var debugCPUCmd = &cobra.Command{
Use: "cpu",
Short: "Profile the daemon's CPU usage",
Long: `Starts and stops CPU profiling in the running daemon without restarting it.
The profile is included in the next debug bundle as cpu.prof.
Profiling is not time limited: it keeps running, and keeps costing CPU, until
"netbird debug cpu stop" is run.`,
}
var debugCPUStartCmd = &cobra.Command{
Use: "start",
Short: "Start CPU profiling in the daemon",
Example: " netbird debug cpu start",
Args: cobra.NoArgs,
RunE: debugCPUStart,
}
var debugCPUStopCmd = &cobra.Command{
Use: "stop",
Short: "Stop CPU profiling in the daemon",
Long: `Stops CPU profiling. The captured profile stays in the daemon until the next
debug bundle is created, which includes it as cpu.prof.`,
Example: " netbird debug cpu stop && netbird debug bundle",
Args: cobra.NoArgs,
RunE: debugCPUStop,
}
func debugCPUStart(cmd *cobra.Command, _ []string) error {
conn, err := getClient(cmd)
if err != nil {
return err
}
defer func() {
if err := conn.Close(); err != nil {
log.Errorf(errCloseConnection, err)
}
}()
if _, err := proto.NewDaemonServiceClient(conn).StartCPUProfile(cmd.Context(), &proto.StartCPUProfileRequest{}); err != nil {
return fmt.Errorf("start CPU profiling: %v", status.Convert(err).Message())
}
cmd.Println("CPU profiling started and runs until stopped. Run `netbird debug cpu stop` and then `netbird debug bundle` to collect it.")
return nil
}
func debugCPUStop(cmd *cobra.Command, _ []string) error {
conn, err := getClient(cmd)
if err != nil {
return err
}
defer func() {
if err := conn.Close(); err != nil {
log.Errorf(errCloseConnection, err)
}
}()
if _, err := proto.NewDaemonServiceClient(conn).StopCPUProfile(cmd.Context(), &proto.StopCPUProfileRequest{}); err != nil {
return fmt.Errorf("stop CPU profiling: %v", status.Convert(err).Message())
}
cmd.Println("CPU profiling stopped. Run `netbird debug bundle` to include cpu.prof.")
return nil
}
func init() {
debugCPUCmd.AddCommand(debugCPUStartCmd)
debugCPUCmd.AddCommand(debugCPUStopCmd)
debugCmd.AddCommand(debugCPUCmd)
}
+164
View File
@@ -0,0 +1,164 @@
package cmd
import (
"bytes"
"context"
"os/user"
"strings"
"testing"
"github.com/spf13/cobra"
"github.com/spf13/pflag"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/profilemanager"
)
// startDebugTestDaemon starts an in-process daemon with an isolated profile
// directory and returns the address the CLI should dial.
func startDebugTestDaemon(t *testing.T) string {
t.Helper()
tempDir := t.TempDir()
origDefaultProfileDir := profilemanager.DefaultConfigPathDir
origActiveProfileStatePath := profilemanager.ActiveProfileStatePath
origConfigDirOverride := profilemanager.ConfigDirOverride
origDaemonAddr := daemonAddr
t.Cleanup(func() {
profilemanager.DefaultConfigPathDir = origDefaultProfileDir
profilemanager.ActiveProfileStatePath = origActiveProfileStatePath
profilemanager.ConfigDirOverride = origConfigDirOverride
daemonAddr = origDaemonAddr
})
profilemanager.DefaultConfigPathDir = tempDir
profilemanager.ActiveProfileStatePath = tempDir + "/active_profile.json"
profilemanager.ConfigDirOverride = tempDir
currUser, err := user.Current()
require.NoError(t, err)
sm := profilemanager.ServiceManager{}
created, err := sm.AddProfile("test1", currUser.Username)
require.NoError(t, err)
require.NoError(t, sm.SetActiveProfileState(&profilemanager.ActiveProfileState{
ID: created.ID,
Username: currUser.Username,
}))
ctx, cancel := context.WithCancel(internal.CtxInitState(context.Background()))
srv, lis := startClientDaemon(t, ctx, "", tempDir+"/config.json")
t.Cleanup(func() {
cancel()
srv.Stop()
})
return "tcp://" + lis.Addr().String()
}
// runDebugCmd runs `netbird debug <args>` against the daemon at addr and
// returns everything the command printed.
func runDebugCmd(addr string, args ...string) (string, error) {
daemonAddr = addr
var out bytes.Buffer
rootCmd.SetOut(&out)
rootCmd.SetErr(&out)
rootCmd.SetArgs(append(append([]string{"debug"}, args...), "--daemon-addr", addr, "--log-file", ""))
err := rootCmd.Execute()
rootCmd.SetOut(nil)
rootCmd.SetErr(nil)
rootCmd.SetArgs(nil)
resetFlags(rootCmd)
return out.String(), err
}
// resetFlags puts every flag of the command and its subcommands back to its
// default so a value parsed in one run does not leak into the next in-process
// execution.
func resetFlags(cmd *cobra.Command) {
reset := func(f *pflag.Flag) {
// Set appends to a slice flag and would parse the "[a,b]" default
// text as elements, so slices are replaced instead.
if sv, ok := f.Value.(pflag.SliceValue); ok {
var def []string
if trimmed := strings.Trim(f.DefValue, "[]"); trimmed != "" {
def = strings.Split(trimmed, ",")
}
_ = sv.Replace(def)
} else {
_ = f.Value.Set(f.DefValue)
}
f.Changed = false
}
cmd.Flags().VisitAll(reset)
cmd.PersistentFlags().VisitAll(reset)
// Commands pin their writers to the buffer of the run that first used
// them, so a later run would print into the old buffer.
cmd.SetOut(nil)
cmd.SetErr(nil)
for _, sub := range cmd.Commands() {
resetFlags(sub)
}
}
// TestResetFlagsSliceDefault guards against Set("[]") on slice flags, which
// stores a literal "[]" element instead of the empty default.
func TestResetFlagsSliceDefault(t *testing.T) {
cmd := &cobra.Command{Use: "x"}
var env, withDefault []string
cmd.Flags().StringSliceVar(&env, "env", nil, "")
cmd.Flags().StringSliceVar(&withDefault, "names", []string{"a", "b"}, "")
require.NoError(t, cmd.Flags().Parse([]string{"--env", "K=V", "--names", "c"}))
resetFlags(cmd)
assert.Empty(t, env, "slice flag with no default must reset to empty")
assert.Equal(t, []string{"a", "b"}, withDefault, "slice flag must reset to its default")
}
func TestDebugCPUStartStop(t *testing.T) {
addr := startDebugTestDaemon(t)
run := func(args ...string) error {
_, err := runDebugCmd(addr, append([]string{"cpu"}, args...)...)
return err
}
require.Error(t, run("stop"), "stop without a running profile must fail")
require.NoError(t, run("start"))
assert.Error(t, run("start"), "second start must be rejected while profiling")
require.NoError(t, run("stop"))
assert.Error(t, run("stop"), "second stop must be rejected")
assert.NoError(t, run("start"), "profiling can be started again after a stop")
assert.NoError(t, run("stop"))
}
// TestDebugForKeepsRunningCPUProfile covers `debug for` started while a
// profile from `debug cpu start` is running: it must say so, leave the
// profile alone, and still create the bundle.
func TestDebugForKeepsRunningCPUProfile(t *testing.T) {
addr := startDebugTestDaemon(t)
_, err := runDebugCmd(addr, "cpu", "start")
require.NoError(t, err)
out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown")
require.NoError(t, err, "output: %s", out)
assert.Contains(t, out, "CPU profiling is already running", "the conflict must be explained")
assert.NotContains(t, out, "rpc error", "the raw RPC error must not reach the user")
assert.Contains(t, out, "Local file:", "the bundle must still be created")
_, err = runDebugCmd(addr, "cpu", "stop")
assert.NoError(t, err, "the profile started by the user must still be running")
}
func TestDebugForNoUpDown(t *testing.T) {
addr := startDebugTestDaemon(t)
out, err := runDebugCmd(addr, "for", "1s", "-S=false", "--no-updown")
require.NoError(t, err, "output: %s", out)
assert.NotContains(t, out, "netbird down", "--no-updown must not bring the daemon down")
assert.NotContains(t, out, "netbird up", "--no-updown must not bring the daemon up")
assert.Contains(t, out, "Local file:", "the bundle must still be created")
}
+13
View File
@@ -0,0 +1,13 @@
package cmd
// remoteJobsAllowedFlag opts this peer into running remote jobs (e.g. debug
// bundles) requested by the management server. It defaults to false: remote
// jobs are an explicit opt-in, and enabling it is a privileged change (see the
// daemon gate in client/server), mirroring the SSH server opt-in.
const remoteJobsAllowedFlag = "allow-remote-jobs"
var remoteJobsAllowed bool
func init() {
upCmd.PersistentFlags().BoolVar(&remoteJobsAllowed, remoteJobsAllowedFlag, false, "Allow the management server to run remote jobs (e.g. debug bundles) on this peer")
}
+24 -20
View File
@@ -4,7 +4,6 @@ import (
"context"
"fmt"
"os"
"os/user"
"strings"
log "github.com/sirupsen/logrus"
@@ -16,6 +15,7 @@ import (
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/auth"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/server"
@@ -53,7 +53,7 @@ var loginCmd = &cobra.Command{
// nolint
ctx = context.WithValue(ctx, system.DeviceNameCtxKey, hostName)
}
username, err := user.Current()
username, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %v", err)
}
@@ -74,7 +74,7 @@ var loginCmd = &cobra.Command{
if providedSetupKey != "" {
return fmt.Errorf("--extend cannot be combined with a setup key; setup keys can only enrol new peers")
}
if err := doExtendSession(ctx, cmd); err != nil {
if err := doExtendSession(ctx, cmd, activeProf); err != nil {
return fmt.Errorf("extend session failed: %v", err)
}
return nil
@@ -92,7 +92,7 @@ var loginCmd = &cobra.Command{
return fmt.Errorf("daemon login failed: %v", err)
}
cmd.Println("Logging successfully")
cmd.Println("Login successful")
return nil
},
@@ -176,7 +176,7 @@ func doDaemonLogin(ctx context.Context, cmd *cobra.Command, providedSetupKey str
// (browser + verification URL) and the resulting JWT is forwarded to the
// management server's ExtendAuthSession RPC. The tunnel stays up
// throughout — no Down/Up, no network-map resync.
func doExtendSession(ctx context.Context, cmd *cobra.Command) error {
func doExtendSession(ctx context.Context, cmd *cobra.Command, activeProf *profilemanager.Profile) error {
conn, err := DialClientGRPCServer(ctx, daemonAddr)
if err != nil {
//nolint
@@ -190,14 +190,12 @@ func doExtendSession(ctx context.Context, cmd *cobra.Command) error {
// the CLI runs in the user's session, the daemon does not: tell it what we can see
req := &proto.RequestExtendAuthSessionRequest{HasGraphicalSession: util.HasGraphicalSession()}
// Pre-fill the IdP login hint from the active profile so the user
// Pre-fill the IdP login hint from the resolved profile so the user
// doesn't have to retype their email. Best-effort: we still proceed
// without a hint if the lookup fails.
pm := profilemanager.NewProfileManager()
if active, perr := pm.GetActiveProfile(); perr == nil {
if profState, sperr := pm.GetProfileState(active.ID); sperr == nil && profState.Email != "" {
req.Hint = &profState.Email
}
if profState, perr := pm.GetProfileState(activeProf.ID); perr == nil && profState.Email != "" {
req.Hint = &profState.Email
}
startResp, err := client.RequestExtendAuthSession(ctx, req)
@@ -235,9 +233,11 @@ func getActiveProfile(ctx context.Context, pm *profilemanager.ProfileManager, pr
// switch profile if provided
if profileName != "" {
if err := switchProfileOnDaemon(ctx, pm, profileName, username); err != nil {
prof, err := switchProfileOnDaemon(ctx, pm, profileName, username)
if err != nil {
return nil, fmt.Errorf("switch profile: %v", err)
}
return prof, nil
}
activeProf, err := pm.GetActiveProfile()
@@ -251,20 +251,19 @@ func getActiveProfile(ctx context.Context, pm *profilemanager.ProfileManager, pr
return activeProf, nil
}
func switchProfileOnDaemon(ctx context.Context, pm *profilemanager.ProfileManager, handle string, username string) error {
func switchProfileOnDaemon(ctx context.Context, pm *profilemanager.ProfileManager, handle string, username string) (*profilemanager.Profile, error) {
resolvedID, err := switchProfile(ctx, handle, username)
if err != nil {
return fmt.Errorf("switch profile on daemon: %v", err)
return nil, fmt.Errorf("switch profile on daemon: %v", err)
}
if err := pm.SwitchProfile(resolvedID); err != nil {
return fmt.Errorf("switch profile: %v", err)
return nil, fmt.Errorf("switch profile: %v", err)
}
conn, err := DialClientGRPCServer(ctx, daemonAddr)
if err != nil {
log.Errorf("failed to connect to service CLI interface %v", err)
return err
return nil, fmt.Errorf("connect to service CLI interface: %w", err)
}
defer conn.Close()
@@ -272,17 +271,17 @@ func switchProfileOnDaemon(ctx context.Context, pm *profilemanager.ProfileManage
status, err := client.Status(ctx, &proto.StatusRequest{})
if err != nil {
return fmt.Errorf("unable to get daemon status: %v", err)
return nil, fmt.Errorf("unable to get daemon status: %v", err)
}
if status.Status == string(internal.StatusConnected) {
if _, err := client.Down(ctx, &proto.DownRequest{}); err != nil {
log.Errorf("call service down method: %v", err)
return err
return nil, err
}
}
return nil
return &profilemanager.Profile{ID: resolvedID}, nil
}
// switchProfile asks the daemon to switch to the profile identified by
@@ -332,6 +331,11 @@ func doForegroundLogin(ctx context.Context, cmd *cobra.Command, setupKey string,
if err != nil {
return fmt.Errorf("read config file %s: %v", configFilePath, err)
}
// CLI standalone login: profilemanager no longer auto-applies MDM,
// so layer in the OS-native policy here. Desktop builds construct
// a Loader with no fetcher — the build-tagged loadPlatform reads
// the registry/plist directly.
config.ApplyMDMPolicy(mdm.NewLoader(nil).Load())
// Mirror runInForegroundMode: recover residual state (DNS, firewall,
// ssh config, legacy routing) from a previous unclean shutdown and
@@ -345,7 +349,7 @@ func doForegroundLogin(ctx context.Context, cmd *cobra.Command, setupKey string,
if err != nil {
return fmt.Errorf("foreground login failed: %v", err)
}
cmd.Println("Logging successfully")
cmd.Println("Login successful")
return nil
}
+2 -2
View File
@@ -3,11 +3,11 @@ package cmd
import (
"context"
"fmt"
"os/user"
"time"
"github.com/spf13/cobra"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/proto"
)
@@ -37,7 +37,7 @@ var logoutCmd = &cobra.Command{
if profileName != "" {
req.ProfileName = &profileName
currUser, err := user.Current()
currUser, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %v", err)
}
+5 -6
View File
@@ -4,7 +4,6 @@ import (
"context"
"errors"
"fmt"
"os/user"
"strings"
"text/tabwriter"
"time"
@@ -97,7 +96,7 @@ func listProfilesFunc(cmd *cobra.Command, _ []string) error {
}
defer conn.Close()
currUser, err := user.Current()
currUser, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %w", err)
}
@@ -138,7 +137,7 @@ func addProfileFunc(cmd *cobra.Command, args []string) error {
return err
}
currUser, err := user.Current()
currUser, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %w", err)
}
@@ -179,7 +178,7 @@ func renameProfileFunc(cmd *cobra.Command, args []string) error {
}
defer conn.Close()
currUser, err := user.Current()
currUser, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %w", err)
}
@@ -233,7 +232,7 @@ func removeProfileFunc(cmd *cobra.Command, args []string) error {
}
defer conn.Close()
currUser, err := user.Current()
currUser, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %w", err)
}
@@ -261,7 +260,7 @@ func selectProfileFunc(cmd *cobra.Command, args []string) error {
profileManager := profilemanager.NewProfileManager()
handle := args[0]
currUser, err := user.Current()
currUser, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %w", err)
}
+7
View File
@@ -23,6 +23,7 @@ import (
"github.com/netbirdio/netbird/client/anonymize"
daddr "github.com/netbirdio/netbird/client/internal/daemonaddr"
"github.com/netbirdio/netbird/client/internal/localmetrics"
"github.com/netbirdio/netbird/client/internal/profilemanager"
)
@@ -31,6 +32,8 @@ const (
dnsResolverAddress = "dns-resolver-address"
enableRosenpassFlag = "enable-rosenpass"
rosenpassPermissiveFlag = "rosenpass-permissive"
enableLocalMetricsFlag = "enable-local-metrics"
localMetricsAddressFlag = "local-metrics-address"
preSharedKeyFlag = "preshared-key"
interfaceNameFlag = "interface-name"
wireguardPortFlag = "wireguard-port"
@@ -80,6 +83,8 @@ var (
updateSettingsDisabled bool
captureEnabled bool
networksDisabled bool
localMetricsEnabled bool
localMetricsAddr string
rootCmd = &cobra.Command{
Use: "netbird",
@@ -215,6 +220,8 @@ func init() {
upCmd.PersistentFlags().BoolVar(&rosenpassEnabled, enableRosenpassFlag, false, "[Experimental] Enable Rosenpass feature. If enabled, the connection will be post-quantum secured via Rosenpass.")
upCmd.PersistentFlags().BoolVar(&rosenpassPermissive, rosenpassPermissiveFlag, false, "[Experimental] Enable Rosenpass in permissive mode to allow this peer to accept WireGuard connections without requiring Rosenpass functionality from peers that do not have Rosenpass enabled.")
upCmd.PersistentFlags().BoolVar(&autoConnectDisabled, disableAutoConnectFlag, false, "Disables auto-connect feature. If enabled, then the client won't connect automatically when the service starts.")
upCmd.PersistentFlags().BoolVar(&localMetricsEnabled, enableLocalMetricsFlag, false, "Enables a local Prometheus /metrics endpoint exposing connection state (peers, latency, P2P vs relay).")
upCmd.PersistentFlags().StringVar(&localMetricsAddr, localMetricsAddressFlag, localmetrics.DefaultListenAddress, "Listen address of the local Prometheus /metrics endpoint.")
upCmd.PersistentFlags().BoolVar(&lazyConnEnabled, enableLazyConnectionFlag, false, "Deprecated: no longer used. Lazy connections are controlled by the server and the NB_LAZY_CONN environment variable.")
_ = upCmd.PersistentFlags().MarkDeprecated(enableLazyConnectionFlag, "no longer used; lazy connections are controlled by the server and the NB_LAZY_CONN environment variable")
+50
View File
@@ -7,6 +7,7 @@ import (
"fmt"
"net/http"
"runtime"
"slices"
"strings"
"sync"
@@ -25,6 +26,30 @@ var serviceCmd = &cobra.Command{
const defaultJSONSocket = "unix:///var/run/netbird-http.sock"
// forbiddenServiceEnvVars are the environment variables the service is never
// registered with, keyed in upper case since these are Windows names. Each one
// decides where the daemon resolves something it then uses with the privileges
// of the account it runs under — LocalSystem on Windows, root elsewhere: the
// executables it runs (PATH, PATHEXT, COMSPEC, SystemRoot, windir) or the
// directory it writes temporary files in (TEMP, TMP). The daemon needs none of
// them, and the utilities it shells out to are resolved by absolute path.
var forbiddenServiceEnvVars = map[string]struct{}{
"PATH": {},
"PATHEXT": {},
"SYSTEMROOT": {},
"WINDIR": {},
"COMSPEC": {},
"TEMP": {},
"TMP": {},
}
// forbiddenServiceEnvPrefixes are the dynamic-loader families, refused whole
// rather than by name: LD_PRELOAD, DYLD_INSERT_LIBRARIES and their siblings all
// reach the loader of the process, the set differs per platform and libc, and
// new members arrive with new OS releases. Listing them one by one is a list
// that is wrong the moment it is written.
var forbiddenServiceEnvPrefixes = []string{"LD_", "DYLD_"}
var (
serviceName string
serviceEnvVars []string
@@ -127,8 +152,33 @@ func parseServiceEnvVars(envVars []string) (map[string]string, error) {
return nil, fmt.Errorf("empty environment variable key in: %s", env)
}
if isForbiddenServiceEnvVar(key) {
return nil, fmt.Errorf("environment variable %s cannot be set on the service: it decides where the service resolves the executables, libraries or temporary files it uses", key)
}
envMap[key] = value
}
return envMap, nil
}
// isForbiddenServiceEnvVar reports whether name is one the service must not be
// registered with.
//
// The names are matched case-insensitively only on Windows, where they are the
// same variable however they are spelled. Elsewhere the environment is
// case-sensitive, so Path and PATH are two different variables and only the
// exact spelling is the one the loader reads.
func isForbiddenServiceEnvVar(name string) bool {
if runtime.GOOS == "windows" {
name = strings.ToUpper(name)
}
if _, forbidden := forbiddenServiceEnvVars[name]; forbidden {
return true
}
return slices.ContainsFunc(forbiddenServiceEnvPrefixes, func(prefix string) bool {
return strings.HasPrefix(name, prefix)
})
}
+4 -2
View File
@@ -41,13 +41,15 @@ func daemonServerOptions(network string) []grpc.ServerOption {
if network == "tcp" {
log.Warnf("daemon is listening on TCP (%s): callers carry no verifiable identity over TCP, "+
"so privileged operations (SSH root login, SSH auth, enabling the SSH server, management URL changes, "+
"deregistration) will be denied. Use a unix socket, or npipe:// on Windows", daemonAddr)
"deregistration) will be denied, and the SSH JWT cache is neither filled nor served. "+
"Use a unix socket, or npipe:// on Windows", daemonAddr)
return nil
}
creds := ipcauth.NewTransportCredentials() //nolint:staticcheck
if creds == nil { //nolint:staticcheck // nil only on platforms without a peer-identity primitive
log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied", runtime.GOOS)
log.Warnf("daemon IPC has no peer-identity primitive on %s: privileged operations will be denied "+
"and the SSH JWT cache is neither filled nor served", runtime.GOOS)
return nil
}
+50 -6
View File
@@ -14,6 +14,7 @@ import (
"github.com/netbirdio/netbird/client/configs"
"github.com/netbirdio/netbird/client/internal/daemonaddr"
"github.com/netbirdio/netbird/client/internal/elevate"
"github.com/netbirdio/netbird/util"
)
@@ -43,10 +44,33 @@ func serviceParamsPath() string {
// loadServiceParams reads saved service parameters from disk.
// Returns nil with no error if the file does not exist.
//
// The file is read by an elevated install and decides the arguments and the
// environment of the service it then registers, so it is used only when its
// ownership and permissions are the ones saveServiceParams leaves behind. That
// restricted ACL is applied when the file is written, which is not necessarily
// before it is first read, so this is checked rather than assumed. A file that
// fails the check is treated as absent, and the install proceeds with its
// defaults.
func loadServiceParams() (*serviceParams, error) {
path := serviceParamsPath()
data, err := os.ReadFile(path)
// Resolve links first so the checks apply to the file that is actually read.
// Since the check covers every directory above it as well, nobody who fails
// it can swap the file between here and the read below.
resolved, err := filepath.EvalSymlinks(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil //nolint:nilnil
}
return nil, fmt.Errorf("resolve service params %s: %w", path, err)
}
if err := elevate.CheckOnlyOwnerWritable(resolved); err != nil {
return nil, fmt.Errorf("refusing to read service params from %s: %w", resolved, err)
}
data, err := os.ReadFile(resolved)
if err != nil {
if os.IsNotExist(err) {
return nil, nil //nolint:nilnil
@@ -182,10 +206,16 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
// If --service-env was explicitly set to empty, all saved env vars are cleared.
// If --service-env was not set, saved env vars are used entirely.
func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
// A forbidden name explicitly passed on the command line is an error the
// operator is told about, but one restored from a file written by an older
// version is dropped: an install that refuses to run would leave the host
// without a daemon over a variable nobody is asking for any more.
saved := dropForbiddenServiceEnvVars(cmd, params.ServiceEnvVars)
if !cmd.Flags().Changed("service-env") {
if len(params.ServiceEnvVars) > 0 {
if len(saved) > 0 {
// No explicit env vars: rebuild serviceEnvVars from saved params.
serviceEnvVars = envMapToSlice(params.ServiceEnvVars)
serviceEnvVars = envMapToSlice(saved)
}
return
}
@@ -204,13 +234,13 @@ func applyServiceEnvParams(cmd *cobra.Command, params *serviceParams) {
return
}
if len(params.ServiceEnvVars) == 0 {
if len(saved) == 0 {
return
}
// Merge saved values underneath explicit ones.
merged := make(map[string]string, len(params.ServiceEnvVars)+len(explicit))
maps.Copy(merged, params.ServiceEnvVars)
merged := make(map[string]string, len(saved)+len(explicit))
maps.Copy(merged, saved)
maps.Copy(merged, explicit) // explicit wins on conflict
serviceEnvVars = envMapToSlice(merged)
}
@@ -233,6 +263,20 @@ var resetParamsCmd = &cobra.Command{
},
}
// dropForbiddenServiceEnvVars returns the saved entries that may still be
// registered on the service, reporting every one it leaves behind.
func dropForbiddenServiceEnvVars(cmd *cobra.Command, saved map[string]string) map[string]string {
kept := make(map[string]string, len(saved))
for key, value := range saved {
if isForbiddenServiceEnvVar(key) {
cmd.PrintErrf("Warning: ignoring saved service environment variable %s: it decides where the service resolves the executables, libraries or temporary files it uses\n", key)
continue
}
kept[key] = value
}
return kept
}
// envMapToSlice converts a map of env vars to a KEY=VALUE slice.
func envMapToSlice(m map[string]string) []string {
s := make([]string, 0, len(m))
+54
View File
@@ -9,6 +9,7 @@ import (
"go/token"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
@@ -353,6 +354,59 @@ func TestApplyServiceEnvParams_NotChanged(t *testing.T) {
assert.Equal(t, map[string]string{"FROM_SAVED": "val"}, result)
}
func TestParseServiceEnvVars_RejectsForbiddenNames(t *testing.T) {
for _, env := range []string{"PATH=C:\\somewhere", "LD_PRELOAD=/tmp/lib.so", "DYLD_FALLBACK_LIBRARY_PATH=/tmp"} {
_, err := parseServiceEnvVars([]string{"KEEP=me", env})
require.Errorf(t, err, "%s selects what the service resolves and must be refused", env)
}
}
func TestIsForbiddenServiceEnvVar(t *testing.T) {
// The loader families are matched by prefix, so a name nobody has heard of
// yet is refused too.
for _, name := range []string{
"PATH", "PATHEXT", "COMSPEC", "SYSTEMROOT", "WINDIR", "TEMP", "TMP",
"LD_PRELOAD", "LD_AUDIT", "DYLD_INSERT_LIBRARIES", "DYLD_FALLBACK_FRAMEWORK_PATH",
} {
assert.Truef(t, isForbiddenServiceEnvVar(name), "%s must be refused", name)
}
// The prefix must not swallow names that merely start with the same letters.
for _, name := range []string{"NB_LOG_LEVEL", "NB_WG_DEBUG", "HTTPS_PROXY", "LDAP_URL", "DYLDX"} {
assert.Falsef(t, isForbiddenServiceEnvVar(name), "%s has no reason to be refused", name)
}
// On Windows a variable is the same one however it is spelled; elsewhere
// Path and PATH are two variables and only the exact one is read.
if runtime.GOOS == "windows" {
assert.True(t, isForbiddenServiceEnvVar("Path"))
assert.True(t, isForbiddenServiceEnvVar("ld_preload"))
} else {
assert.False(t, isForbiddenServiceEnvVar("Path"))
assert.False(t, isForbiddenServiceEnvVar("ld_preload"))
}
}
func TestApplyServiceEnvParams_DropsForbiddenSavedNames(t *testing.T) {
origServiceEnvVars := serviceEnvVars
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
serviceEnvVars = nil
cmd := &cobra.Command{}
cmd.Flags().StringSlice("service-env", nil, "")
saved := &serviceParams{
ServiceEnvVars: map[string]string{"PATH": "C:\\attacker", "NB_LOG_FORMAT": "json"},
}
applyServiceEnvParams(cmd, saved)
result, err := parseServiceEnvVars(serviceEnvVars)
require.NoError(t, err, "a saved PATH must be dropped rather than fail the install")
assert.Equal(t, map[string]string{"NB_LOG_FORMAT": "json"}, result)
}
func TestApplyServiceEnvParams_ExplicitEmptyClears(t *testing.T) {
origServiceEnvVars := serviceEnvVars
t.Cleanup(func() { serviceEnvVars = origServiceEnvVars })
+57
View File
@@ -0,0 +1,57 @@
//go:build !windows && !ios && !android
package cmd
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/client/configs"
)
// The Windows equivalent of this is the ACL check in
// elevate.CheckOnlyOwnerWritable, covered by that package's own tests; here the
// point is that loadServiceParams asks the question at all.
func TestLoadServiceParams_RefusesWorldWritableFile(t *testing.T) {
tmpDir := t.TempDir()
original := configs.StateDir
t.Cleanup(func() { configs.StateDir = original })
configs.StateDir = tmpDir
path := filepath.Join(tmpDir, serviceParamsFile)
require.NoError(t, os.WriteFile(path, []byte(`{"log_level":"debug"}`), 0o666))
// WriteFile is subject to the umask, so set the bits that matter explicitly.
require.NoError(t, os.Chmod(path, 0o666))
params, err := loadServiceParams()
require.Error(t, err, "a service.json anyone can rewrite must not be trusted")
assert.Nil(t, params)
require.NoError(t, os.Chmod(path, 0o600))
params, err = loadServiceParams()
require.NoError(t, err)
require.NotNil(t, params)
assert.Equal(t, "debug", params.LogLevel)
}
func TestLoadServiceParams_RefusesWorldWritableDirectory(t *testing.T) {
tmpDir := t.TempDir()
stateDir := filepath.Join(tmpDir, "state")
require.NoError(t, os.Mkdir(stateDir, 0o777))
require.NoError(t, os.Chmod(stateDir, 0o777))
original := configs.StateDir
t.Cleanup(func() { configs.StateDir = original })
configs.StateDir = stateDir
require.NoError(t, os.WriteFile(filepath.Join(stateDir, serviceParamsFile), []byte(`{}`), 0o600))
params, err := loadServiceParams()
require.Error(t, err, "a service.json in a directory anyone can replace entries in must not be trusted")
assert.Nil(t, params)
}
+172 -63
View File
@@ -2,10 +2,10 @@ package cmd
import (
"context"
"errors"
"fmt"
"net"
"net/netip"
"os/user"
"runtime"
"strings"
"time"
@@ -21,6 +21,7 @@ import (
"github.com/netbirdio/netbird/client/internal"
"github.com/netbirdio/netbird/client/internal/peer"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/mdm"
nbnet "github.com/netbirdio/netbird/client/net"
"github.com/netbirdio/netbird/client/proto"
"github.com/netbirdio/netbird/client/server"
@@ -48,6 +49,8 @@ const (
profileNameDesc = "profile name to use for the login. If not specified, the last used profile will be used."
)
var errDaemonActiveProfileUnsupported = errors.New("daemon does not support active profile lookup")
var (
foregroundMode bool
dnsLabels []string
@@ -122,23 +125,25 @@ func upFunc(cmd *cobra.Command, args []string) error {
pm := profilemanager.NewProfileManager()
username, err := user.Current()
username, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %v", err)
}
var activeProf *profilemanager.Profile
var profileSwitched bool
// switch profile if provided
if profileName != "" {
if err := switchOrCreateProfile(cmd.Context(), pm, profileName, username.Username); err != nil {
activeProf, err = switchOrCreateProfile(cmd.Context(), pm, profileName, username.Username)
if err != nil {
return fmt.Errorf("switch profile: %v", err)
}
profileSwitched = true
}
activeProf, err := pm.GetActiveProfile()
if err != nil {
return fmt.Errorf("get active profile: %v", err)
} else {
activeProf, err = pm.GetActiveProfile()
if err != nil {
return fmt.Errorf("get active profile: %v", err)
}
}
if foregroundMode {
@@ -150,13 +155,15 @@ func upFunc(cmd *cobra.Command, args []string) error {
// switchOrCreateProfile switches the active profile to the one identified by
// handle, creating it first when it does not exist yet. This restores the
// pre-0.73 behaviour where `netbird up --profile <name>` auto-creates a
// missing profile instead of failing.
func switchOrCreateProfile(ctx context.Context, pm *profilemanager.ProfileManager, handle, username string) error {
// missing profile instead of failing. Returns the daemon-resolved profile so
// callers act on it directly instead of re-reading the local state, which is
// not updated under sudo.
func switchOrCreateProfile(ctx context.Context, pm *profilemanager.ProfileManager, handle, username string) (*profilemanager.Profile, error) {
resolvedID, err := switchProfile(ctx, handle, username)
if err != nil {
st, ok := gstatus.FromError(err)
if !ok || st.Code() != codes.NotFound {
return err
return nil, err
}
// Don't fail immediately on a create error: a concurrent run may
// have created the profile between the NotFound above and this
@@ -165,16 +172,16 @@ func switchOrCreateProfile(ctx context.Context, pm *profilemanager.ProfileManage
_, createErr := createProfile(ctx, handle, username)
if resolvedID, err = switchProfile(ctx, handle, username); err != nil {
if createErr != nil {
return fmt.Errorf("create profile: %w", createErr)
return nil, fmt.Errorf("create profile: %w", createErr)
}
return err
return nil, err
}
}
if err := pm.SwitchProfile(resolvedID); err != nil {
return err
return nil, err
}
return nil
return &profilemanager.Profile{ID: resolvedID}, nil
}
// createProfile dials the daemon and creates a new profile with the given
@@ -228,6 +235,10 @@ func runInForegroundMode(ctx context.Context, cmd *cobra.Command, activeProf *pr
if err != nil {
return fmt.Errorf("get config file: %v", err)
}
// CLI foreground path runs without the daemon Server: layer in the
// active MDM policy explicitly so a forced ManagementURL / PSK /
// other managed key actually takes effect on this run.
config.ApplyMDMPolicy(mdm.NewLoader(nil).Load())
_, _ = profilemanager.UpdateOldManagementURL(ctx, config, configFilePath)
@@ -302,6 +313,30 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
return fmt.Errorf("unable to get daemon status: %v", err)
}
// Under sudo the invoking user's local active-profile mirror is never
// written (the SwitchProfile write is a no-op), and plain root has no
// invoking user at all — so the mirror read into activeProf above is stale
// or defaulted and must not drive the daemon. With no --profile to make the
// choice explicit, take the profile the daemon already holds for this user
// instead: it stays on the user's current profile rather than silently
// switching to the mirror's default, and refuses when the daemon is on
// another user's profile.
if profileName == "" && !profilemanager.MirrorIsAuthoritative() {
u, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %v", err)
}
resolved, err := daemonActiveProfileForUser(ctx, client, u.Username)
switch {
case errors.Is(err, errDaemonActiveProfileUnsupported):
log.Warnf("keeping the locally resolved profile: %v", err)
case err != nil:
return err
default:
activeProf = resolved
}
}
if status.Status == string(internal.StatusConnected) {
if !profileSwitched {
cmd.Println("Already connected")
@@ -314,7 +349,7 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
}
}
username, err := user.Current()
username, err := profilemanager.InvokingUser()
if err != nil {
return fmt.Errorf("get current user: %v", err)
}
@@ -398,26 +433,21 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ
return nil
}
func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, profileName, username string) *proto.SetConfigRequest {
var req proto.SetConfigRequest
req.ProfileName = profileName
req.Username = username
req.ManagementUrl = managementURL
req.AdminURL = adminURL
req.NatExternalIPs = natExternalIPs
req.CustomDNSAddress = customDNSAddressConverted
req.ExtraIFaceBlacklist = extraIFaceBlackList
req.DnsLabels = dnsLabelsValidated.ToPunycodeList()
req.CleanDNSLabels = dnsLabels != nil && len(dnsLabels) == 0
req.CleanNATExternalIPs = natExternalIPs != nil && len(natExternalIPs) == 0
if cmd.Flag(enableRosenpassFlag).Changed {
req.RosenpassEnabled = &rosenpassEnabled
}
if cmd.Flag(rosenpassPermissiveFlag).Changed {
req.RosenpassPermissive = &rosenpassPermissive
// setBoolPtrIfChanged points dst at a copy of val when the named bool flag was
// explicitly set on cmd. It collapses the repeated
// "if cmd.Flag(x).Changed { field = &val }" pattern in the request builders into
// a single call, keeping their cognitive complexity within bounds.
func setBoolPtrIfChanged(cmd *cobra.Command, name string, dst **bool, val bool) {
if cmd.Flag(name).Changed {
dst2 := val
*dst = &dst2
}
}
// setSSHSetConfigFields copies the SSH server flags the user actually
// passed into req, leaving the rest unset so the daemon keeps the
// persisted values.
func setSSHSetConfigFields(req *proto.SetConfigRequest, cmd *cobra.Command) {
if cmd.Flag(serverSSHAllowedFlag).Changed {
req.ServerSSHAllowed = &serverSSHAllowed
}
@@ -440,6 +470,31 @@ func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, pro
sshJWTCacheTTL32 := int32(sshJWTCacheTTL)
req.SshJWTCacheTTL = &sshJWTCacheTTL32
}
}
func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, profileName, username string) *proto.SetConfigRequest {
var req proto.SetConfigRequest
req.ProfileName = profileName
req.Username = username
req.ManagementUrl = managementURL
req.AdminURL = adminURL
req.NatExternalIPs = natExternalIPs
req.CustomDNSAddress = customDNSAddressConverted
req.ExtraIFaceBlacklist = extraIFaceBlackList
req.DnsLabels = dnsLabelsValidated.ToPunycodeList()
req.CleanDNSLabels = dnsLabels != nil && len(dnsLabels) == 0
req.CleanNATExternalIPs = natExternalIPs != nil && len(natExternalIPs) == 0
if cmd.Flag(enableRosenpassFlag).Changed {
req.RosenpassEnabled = &rosenpassEnabled
}
if cmd.Flag(rosenpassPermissiveFlag).Changed {
req.RosenpassPermissive = &rosenpassPermissive
}
setSSHSetConfigFields(&req, cmd)
setBoolPtrIfChanged(cmd, remoteJobsAllowedFlag, &req.RemoteJobsAllowed, remoteJobsAllowed)
if cmd.Flag(interfaceNameFlag).Changed {
if err := parseInterfaceName(interfaceName); err != nil {
log.Errorf("parse interface name: %v", err)
@@ -499,6 +554,13 @@ func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, pro
req.DisableIpv6 = &disableIPv6
}
if cmd.Flag(enableLocalMetricsFlag).Changed {
req.EnableLocalMetrics = &localMetricsEnabled
}
if cmd.Flag(localMetricsAddressFlag).Changed {
req.LocalMetricsAddress = &localMetricsAddr
}
return &req
}
@@ -523,6 +585,7 @@ func setupConfig(customDNSAddressConverted []byte, cmd *cobra.Command, configFil
if cmd.Flag(serverSSHAllowedFlag).Changed {
ic.ServerSSHAllowed = &serverSSHAllowed
}
setBoolPtrIfChanged(cmd, remoteJobsAllowedFlag, &ic.RemoteJobsAllowed, remoteJobsAllowed)
if cmd.Flag(enableSSHRootFlag).Changed {
ic.EnableSSHRoot = &enableSSHRoot
@@ -616,9 +679,45 @@ func setupConfig(customDNSAddressConverted []byte, cmd *cobra.Command, configFil
ic.DisableIPv6 = &disableIPv6
}
if cmd.Flag(enableLocalMetricsFlag).Changed {
ic.LocalMetricsEnabled = &localMetricsEnabled
}
if cmd.Flag(localMetricsAddressFlag).Changed {
ic.LocalMetricsAddress = &localMetricsAddr
}
return &ic, nil
}
// setSSHLoginFields copies the SSH server flags the user actually passed
// into req, leaving the rest unset so the daemon keeps the persisted
// values.
func setSSHLoginFields(req *proto.LoginRequest, cmd *cobra.Command) {
if cmd.Flag(serverSSHAllowedFlag).Changed {
req.ServerSSHAllowed = &serverSSHAllowed
}
if cmd.Flag(enableSSHRootFlag).Changed {
req.EnableSSHRoot = &enableSSHRoot
}
if cmd.Flag(enableSSHSFTPFlag).Changed {
req.EnableSSHSFTP = &enableSSHSFTP
}
if cmd.Flag(enableSSHLocalPortForwardFlag).Changed {
req.EnableSSHLocalPortForwarding = &enableSSHLocalPortForward
}
if cmd.Flag(enableSSHRemotePortForwardFlag).Changed {
req.EnableSSHRemotePortForwarding = &enableSSHRemotePortForward
}
if cmd.Flag(disableSSHAuthFlag).Changed {
req.DisableSSHAuth = &disableSSHAuth
}
if cmd.Flag(sshJWTCacheTTLFlag).Changed {
sshJWTCacheTTL32 := int32(sshJWTCacheTTL)
req.SshJWTCacheTTL = &sshJWTCacheTTL32
}
}
func setupLoginRequest(providedSetupKey string, customDNSAddressConverted []byte, cmd *cobra.Command) (*proto.LoginRequest, error) {
loginRequest := proto.LoginRequest{
SetupKey: providedSetupKey,
@@ -645,39 +744,21 @@ func setupLoginRequest(providedSetupKey string, customDNSAddressConverted []byte
loginRequest.RosenpassPermissive = &rosenpassPermissive
}
if cmd.Flag(serverSSHAllowedFlag).Changed {
loginRequest.ServerSSHAllowed = &serverSSHAllowed
}
if cmd.Flag(enableSSHRootFlag).Changed {
loginRequest.EnableSSHRoot = &enableSSHRoot
}
if cmd.Flag(enableSSHSFTPFlag).Changed {
loginRequest.EnableSSHSFTP = &enableSSHSFTP
}
if cmd.Flag(enableSSHLocalPortForwardFlag).Changed {
loginRequest.EnableSSHLocalPortForwarding = &enableSSHLocalPortForward
}
if cmd.Flag(enableSSHRemotePortForwardFlag).Changed {
loginRequest.EnableSSHRemotePortForwarding = &enableSSHRemotePortForward
}
if cmd.Flag(disableSSHAuthFlag).Changed {
loginRequest.DisableSSHAuth = &disableSSHAuth
}
if cmd.Flag(sshJWTCacheTTLFlag).Changed {
sshJWTCacheTTL32 := int32(sshJWTCacheTTL)
loginRequest.SshJWTCacheTTL = &sshJWTCacheTTL32
}
setSSHLoginFields(&loginRequest, cmd)
setBoolPtrIfChanged(cmd, remoteJobsAllowedFlag, &loginRequest.RemoteJobsAllowed, remoteJobsAllowed)
if cmd.Flag(disableAutoConnectFlag).Changed {
loginRequest.DisableAutoConnect = &autoConnectDisabled
}
if cmd.Flag(enableLocalMetricsFlag).Changed {
loginRequest.EnableLocalMetrics = &localMetricsEnabled
}
if cmd.Flag(localMetricsAddressFlag).Changed {
loginRequest.LocalMetricsAddress = &localMetricsAddr
}
if cmd.Flag(interfaceNameFlag).Changed {
if err := parseInterfaceName(interfaceName); err != nil {
return nil, err
@@ -849,3 +930,31 @@ func isValidAddrPort(input string) bool {
_, err := netip.ParseAddrPort(input)
return err == nil
}
// daemonActiveProfileForUser returns the profile the daemon currently holds for
// username, for the no --profile case where the local mirror is not
// authoritative (sudo or plain root). It returns that profile when the daemon
// owns it for this user or when the profile is unowned (empty username, as on a
// fresh install), so the caller acts on the daemon's real state instead of the
// stale mirror. It denies with a --profile hint when the daemon is on another
// user's profile, when the lookup fails, or when the daemon reports no active
// profile. Returns errDaemonActiveProfileUnsupported when the daemon predates
// the RPC; the caller keeps the mirror-derived profile in that case.
func daemonActiveProfileForUser(ctx context.Context, client proto.DaemonServiceClient, username string) (*profilemanager.Profile, error) {
active, err := client.GetActiveProfile(ctx, &proto.GetActiveProfileRequest{})
if err != nil {
if st, ok := gstatus.FromError(err); ok && st.Code() == codes.Unimplemented {
return nil, fmt.Errorf("%w: %v", errDaemonActiveProfileUnsupported, err)
}
return nil, fmt.Errorf("pass --profile to choose the profile explicitly: the daemon's active profile could not be verified: %v", err)
}
if active.GetId() == "" {
return nil, fmt.Errorf("pass --profile to choose the profile explicitly: the daemon reported no active profile")
}
if active.GetUsername() != "" && active.GetUsername() != username {
return nil, fmt.Errorf(
"pass --profile to choose the profile explicitly: the daemon's active profile is %q (user %q) but this invocation runs for %q",
active.GetProfileName(), active.GetUsername(), username)
}
return &profilemanager.Profile{ID: profilemanager.ID(active.GetId())}, nil
}
+88
View File
@@ -0,0 +1,88 @@
package cmd
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
gstatus "google.golang.org/grpc/status"
"github.com/netbirdio/netbird/client/internal/profilemanager"
"github.com/netbirdio/netbird/client/proto"
)
type fakeActiveProfileClient struct {
proto.DaemonServiceClient
resp *proto.GetActiveProfileResponse
err error
}
func (f *fakeActiveProfileClient) GetActiveProfile(_ context.Context, _ *proto.GetActiveProfileRequest, _ ...grpc.CallOption) (*proto.GetActiveProfileResponse, error) {
return f.resp, f.err
}
func TestDaemonActiveProfileForUserReturnsOwnProfile(t *testing.T) {
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "default", ProfileName: "default", Username: "root"}}
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
require.NoError(t, err)
require.NotNil(t, prof)
assert.Equal(t, profilemanager.ID("default"), prof.ID)
}
func TestDaemonActiveProfileForUserReturnsUnownedProfile(t *testing.T) {
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "default", ProfileName: "default", Username: ""}}
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
require.NoError(t, err)
require.NotNil(t, prof)
assert.Equal(t, profilemanager.ID("default"), prof.ID)
}
func TestDaemonActiveProfileForUserKeepsDaemonProfileOverStaleMirror(t *testing.T) {
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "ab12", ProfileName: "work", Username: "misha"}}
prof, err := daemonActiveProfileForUser(context.Background(), client, "misha")
require.NoError(t, err)
require.NotNil(t, prof)
assert.Equal(t, profilemanager.ID("ab12"), prof.ID)
}
func TestDaemonActiveProfileForUserRejectsOtherUsersProfile(t *testing.T) {
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "ab12", ProfileName: "work", Username: "misha"}}
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
require.Error(t, err)
assert.Nil(t, prof)
assert.Contains(t, err.Error(), "--profile")
}
func TestDaemonActiveProfileForUserRejectsOtherUsersDefaultProfile(t *testing.T) {
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{Id: "default", ProfileName: "default", Username: "misha"}}
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
require.Error(t, err)
assert.Nil(t, prof)
assert.Contains(t, err.Error(), "--profile")
}
func TestDaemonActiveProfileForUserRejectsLookupError(t *testing.T) {
client := &fakeActiveProfileClient{err: gstatus.Error(codes.Internal, "boom")}
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
require.Error(t, err)
assert.Nil(t, prof)
assert.Contains(t, err.Error(), "--profile")
}
func TestDaemonActiveProfileForUserRejectsEmptyResponse(t *testing.T) {
client := &fakeActiveProfileClient{resp: &proto.GetActiveProfileResponse{}}
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
require.Error(t, err)
assert.Nil(t, prof)
assert.Contains(t, err.Error(), "--profile")
}
func TestDaemonActiveProfileForUserKeepsMirrorWhenDaemonWithoutRPC(t *testing.T) {
client := &fakeActiveProfileClient{err: gstatus.Error(codes.Unimplemented, "unknown method")}
prof, err := daemonActiveProfileForUser(context.Background(), client, "root")
require.ErrorIs(t, err, errDaemonActiveProfileUnsupported)
assert.Nil(t, prof)
}