mirror of
https://github.com/netbirdio/netbird.git
synced 2026-07-26 10:21:28 +02:00
Compare commits
13 Commits
grpc-acl
...
refactor/p
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
753925032a | ||
|
|
9be5f238da | ||
|
|
791e8b33ae | ||
|
|
917e85f648 | ||
|
|
9e75d5c732 | ||
|
|
1afc9bcac7 | ||
|
|
78f3165e85 | ||
|
|
4d76fd3c80 | ||
|
|
07ffbc9424 | ||
|
|
467b2a1712 | ||
|
|
cb4088484e | ||
|
|
7620599961 | ||
|
|
5f98524e02 |
@@ -237,7 +237,7 @@ task dev
|
||||
Pass daemon flags after `--`:
|
||||
|
||||
```
|
||||
task dev -- --daemon-addr=npipe://netbird
|
||||
task dev -- --daemon-addr=tcp://127.0.0.1:41731
|
||||
```
|
||||
|
||||
Production build (frontend assets embedded into the binary, output in `client/ui/bin/`):
|
||||
|
||||
@@ -247,9 +247,6 @@ func (c *Client) DebugBundle(platformFiles PlatformFiles, anonymize bool) (strin
|
||||
deps.SyncResponse = resp
|
||||
|
||||
if e := cc.Engine(); e != nil {
|
||||
deps.RefreshStatus = func() {
|
||||
e.RunHealthProbes(context.Background(), true)
|
||||
}
|
||||
if cm := e.GetClientMetrics(); cm != nil {
|
||||
deps.ClientMetrics = cm
|
||||
}
|
||||
|
||||
@@ -145,7 +145,7 @@ func (pm *ProfileManager) SwitchProfile(id string) error {
|
||||
// AddProfile creates a new profile
|
||||
func (pm *ProfileManager) AddProfile(profileName string) error {
|
||||
// Use ServiceManager (creates profile in profiles/ directory)
|
||||
profile, err := pm.serviceMgr.AddProfile(profileName, androidUsername, nil)
|
||||
profile, err := pm.serviceMgr.AddProfile(profileName, androidUsername)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to add profile: %w", err)
|
||||
}
|
||||
|
||||
@@ -1,58 +0,0 @@
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
// SetSSHConfigCmdName is the hidden subcommand an elevated process runs to apply
|
||||
// the privileged SSH settings.
|
||||
const SetSSHConfigCmdName = "set-ssh-config"
|
||||
|
||||
// sshConfigElevated is set by runInDaemonMode once the dangerous SSH settings
|
||||
// have been applied by an elevated helper, so the (unprivileged) up flow does
|
||||
// not re-send them and re-trip the daemon gate.
|
||||
var sshConfigElevated bool
|
||||
|
||||
// wantsDangerousSSH reports whether this invocation is trying to ENABLE SSH root
|
||||
// login or DISABLE SSH authentication.
|
||||
func wantsDangerousSSH(cmd *cobra.Command) bool {
|
||||
return (cmd.Flags().Changed(enableSSHRootFlag) && enableSSHRoot) ||
|
||||
(cmd.Flags().Changed(disableSSHAuthFlag) && disableSSHAuth)
|
||||
}
|
||||
|
||||
// buildSetSSHConfigArgs builds the argument list for an elevated
|
||||
// `netbird set-ssh-config` invocation.
|
||||
func buildSetSSHConfigArgs(profileName, username string, enableSSHRoot, disableSSHAuth *bool, daemonAddr string) []string {
|
||||
args := []string{SetSSHConfigCmdName}
|
||||
if profileName != "" {
|
||||
args = append(args, "--profile", profileName)
|
||||
}
|
||||
if username != "" {
|
||||
args = append(args, "--username", username)
|
||||
}
|
||||
if enableSSHRoot != nil && *enableSSHRoot {
|
||||
args = append(args, "--"+enableSSHRootFlag)
|
||||
}
|
||||
if disableSSHAuth != nil && *disableSSHAuth {
|
||||
args = append(args, "--"+disableSSHAuthFlag)
|
||||
}
|
||||
if daemonAddr != "" {
|
||||
args = append(args, "--daemon-addr", daemonAddr)
|
||||
}
|
||||
return args
|
||||
}
|
||||
|
||||
// ElevateSSHConfig re-runs the current executable's `set-ssh-config` command with
|
||||
// root/administrator privileges via the platform's prompt (pkexec/UAC/osascript),
|
||||
// so the elevated process connects to the daemon with a privileged identity and
|
||||
// the daemon's requirePrivilegedForDangerousSSH gate passes.
|
||||
func ElevateSSHConfig(profileName, username string, enableSSHRoot, disableSSHAuth *bool) error {
|
||||
exe, err := os.Executable()
|
||||
if err != nil {
|
||||
return fmt.Errorf("locate executable for elevation: %w", err)
|
||||
}
|
||||
return runElevated(exe, buildSetSSHConfigArgs(profileName, username, enableSSHRoot, disableSSHAuth, daemonAddr))
|
||||
}
|
||||
@@ -1,16 +0,0 @@
|
||||
//go:build darwin
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
)
|
||||
|
||||
// isProcessPrivileged reports whether the current process runs as root.
|
||||
func isProcessPrivileged() bool { return os.Geteuid() == 0 }
|
||||
|
||||
// TODO(ssh-elevation): implement the osascript admin prompt.
|
||||
func runElevated(_ string, _ []string) error {
|
||||
return fmt.Errorf("automatic privilege elevation is not yet implemented on macOS, re-run with sudo")
|
||||
}
|
||||
@@ -1,38 +0,0 @@
|
||||
//go:build linux
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// isProcessPrivileged reports whether the current process runs as root.
|
||||
func isProcessPrivileged() bool { return os.Geteuid() == 0 }
|
||||
|
||||
// runElevated re-runs exe with args as root. On a graphical session it uses
|
||||
// pkexec, which drives the desktop's polkit authentication agent (GUI prompt).
|
||||
func runElevated(exe string, args []string) error {
|
||||
if !hasGraphicalSession() {
|
||||
return fmt.Errorf("cannot request privilege elevation without a graphical session. re-run as root: sudo %s %s", exe, strings.Join(args, " "))
|
||||
}
|
||||
pkexec, err := exec.LookPath("pkexec")
|
||||
if err != nil {
|
||||
return fmt.Errorf("pkexec not found for privilege elevation, re-run as root: sudo %s %s", exe, strings.Join(args, " "))
|
||||
}
|
||||
|
||||
c := exec.Command(pkexec, append([]string{exe}, args...)...)
|
||||
c.Stdin, c.Stdout, c.Stderr = os.Stdin, os.Stdout, os.Stderr
|
||||
if err := c.Run(); err != nil {
|
||||
return fmt.Errorf("elevated %s failed (elevation cancelled or denied?): %w", SetSSHConfigCmdName, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// hasGraphicalSession heuristically reports whether a desktop session is present
|
||||
// that polkit can prompt in.
|
||||
func hasGraphicalSession() bool {
|
||||
return os.Getenv("DISPLAY") != "" || os.Getenv("WAYLAND_DISPLAY") != ""
|
||||
}
|
||||
@@ -1,17 +0,0 @@
|
||||
//go:build !linux && !darwin && !windows
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// isProcessPrivileged reports whether the current process runs as root.
|
||||
func isProcessPrivileged() bool { return os.Geteuid() == 0 }
|
||||
|
||||
// runElevated has no automatic elevation mechanism on these platforms
|
||||
func runElevated(exe string, args []string) error {
|
||||
return fmt.Errorf("automatic privilege elevation is not supported on this platform, re-run as root: sudo %s %s", exe, strings.Join(args, " "))
|
||||
}
|
||||
@@ -1,99 +0,0 @@
|
||||
//go:build !ios && !android
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestBuildSetSSHConfigArgs(t *testing.T) {
|
||||
tr, fa := true, false
|
||||
|
||||
t.Run("enable root only, with daemon-addr", func(t *testing.T) {
|
||||
got := buildSetSSHConfigArgs("prof", "alice", &tr, nil, "unix:///x.sock")
|
||||
assert.Equal(t, []string{
|
||||
SetSSHConfigCmdName, "--profile", "prof", "--username", "alice",
|
||||
"--" + enableSSHRootFlag, "--daemon-addr", "unix:///x.sock",
|
||||
}, got)
|
||||
})
|
||||
|
||||
t.Run("disable auth only, no daemon-addr", func(t *testing.T) {
|
||||
got := buildSetSSHConfigArgs("prof", "alice", nil, &tr, "")
|
||||
assert.Equal(t, []string{
|
||||
SetSSHConfigCmdName, "--profile", "prof", "--username", "alice",
|
||||
"--" + disableSSHAuthFlag,
|
||||
}, got)
|
||||
})
|
||||
|
||||
t.Run("false pointers omit the flags", func(t *testing.T) {
|
||||
got := buildSetSSHConfigArgs("", "", &fa, &fa, "")
|
||||
assert.Equal(t, []string{SetSSHConfigCmdName}, got)
|
||||
})
|
||||
|
||||
t.Run("both enabled", func(t *testing.T) {
|
||||
got := buildSetSSHConfigArgs("p", "u", &tr, &tr, "")
|
||||
assert.Equal(t, []string{
|
||||
SetSSHConfigCmdName, "--profile", "p", "--username", "u",
|
||||
"--" + enableSSHRootFlag, "--" + disableSSHAuthFlag,
|
||||
}, got)
|
||||
})
|
||||
}
|
||||
|
||||
func TestBuildSetSSHConfigRequest(t *testing.T) {
|
||||
tr := true
|
||||
|
||||
req := buildSetSSHConfigRequest("p", "u", &tr, nil)
|
||||
assert.Equal(t, "p", req.ProfileName)
|
||||
assert.Equal(t, "u", req.Username)
|
||||
if assert.NotNil(t, req.EnableSSHRoot) {
|
||||
assert.True(t, *req.EnableSSHRoot)
|
||||
}
|
||||
assert.Nil(t, req.DisableSSHAuth, "an unset flag must leave the daemon value untouched")
|
||||
}
|
||||
|
||||
func TestWantsDangerousSSH(t *testing.T) {
|
||||
origRoot, origAuth := enableSSHRoot, disableSSHAuth
|
||||
t.Cleanup(func() { enableSSHRoot, disableSSHAuth = origRoot, origAuth })
|
||||
|
||||
newCmd := func() *cobra.Command {
|
||||
enableSSHRoot, disableSSHAuth = false, false
|
||||
c := &cobra.Command{Use: "x"}
|
||||
c.Flags().BoolVar(&enableSSHRoot, enableSSHRootFlag, false, "")
|
||||
c.Flags().BoolVar(&disableSSHAuth, disableSSHAuthFlag, false, "")
|
||||
return c
|
||||
}
|
||||
|
||||
// wantsDangerousSSH fires only in the privileged direction.
|
||||
t.Run("enable root true is dangerous", func(t *testing.T) {
|
||||
c := newCmd()
|
||||
require.NoError(t, c.Flags().Set(enableSSHRootFlag, "true"))
|
||||
assert.True(t, wantsDangerousSSH(c))
|
||||
})
|
||||
|
||||
t.Run("enable root false is not dangerous", func(t *testing.T) {
|
||||
c := newCmd()
|
||||
require.NoError(t, c.Flags().Set(enableSSHRootFlag, "false"))
|
||||
assert.False(t, wantsDangerousSSH(c))
|
||||
})
|
||||
|
||||
t.Run("disable auth true is dangerous", func(t *testing.T) {
|
||||
c := newCmd()
|
||||
require.NoError(t, c.Flags().Set(disableSSHAuthFlag, "true"))
|
||||
assert.True(t, wantsDangerousSSH(c))
|
||||
})
|
||||
|
||||
t.Run("disable auth false is not dangerous", func(t *testing.T) {
|
||||
c := newCmd()
|
||||
require.NoError(t, c.Flags().Set(disableSSHAuthFlag, "false"))
|
||||
assert.False(t, wantsDangerousSSH(c))
|
||||
})
|
||||
|
||||
t.Run("nothing changed is not dangerous", func(t *testing.T) {
|
||||
c := newCmd()
|
||||
assert.False(t, wantsDangerousSSH(c))
|
||||
})
|
||||
}
|
||||
@@ -1,23 +0,0 @@
|
||||
//go:build windows
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
// isProcessPrivileged reports whether the current process token is elevated
|
||||
// (running as administrator / LocalSystem).
|
||||
func isProcessPrivileged() bool {
|
||||
return windows.GetCurrentProcessToken().IsElevated()
|
||||
}
|
||||
|
||||
// runElevated should re-run exe with args elevated via a UAC prompt
|
||||
// (ShellExecuteEx with the "runas" verb).
|
||||
//
|
||||
// TODO(ssh-elevation): implement ShellExecuteEx("runas") + wait for the child.
|
||||
func runElevated(_ string, _ []string) error {
|
||||
return fmt.Errorf("automatic privilege elevation is not yet implemented on Windows, re-run netbird as administrator")
|
||||
}
|
||||
@@ -17,9 +17,7 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal"
|
||||
"github.com/netbirdio/netbird/client/internal/auth"
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
nbnet "github.com/netbirdio/netbird/client/net"
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
"github.com/netbirdio/netbird/client/server"
|
||||
"github.com/netbirdio/netbird/client/system"
|
||||
"github.com/netbirdio/netbird/util"
|
||||
)
|
||||
@@ -333,14 +331,6 @@ func doForegroundLogin(ctx context.Context, cmd *cobra.Command, setupKey string,
|
||||
return fmt.Errorf("read config file %s: %v", configFilePath, err)
|
||||
}
|
||||
|
||||
// Mirror runInForegroundMode: recover residual state (DNS, firewall,
|
||||
// ssh config, legacy routing) from a previous unclean shutdown and
|
||||
// enable advanced routing before dialing management.
|
||||
if err := server.RestoreResidualState(ctx, profilemanager.NewServiceManager(configFilePath).GetStatePath()); err != nil {
|
||||
log.Warnf("failed to restore residual state: %v", err)
|
||||
}
|
||||
nbnet.Init()
|
||||
|
||||
err = foregroundLogin(ctx, cmd, config, setupKey, activeProf.ID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("foreground login failed: %v", err)
|
||||
|
||||
@@ -1,116 +0,0 @@
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
"github.com/netbirdio/netbird/util"
|
||||
)
|
||||
|
||||
var ownerCmd = &cobra.Command{
|
||||
Use: "owner",
|
||||
Short: "Manage who may control the NetBird daemon",
|
||||
Long: `Manage the daemon-wide owners.
|
||||
|
||||
Owners are enforced on the daemon and stored in the service parameters. All owners
|
||||
may control the daemon and use the shared default profile (plus root/administrator),
|
||||
every other profile stays isolated to the user that created it. An unowned daemon
|
||||
is claimed by the first caller (trust-on-first-use).`,
|
||||
}
|
||||
|
||||
var ownerAddCmd = &cobra.Command{
|
||||
Use: "add <principal>",
|
||||
Short: "Add a daemon owner principal",
|
||||
Long: `Add a daemon-wide owner principal. Principals are typed:
|
||||
uid:1000 a Unix user ID
|
||||
gid:1000 a Unix group ID
|
||||
group:netbird-admins a Unix group name (resolved via NSS/getent)
|
||||
sid:S-1-5-21-... a Windows user or group SID
|
||||
|
||||
Requires root/administrator or an existing owner.`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return withDaemon(cmd, func(ctx context.Context, c proto.DaemonServiceClient) error {
|
||||
if _, err := c.AddOwner(ctx, &proto.AddOwnerRequest{Principal: args[0]}); err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.Printf("Added daemon owner %q\n", args[0])
|
||||
return nil
|
||||
})
|
||||
},
|
||||
}
|
||||
|
||||
var ownerResetCmd = &cobra.Command{
|
||||
Use: "reset",
|
||||
Short: "Clear the daemon owner list (root/administrator only)",
|
||||
Long: `Clear the daemon-wide owner list, returning the daemon to the unowned
|
||||
state. The next caller then claims ownership (trust-on-first-use). Requires
|
||||
root/administrator.`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return withDaemon(cmd, func(ctx context.Context, c proto.DaemonServiceClient) error {
|
||||
if _, err := c.ResetOwner(ctx, &proto.ResetOwnerRequest{}); err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.Println("Daemon owner list cleared, the next caller will claim ownership")
|
||||
return nil
|
||||
})
|
||||
},
|
||||
}
|
||||
|
||||
var ownerShareCmd = &cobra.Command{
|
||||
Use: "share",
|
||||
Short: "Mark the daemon shared (any local user may control it)",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return withDaemon(cmd, func(ctx context.Context, c proto.DaemonServiceClient) error {
|
||||
if _, err := c.ShareProfile(ctx, &proto.ShareProfileRequest{Shared: true}); err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.Println("Daemon is now shared with all local users")
|
||||
return nil
|
||||
})
|
||||
},
|
||||
}
|
||||
|
||||
var ownerUnshareCmd = &cobra.Command{
|
||||
Use: "unshare",
|
||||
Short: "Stop sharing the daemon (restrict to its owners)",
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return withDaemon(cmd, func(ctx context.Context, c proto.DaemonServiceClient) error {
|
||||
if _, err := c.ShareProfile(ctx, &proto.ShareProfileRequest{Shared: false}); err != nil {
|
||||
return err
|
||||
}
|
||||
cmd.Println("Daemon is no longer shared")
|
||||
return nil
|
||||
})
|
||||
},
|
||||
}
|
||||
|
||||
// withDaemon runs fn with a connected daemon client, handling setup and teardown.
|
||||
func withDaemon(cmd *cobra.Command, fn func(context.Context, proto.DaemonServiceClient) error) error {
|
||||
SetFlagsFromEnvVars(rootCmd)
|
||||
cmd.SetOut(cmd.OutOrStdout())
|
||||
if err := util.InitLog(logLevel, util.LogConsole); err != nil {
|
||||
log.Errorf("failed initializing log %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(cmd.Context(), 20*time.Second)
|
||||
defer cancel()
|
||||
|
||||
conn, err := DialClientGRPCServer(ctx, daemonAddr)
|
||||
if err != nil {
|
||||
log.Errorf("failed to connect to service CLI interface %v", err)
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
if cerr := conn.Close(); cerr != nil {
|
||||
log.Debugf("close daemon connection: %v", cerr)
|
||||
}
|
||||
}()
|
||||
|
||||
return fn(ctx, proto.NewDaemonServiceClient(conn))
|
||||
}
|
||||
@@ -117,12 +117,10 @@ func listProfilesFunc(cmd *cobra.Command, _ []string) error {
|
||||
} else {
|
||||
fmt.Fprintln(tw, "NAME\tACTIVE")
|
||||
}
|
||||
anyActive := false
|
||||
for _, profile := range resp.Profiles {
|
||||
marker := ""
|
||||
if profile.IsActive {
|
||||
marker = "✓"
|
||||
anyActive = true
|
||||
}
|
||||
name := profilemanager.StripCtrlChars(profile.Name)
|
||||
id := profilemanager.ID(profile.Id)
|
||||
@@ -132,19 +130,7 @@ func listProfilesFunc(cmd *cobra.Command, _ []string) error {
|
||||
fmt.Fprintf(tw, "%s\t%s\n", name, marker)
|
||||
}
|
||||
}
|
||||
if err := tw.Flush(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// None of the caller's profiles is active: another user may hold the daemon.
|
||||
// Surface it so the empty ACTIVE column is not mistaken for "nothing active".
|
||||
if !anyActive {
|
||||
if active, aerr := daemonClient.GetActiveProfile(cmd.Context(), &proto.GetActiveProfileRequest{}); aerr == nil && active.GetUsername() != "" && !usernamesMatch(active.GetUsername(), currUser.Username) {
|
||||
cmd.Printf("\nActive profile belongs to another user: %s (user %s)\n",
|
||||
profilemanager.StripCtrlChars(active.GetProfileName()), active.GetUsername())
|
||||
}
|
||||
}
|
||||
return nil
|
||||
return tw.Flush()
|
||||
}
|
||||
|
||||
func addProfileFunc(cmd *cobra.Command, args []string) error {
|
||||
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"net"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path"
|
||||
@@ -144,10 +143,10 @@ func init() {
|
||||
|
||||
defaultDaemonAddr := "unix:///var/run/netbird.sock"
|
||||
if runtime.GOOS == "windows" {
|
||||
defaultDaemonAddr = windowsPipeDaemonAddr
|
||||
defaultDaemonAddr = "tcp://127.0.0.1:41731"
|
||||
}
|
||||
|
||||
rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp|npipe]://[path|host:port|name]")
|
||||
rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp]://[path|host:port]")
|
||||
rootCmd.PersistentFlags().StringVarP(&managementURL, "management-url", "m", "", fmt.Sprintf("Management Service URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultManagementURL))
|
||||
rootCmd.PersistentFlags().StringVar(&adminURL, "admin-url", "", fmt.Sprintf("Admin Panel URL [http|https]://[host]:[port] (default \"%s\")", profilemanager.DefaultAdminURL))
|
||||
rootCmd.PersistentFlags().StringVarP(&logLevel, "log-level", "l", "info", "sets NetBird log level")
|
||||
@@ -173,11 +172,6 @@ func init() {
|
||||
rootCmd.AddCommand(profileCmd)
|
||||
rootCmd.AddCommand(exposeCmd)
|
||||
|
||||
rootCmd.AddCommand(ownerCmd)
|
||||
ownerCmd.AddCommand(ownerAddCmd, ownerResetCmd, ownerShareCmd, ownerUnshareCmd)
|
||||
|
||||
rootCmd.AddCommand(setSSHConfigCmd)
|
||||
|
||||
networksCMD.AddCommand(routesListCmd)
|
||||
networksCMD.AddCommand(routesSelectCmd, routesDeselectCmd)
|
||||
|
||||
@@ -270,32 +264,17 @@ func FlagNameToEnvVar(cmdFlag string, prefix string) string {
|
||||
return prefix + upper
|
||||
}
|
||||
|
||||
// daemonDialTarget returns the gRPC dial target and base options for the daemon
|
||||
// address, handling the npipe scheme (Windows named pipe, via a context dialer)
|
||||
// and unix/tcp.
|
||||
func daemonDialTarget(addr string) (string, []grpc.DialOption) {
|
||||
opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())}
|
||||
target := strings.TrimPrefix(addr, "tcp://")
|
||||
if strings.HasPrefix(addr, "npipe://") {
|
||||
path := pipePath(strings.TrimPrefix(addr, "npipe://"))
|
||||
opts = append(opts, grpc.WithContextDialer(func(dialCtx context.Context, _ string) (net.Conn, error) {
|
||||
return dialNamedPipe(dialCtx, path)
|
||||
}))
|
||||
target = "passthrough:///netbird-daemon-pipe"
|
||||
}
|
||||
return target, opts
|
||||
}
|
||||
|
||||
// DialClientGRPCServer returns client connection to the daemon server.
|
||||
func DialClientGRPCServer(ctx context.Context, addr string, opts ...grpc.DialOption) (*grpc.ClientConn, error) {
|
||||
func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, error) {
|
||||
ctx, cancel := context.WithTimeout(ctx, time.Second*10)
|
||||
defer cancel()
|
||||
|
||||
target, dialOpts := daemonDialTarget(addr)
|
||||
dialOpts = append(dialOpts, grpc.WithBlock())
|
||||
dialOpts = append(dialOpts, opts...)
|
||||
|
||||
return grpc.DialContext(ctx, target, dialOpts...)
|
||||
return grpc.DialContext(
|
||||
ctx,
|
||||
strings.TrimPrefix(addr, "tcp://"),
|
||||
grpc.WithTransportCredentials(insecure.NewCredentials()),
|
||||
grpc.WithBlock(),
|
||||
)
|
||||
}
|
||||
|
||||
// WithBackOff execute function in backoff cycle.
|
||||
|
||||
@@ -30,12 +30,6 @@ var (
|
||||
serviceEnvVars []string
|
||||
jsonSocket string
|
||||
enableJSONSocket bool
|
||||
// owners seeds the daemon-wide owner set at install time (--owner). At runtime
|
||||
// the daemon reads and writes owners in service.json directly.
|
||||
owners []string
|
||||
// daemonShared carries the persisted daemon shared flag across
|
||||
// install/reconfigure round-trips (set at runtime via `netbird owner share`).
|
||||
daemonShared bool
|
||||
)
|
||||
|
||||
type program struct {
|
||||
@@ -60,8 +54,7 @@ func init() {
|
||||
serviceCmd.PersistentFlags().BoolVar(&captureEnabled, "enable-capture", false, "Enables packet capture via 'netbird debug capture'. To persist, use: netbird service install --enable-capture")
|
||||
serviceCmd.PersistentFlags().BoolVar(&networksDisabled, "disable-networks", false, "Disables network selection. If enabled, the client will not allow listing, selecting, or deselecting networks. To persist, use: netbird service install --disable-networks")
|
||||
serviceCmd.PersistentFlags().BoolVar(&enableJSONSocket, "enable-json-socket", false, "Enables the HTTP/JSON API socket served by grpc-gateway. To persist, use: netbird service install --enable-json-socket")
|
||||
serviceCmd.PersistentFlags().StringVar(&jsonSocket, "json-socket", defaultJSONSocket, "HTTP/JSON API socket address [unix|tcp|npipe]://[path|host:port|name]. Requires --enable-json-socket to serve. To persist, use: netbird service install --enable-json-socket --json-socket")
|
||||
serviceCmd.PersistentFlags().StringSliceVar(&owners, "owner", nil, "Principal(s) allowed to control the daemon and its default profile: uid:1000, gid:1000, group:netbird-admins (NSS), or sid:S-1-5-... (Windows). Repeatable. Other profiles stay isolated per user. To persist: netbird service install --owner uid:1000")
|
||||
serviceCmd.PersistentFlags().StringVar(&jsonSocket, "json-socket", defaultJSONSocket, "HTTP/JSON API socket address [unix|tcp]://[path|host:port]. Requires --enable-json-socket to serve. To persist, use: netbird service install --enable-json-socket --json-socket")
|
||||
|
||||
rootCmd.PersistentFlags().StringVarP(&serviceName, "service", "s", defaultServiceName, "Netbird system service name")
|
||||
serviceEnvDesc := `Sets extra environment variables for the service. ` +
|
||||
|
||||
@@ -5,7 +5,6 @@ package cmd
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"runtime"
|
||||
"time"
|
||||
|
||||
"github.com/kardianos/service"
|
||||
@@ -14,32 +13,12 @@ import (
|
||||
"github.com/spf13/cobra"
|
||||
"google.golang.org/grpc"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
"github.com/netbirdio/netbird/client/server"
|
||||
"github.com/netbirdio/netbird/client/system"
|
||||
"github.com/netbirdio/netbird/util"
|
||||
)
|
||||
|
||||
// daemonServerOptions installs peer-identity transport credentials and the
|
||||
// authorization interceptor on the daemon ipc if supported.
|
||||
func daemonServerOptions(network string, interceptor *ipcauth.Interceptor) []grpc.ServerOption {
|
||||
creds := ipcauth.NewTransportCredentials()
|
||||
if creds == nil {
|
||||
log.Warnf("daemon ipc has no peer-identity primitive on %s, per-caller authorization is disabled", runtime.GOOS)
|
||||
return nil
|
||||
}
|
||||
if network == "tcp" {
|
||||
log.Warnf("daemon is listening on TCP (%s), peer identity cannot be authenticated over TCP, per-caller authorization is disabled", daemonAddr)
|
||||
return nil
|
||||
}
|
||||
return []grpc.ServerOption{
|
||||
grpc.Creds(creds),
|
||||
grpc.ChainUnaryInterceptor(interceptor.UnaryServerInterceptor()),
|
||||
grpc.ChainStreamInterceptor(interceptor.StreamServerInterceptor()),
|
||||
}
|
||||
}
|
||||
|
||||
func validateJSONSocketFlags() error {
|
||||
if serviceCmd.PersistentFlags().Changed("json-socket") && !enableJSONSocket {
|
||||
return fmt.Errorf("--json-socket requires --enable-json-socket to configure the daemon JSON gateway")
|
||||
@@ -58,28 +37,8 @@ func (p *program) Start(svc service.Service) error {
|
||||
// Collect static system and platform information
|
||||
system.UpdateStaticInfoAsync()
|
||||
|
||||
// A daemon installed before named-pipe support uses the old loopback-TCP
|
||||
// address as the daemon address. We migrate to a named pipe so an
|
||||
// upgraded daemon enforces per-caller authorization instead of silently
|
||||
// running on identity-less TCP.
|
||||
if migrated, ok := migrateLegacyDaemonAddr(daemonAddr); ok {
|
||||
log.Infof("legacy daemon address %q predates named-pipe support. listening on %q so per-caller authorization is enforced", daemonAddr, migrated)
|
||||
daemonAddr = migrated
|
||||
}
|
||||
|
||||
network, _, err := parseListenAddress(daemonAddr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse daemon address: %w", err)
|
||||
}
|
||||
|
||||
// Owner-authorization interceptor. The ConfigAdapter is a lazy bridge: the
|
||||
// gRPC server is built before the daemon server instance exists, so we set
|
||||
// the real policy backend below once serverInstance is created.
|
||||
ownerAdapter := &ipcauth.ConfigAdapter{}
|
||||
authInterceptor := ipcauth.NewInterceptor(ownerAdapter, ipcauth.NewDefaultGroupResolver())
|
||||
|
||||
// in any case, even if configuration does not exist we run daemon to serve the CLI gRPC API.
|
||||
p.serv = grpc.NewServer(daemonServerOptions(network, authInterceptor)...)
|
||||
// in any case, even if configuration does not exists we run daemon to serve CLI gRPC API.
|
||||
p.serv = grpc.NewServer()
|
||||
|
||||
daemonListener, err := listenOnAddress(daemonAddr)
|
||||
if err != nil {
|
||||
@@ -115,13 +74,9 @@ func (p *program) Start(svc service.Service) error {
|
||||
}
|
||||
|
||||
serverInstance := server.New(p.ctx, util.FindFirstLogPath(logFiles), configPath, profilesDisabled, updateSettingsDisabled, captureEnabled, networksDisabled)
|
||||
// Daemon-wide owners live in service.json (governs the default profile and
|
||||
// daemon access), wire persistence before serving so owner add / TOFU work.
|
||||
serverInstance.SetDaemonOwnerStore(daemonOwnerStore{})
|
||||
if err := serverInstance.Start(); err != nil {
|
||||
log.Fatalf("failed to start daemon: %v", err)
|
||||
}
|
||||
ownerAdapter.SetBackend(serverInstance)
|
||||
proto.RegisterDaemonServiceServer(p.serv, serverInstance)
|
||||
|
||||
p.serverInstanceMu.Lock()
|
||||
@@ -129,7 +84,6 @@ func (p *program) Start(svc service.Service) error {
|
||||
p.serverInstanceMu.Unlock()
|
||||
|
||||
if jsonListener != nil {
|
||||
log.Warnf("JSON gateway (--enable-json-socket) re-dials the daemon locally. The HTTP client's identity is forwarded so per-caller authorization still applies, but restrict access to %s appropriately", jsonSocket)
|
||||
if err := p.startJSONGateway(jsonListener, daemonAddr); err != nil {
|
||||
log.Fatalf("failed to start daemon JSON server: %v", err)
|
||||
}
|
||||
|
||||
@@ -5,62 +5,27 @@ package cmd
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
)
|
||||
|
||||
// jsonPeerCtxKey keys the HTTP client's kernel identity in the request context.
|
||||
type jsonPeerCtxKey struct{}
|
||||
|
||||
// jsonConnContext reads the connecting HTTP client's identity from the JSON
|
||||
// socket and stashes it so it can be forwarded to the daemon. The gateway
|
||||
// re-dials the daemon as the daemon's own identity, so without this the
|
||||
// daemon would see every JSON request as privileged.
|
||||
func jsonConnContext(ctx context.Context, c net.Conn) context.Context {
|
||||
id, err := ipcauth.ConnIdentity(c)
|
||||
if err != nil {
|
||||
log.Warnf("json gateway: cannot read HTTP client identity, requests won't carry it: %v", err)
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, jsonPeerCtxKey{}, id)
|
||||
}
|
||||
|
||||
// jsonForwardIdentity injects the stashed HTTP client identity as gRPC metadata
|
||||
// on the gateway's re-dial to the daemon. The daemon trusts it only because the
|
||||
// dial arrives as the daemon's own (self/privileged) identity.
|
||||
func jsonForwardIdentity(ctx context.Context, _ *http.Request) metadata.MD {
|
||||
id, ok := ctx.Value(jsonPeerCtxKey{}).(ipcauth.Identity)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return ipcauth.ForwardIdentityMetadata(id)
|
||||
func grpcGatewayEndpoint(addr string) string {
|
||||
return strings.TrimPrefix(addr, "tcp://")
|
||||
}
|
||||
|
||||
func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint string) error {
|
||||
if jsonListener.network == "tcp" {
|
||||
log.Warnf("json daemon is listening on TCP (%s), peer identity cannot be authenticated over TCP, per-caller authorization is disabled", daemonAddr)
|
||||
}
|
||||
|
||||
mux := runtime.NewServeMux(runtime.WithMetadata(jsonForwardIdentity))
|
||||
|
||||
// Lazy client to the daemon, npipe-aware (grpc.NewClient does not connect
|
||||
// until the first request, so this does not block startup before Serve).
|
||||
target, opts := daemonDialTarget(daemonEndpoint)
|
||||
conn, err := grpc.NewClient(target, opts...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("create daemon client for JSON gateway: %w", err)
|
||||
}
|
||||
if err := proto.RegisterDaemonServiceHandler(p.ctx, mux, conn); err != nil {
|
||||
mux := runtime.NewServeMux()
|
||||
opts := []grpc.DialOption{grpc.WithTransportCredentials(insecure.NewCredentials())}
|
||||
if err := proto.RegisterDaemonServiceHandlerFromEndpoint(p.ctx, mux, grpcGatewayEndpoint(daemonEndpoint), opts); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -70,7 +35,6 @@ func (p *program) startJSONGateway(jsonListener *socketListener, daemonEndpoint
|
||||
BaseContext: func(net.Listener) context.Context {
|
||||
return p.ctx
|
||||
},
|
||||
ConnContext: jsonConnContext,
|
||||
}
|
||||
|
||||
p.jsonServMu.Lock()
|
||||
|
||||
@@ -33,15 +33,6 @@ type serviceParams struct {
|
||||
DisableNetworks bool `json:"disable_networks,omitempty"`
|
||||
EnableJSONSocket bool `json:"enable_json_socket,omitempty"`
|
||||
ServiceEnvVars map[string]string `json:"service_env_vars,omitempty"`
|
||||
// Owners lists the principals allowed to control this profile over the local
|
||||
// IPC, as typed strings: "uid:1000", "gid:1000", "group:netbird-admins"
|
||||
// (Unix, NSS-resolved) or "sid:S-1-5-..." (Windows user or group SID). Empty
|
||||
// with Shared=false means the profile is owned by nobody yet, until claimed
|
||||
Owners []string `json:"owners,omitempty"`
|
||||
|
||||
// Shared, when true, lets any authenticated local caller control this profile
|
||||
// (opt-in). Takes precedence over Owners.
|
||||
Shared bool `json:"shared,omitempty"`
|
||||
}
|
||||
|
||||
// serviceParamsPath returns the path to the service params file.
|
||||
@@ -49,38 +40,6 @@ func serviceParamsPath() string {
|
||||
return filepath.Join(configs.StateDir, serviceParamsFile)
|
||||
}
|
||||
|
||||
// daemonOwnerStore persists the daemon-wide owner set in service.json. It
|
||||
// implements server.DaemonOwnerStore so the daemon can read owners at startup and
|
||||
// mutate them at runtime (owner add, reset, share, TOFU claim) without server
|
||||
// importing cmd. Load-modify-write preserves the other service.json fields.
|
||||
type daemonOwnerStore struct{}
|
||||
|
||||
func (daemonOwnerStore) Load() ([]string, bool, error) {
|
||||
params, err := loadServiceParams()
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
if params == nil {
|
||||
return nil, false, nil
|
||||
}
|
||||
return params.Owners, params.Shared, nil
|
||||
}
|
||||
|
||||
func (daemonOwnerStore) Save(owners []string, shared bool) error {
|
||||
params, err := loadServiceParams()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if params == nil {
|
||||
// No service.json yet (daemon started without `service install`). Seed it
|
||||
// from the running daemon's current parameters so the file stays complete.
|
||||
params = currentServiceParams()
|
||||
}
|
||||
params.Owners = owners
|
||||
params.Shared = shared
|
||||
return saveServiceParams(params)
|
||||
}
|
||||
|
||||
// loadServiceParams reads saved service parameters from disk.
|
||||
// Returns nil with no error if the file does not exist.
|
||||
func loadServiceParams() (*serviceParams, error) {
|
||||
@@ -127,8 +86,6 @@ func currentServiceParams() *serviceParams {
|
||||
EnableCapture: captureEnabled,
|
||||
DisableNetworks: networksDisabled,
|
||||
EnableJSONSocket: enableJSONSocket,
|
||||
Owners: owners,
|
||||
Shared: daemonShared,
|
||||
}
|
||||
|
||||
if len(serviceEnvVars) > 0 {
|
||||
@@ -168,10 +125,6 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
|
||||
|
||||
if !rootCmd.PersistentFlags().Changed("daemon-addr") && params.DaemonAddr != "" {
|
||||
daemonAddr = params.DaemonAddr
|
||||
if migrated, ok := migrateLegacyDaemonAddr(daemonAddr); ok {
|
||||
cmd.Printf("Migrating saved daemon address %q to %q so per-caller authorization can be enforced\n", daemonAddr, migrated)
|
||||
daemonAddr = migrated
|
||||
}
|
||||
}
|
||||
|
||||
if !serviceCmd.PersistentFlags().Changed("json-socket") && params.JSONSocket != "" {
|
||||
@@ -212,14 +165,6 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
|
||||
networksDisabled = params.DisableNetworks
|
||||
}
|
||||
|
||||
// Carry the daemon-wide owner set forward across install/reconfigure so a
|
||||
// runtime owner add or TOFU claim in service.json is not clobbered. --owner
|
||||
// overrides.
|
||||
if !serviceCmd.PersistentFlags().Changed("owner") && len(params.Owners) > 0 {
|
||||
owners = params.Owners
|
||||
}
|
||||
daemonShared = params.Shared
|
||||
|
||||
applyServiceEnvParams(cmd, params)
|
||||
}
|
||||
|
||||
|
||||
@@ -431,15 +431,9 @@ func TestServiceParams_BuildArgsCoversAllFlags(t *testing.T) {
|
||||
installerFile, err := parser.ParseFile(fset, "service_installer.go", nil, 0)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Fields that are handled outside of buildServiceArguments.
|
||||
// - ServiceEnvVars flows through newSVCConfig() EnvVars, not CLI args.
|
||||
// - Owners/Shared are daemon-wide ownership persisted in service.json and
|
||||
// read+mutated by the daemon at runtime (owner add / TOFU claim); they are
|
||||
// deliberately NOT baked into the run args so runtime changes are not lost.
|
||||
// Fields that are handled outside of buildServiceArguments (env vars go through newSVCConfig).
|
||||
fieldsNotInArgs := map[string]bool{
|
||||
"ServiceEnvVars": true,
|
||||
"Owners": true,
|
||||
"Shared": true,
|
||||
}
|
||||
|
||||
buildFields := extractFuncGlobalRefs(t, installerFile, "buildServiceArguments")
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
//go:build !windows
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"runtime"
|
||||
)
|
||||
|
||||
// listenNamedPipe is unsupported off Windows; named pipes are a Windows-only transport.
|
||||
func listenNamedPipe(string) (net.Listener, error) {
|
||||
return nil, fmt.Errorf("named pipe daemon socket is only supported on Windows, not %s", runtime.GOOS)
|
||||
}
|
||||
|
||||
// dialNamedPipe is unsupported off Windows.
|
||||
func dialNamedPipe(context.Context, string) (net.Conn, error) {
|
||||
return nil, fmt.Errorf("named pipe daemon socket is only supported on Windows, not %s", runtime.GOOS)
|
||||
}
|
||||
@@ -1,30 +0,0 @@
|
||||
//go:build windows
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
|
||||
"github.com/Microsoft/go-winio"
|
||||
"golang.org/x/sys/windows"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/ipcauth"
|
||||
)
|
||||
|
||||
// listenNamedPipe creates the daemon control named pipe with a permissive,
|
||||
// local-only SDDL. Any local caller may connect, like Unix socket with 0666.
|
||||
func listenNamedPipe(path string) (net.Listener, error) {
|
||||
return winio.ListenPipe(path, &winio.PipeConfig{
|
||||
SecurityDescriptor: ipcauth.DefaultPipeSDDL(),
|
||||
})
|
||||
}
|
||||
|
||||
// dialNamedPipe connects to the daemon ipc named pipe at SECURITY_IDENTIFICATION.
|
||||
func dialNamedPipe(ctx context.Context, path string) (net.Conn, error) {
|
||||
access := uint32(windows.GENERIC_READ | windows.GENERIC_WRITE)
|
||||
// winio's plain DialPipe connects at SECURITY_ANONYMOUS, under which the
|
||||
// daemon cannot read the caller's token. Identification lets the daemon
|
||||
// read its SID/groups without granting it the ability to act as the caller.
|
||||
return winio.DialPipeAccessImpLevel(ctx, path, access, winio.PipeImpLevelIdentification)
|
||||
}
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"runtime"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
@@ -15,30 +14,6 @@ import (
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
const (
|
||||
windowsPipeDaemonAddr = "npipe://netbird"
|
||||
|
||||
// legacyWindowsDaemonAddr is the loopback-TCP address the Windows daemon used
|
||||
// before named-pipe support.
|
||||
legacyWindowsDaemonAddr = "tcp://127.0.0.1:41731"
|
||||
)
|
||||
|
||||
// migrateLegacyDaemonAddr upgrades the pre-named-pipe Windows daemon address to
|
||||
// the pipe. Existing installs persist daemon addr, so on upgrade the daemon
|
||||
// would otherwise keep listening on TCP and silently run without IPC
|
||||
// authorization. Only the exact legacy default is rewritten, while a
|
||||
// deliberately-chosen custom TCP address is left alone.
|
||||
func migrateLegacyDaemonAddr(addr string) (string, bool) {
|
||||
return migrateLegacyDaemonAddrForOS(runtime.GOOS, addr)
|
||||
}
|
||||
|
||||
func migrateLegacyDaemonAddrForOS(goos, addr string) (string, bool) {
|
||||
if goos == "windows" && addr == legacyWindowsDaemonAddr {
|
||||
return windowsPipeDaemonAddr, true
|
||||
}
|
||||
return addr, false
|
||||
}
|
||||
|
||||
type socketListener struct {
|
||||
net.Listener
|
||||
network string
|
||||
@@ -51,15 +26,6 @@ func listenOnAddress(addr string) (*socketListener, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if network == "npipe" {
|
||||
path := pipePath(address)
|
||||
listener, err := listenNamedPipe(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &socketListener{Listener: listener, network: network, address: path}, nil
|
||||
}
|
||||
|
||||
if network == "unix" {
|
||||
removeStaleUnixSocket(address)
|
||||
}
|
||||
@@ -75,26 +41,17 @@ func listenOnAddress(addr string) (*socketListener, error) {
|
||||
func parseListenAddress(addr string) (string, string, error) {
|
||||
network, address, ok := strings.Cut(addr, "://")
|
||||
if !ok || network == "" || address == "" {
|
||||
return "", "", fmt.Errorf("address must be in [unix|tcp|npipe]://[path|host:port|name] format: %q", addr)
|
||||
return "", "", fmt.Errorf("address must be in [unix|tcp]://[path|host:port] format: %q", addr)
|
||||
}
|
||||
|
||||
switch network {
|
||||
case "unix", "tcp", "npipe":
|
||||
case "unix", "tcp":
|
||||
return network, address, nil
|
||||
default:
|
||||
return "", "", fmt.Errorf("unsupported daemon address protocol: %v", network)
|
||||
}
|
||||
}
|
||||
|
||||
// pipePath maps a daemon-addr npipe name ("npipe://netbird") to a Windows
|
||||
// named-pipe path (\\.\pipe\netbird).
|
||||
func pipePath(name string) string {
|
||||
if strings.HasPrefix(name, `\\`) {
|
||||
return name
|
||||
}
|
||||
return `\\.\pipe\` + name
|
||||
}
|
||||
|
||||
func removeStaleUnixSocket(path string) {
|
||||
stat, err := os.Lstat(path)
|
||||
if err != nil {
|
||||
|
||||
@@ -1,63 +0,0 @@
|
||||
//go:build !ios && !android
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestMigrateLegacyDaemonAddrForOS(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
goos string
|
||||
addr string
|
||||
want string
|
||||
migrate bool
|
||||
}{
|
||||
{
|
||||
name: "windows legacy tcp migrates to pipe",
|
||||
goos: "windows",
|
||||
addr: legacyWindowsDaemonAddr,
|
||||
want: windowsPipeDaemonAddr,
|
||||
migrate: true,
|
||||
},
|
||||
{
|
||||
name: "windows pipe already migrated stays",
|
||||
goos: "windows",
|
||||
addr: windowsPipeDaemonAddr,
|
||||
want: windowsPipeDaemonAddr,
|
||||
migrate: false,
|
||||
},
|
||||
{
|
||||
name: "windows custom tcp left alone",
|
||||
goos: "windows",
|
||||
addr: "tcp://127.0.0.1:9999",
|
||||
want: "tcp://127.0.0.1:9999",
|
||||
migrate: false,
|
||||
},
|
||||
{
|
||||
name: "linux legacy-looking tcp not migrated",
|
||||
goos: "linux",
|
||||
addr: legacyWindowsDaemonAddr,
|
||||
want: legacyWindowsDaemonAddr,
|
||||
migrate: false,
|
||||
},
|
||||
{
|
||||
name: "linux unix socket untouched",
|
||||
goos: "linux",
|
||||
addr: "unix:///var/run/netbird.sock",
|
||||
want: "unix:///var/run/netbird.sock",
|
||||
migrate: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, ok := migrateLegacyDaemonAddrForOS(tc.goos, tc.addr)
|
||||
assert.Equal(t, tc.want, got)
|
||||
assert.Equal(t, tc.migrate, ok)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,73 +0,0 @@
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal"
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
)
|
||||
|
||||
var (
|
||||
setSSHConfigProfile string
|
||||
setSSHConfigUsername string
|
||||
setSSHConfigEnableRoot bool
|
||||
setSSHConfigDisableAuth bool
|
||||
)
|
||||
|
||||
// setSSHConfigCmd applies the privileged SSH settings (enable root login /
|
||||
// disable auth) to a profile over the daemon IPC. Hidden because users
|
||||
// interact via `netbird up` / the UI, not directly.
|
||||
var setSSHConfigCmd = &cobra.Command{
|
||||
Use: SetSSHConfigCmdName,
|
||||
Short: "Apply privileged SSH settings to a profile (internal elevation target)",
|
||||
Hidden: true,
|
||||
RunE: func(cmd *cobra.Command, _ []string) error {
|
||||
ctx := internal.CtxInitState(cmd.Context())
|
||||
|
||||
conn, err := DialClientGRPCServer(ctx, daemonAddr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("connect to daemon: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if cerr := conn.Close(); cerr != nil {
|
||||
log.Warnf("failed closing daemon gRPC client connection: %v", cerr)
|
||||
}
|
||||
}()
|
||||
|
||||
var enableRoot, disableAuth *bool
|
||||
if cmd.Flags().Changed(enableSSHRootFlag) {
|
||||
enableRoot = &setSSHConfigEnableRoot
|
||||
}
|
||||
if cmd.Flags().Changed(disableSSHAuthFlag) {
|
||||
disableAuth = &setSSHConfigDisableAuth
|
||||
}
|
||||
|
||||
req := buildSetSSHConfigRequest(setSSHConfigProfile, setSSHConfigUsername, enableRoot, disableAuth)
|
||||
if _, err := proto.NewDaemonServiceClient(conn).SetConfig(ctx, req); err != nil {
|
||||
return fmt.Errorf("apply SSH config: %w", err)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
// buildSetSSHConfigRequest builds a minimal SetConfigRequest that touches only
|
||||
// the SSH fields that were provided (nil pointers leave the daemon's stored
|
||||
// value untouched, matching setupSetConfigReq's pointer semantics).
|
||||
func buildSetSSHConfigRequest(profileName, username string, enableSSHRoot, disableSSHAuth *bool) *proto.SetConfigRequest {
|
||||
return &proto.SetConfigRequest{
|
||||
ProfileName: profileName,
|
||||
Username: username,
|
||||
EnableSSHRoot: enableSSHRoot,
|
||||
DisableSSHAuth: disableSSHAuth,
|
||||
}
|
||||
}
|
||||
|
||||
func init() {
|
||||
setSSHConfigCmd.Flags().StringVar(&setSSHConfigProfile, "profile", "", "profile name to apply the SSH settings to")
|
||||
setSSHConfigCmd.Flags().StringVar(&setSSHConfigUsername, "username", "", "owning username of the profile")
|
||||
setSSHConfigCmd.Flags().BoolVar(&setSSHConfigEnableRoot, enableSSHRootFlag, false, "enable root login for the SSH server")
|
||||
setSSHConfigCmd.Flags().BoolVar(&setSSHConfigDisableAuth, disableSSHAuthFlag, false, "disable SSH authentication")
|
||||
}
|
||||
@@ -5,8 +5,6 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os/user"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -174,9 +172,10 @@ func getStatus(ctx context.Context, fullPeerStatus bool, shouldRunProbes bool) (
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// getActiveProfileName asks the daemon for the active profile's display name,
|
||||
// annotated with the owning user when the active profile belongs to someone else.
|
||||
// Returns an empty string on any error so status output degrades gracefully.
|
||||
// getActiveProfileName asks the daemon for the active profile's display
|
||||
// name. The daemon runs as root and can read the per-user profile files to
|
||||
// resolve the ID to its human-readable name. Returns an empty string on any
|
||||
// error so status output degrades gracefully.
|
||||
func getActiveProfileName(ctx context.Context) string {
|
||||
conn, err := DialClientGRPCServer(ctx, daemonAddr)
|
||||
if err != nil {
|
||||
@@ -189,22 +188,7 @@ func getActiveProfileName(ctx context.Context) string {
|
||||
return ""
|
||||
}
|
||||
|
||||
name := resp.GetProfileName()
|
||||
if owner := resp.GetUsername(); owner != "" {
|
||||
if curr, uerr := user.Current(); uerr != nil || !usernamesMatch(owner, curr.Username) {
|
||||
name = fmt.Sprintf("%s (user %s)", name, owner)
|
||||
}
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
// usernamesMatch compares usernames case-insensitively on Windows (domain
|
||||
// accounts) and exactly elsewhere.
|
||||
func usernamesMatch(a, b string) bool {
|
||||
if runtime.GOOS == "windows" {
|
||||
return strings.EqualFold(a, b)
|
||||
}
|
||||
return a == b
|
||||
return resp.GetProfileName()
|
||||
}
|
||||
|
||||
func parseFilters() error {
|
||||
|
||||
@@ -21,9 +21,7 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal"
|
||||
"github.com/netbirdio/netbird/client/internal/peer"
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
nbnet "github.com/netbirdio/netbird/client/net"
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
"github.com/netbirdio/netbird/client/server"
|
||||
"github.com/netbirdio/netbird/client/system"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
"github.com/netbirdio/netbird/util"
|
||||
@@ -56,7 +54,6 @@ var (
|
||||
showQR bool
|
||||
profileName string
|
||||
configPath string
|
||||
claimOwner bool
|
||||
|
||||
upCmd = &cobra.Command{
|
||||
Use: "up",
|
||||
@@ -68,7 +65,6 @@ var (
|
||||
|
||||
func init() {
|
||||
upCmd.PersistentFlags().BoolVarP(&foregroundMode, "foreground-mode", "F", false, "start service in foreground")
|
||||
upCmd.PersistentFlags().BoolVar(&claimOwner, "owner", false, "claim ownership of this profile for the current user, restricting daemon control of it to you and root/administrator")
|
||||
upCmd.PersistentFlags().StringVar(&interfaceName, interfaceNameFlag, iface.WgInterfaceDefault, "WireGuard interface name")
|
||||
upCmd.PersistentFlags().Uint16Var(&wireguardPort, wireguardPortFlag, iface.DefaultWgPort, "WireGuard interface listening port")
|
||||
upCmd.PersistentFlags().Uint16Var(&mtu, mtuFlag, iface.DefaultMTU, "Set MTU (Maximum Transmission Unit) for the WireGuard interface")
|
||||
@@ -233,24 +229,6 @@ func runInForegroundMode(ctx context.Context, cmd *cobra.Command, activeProf *pr
|
||||
|
||||
_, _ = profilemanager.UpdateOldManagementURL(ctx, config, configFilePath)
|
||||
|
||||
// Restore residual state left by a previous run that did not shut down
|
||||
// cleanly, mirroring what the daemon does before connecting: it recovers
|
||||
// DNS config (a stale resolv.conf takeover can make the management
|
||||
// hostname unresolvable), firewall rules, ssh config and legacy routing.
|
||||
// Route cleanup itself happens at engine start; nbnet.Init() below lets
|
||||
// the management dial bypass a leftover fwmark rule until then.
|
||||
// Foreground mode is particularly exposed in containers: a crashed
|
||||
// container restarts inside the same (pod) network namespace, so stale
|
||||
// state survives while the process does not.
|
||||
if err := server.RestoreResidualState(ctx, profilemanager.NewServiceManager(configPath).GetStatePath()); err != nil {
|
||||
log.Warnf("failed to restore residual state: %v", err)
|
||||
}
|
||||
|
||||
// Enable advanced routing (as the daemon does on startup) so the
|
||||
// management dial bypasses a leftover fwmark rule instead of being
|
||||
// shunted into a stale routing table.
|
||||
nbnet.Init()
|
||||
|
||||
err = foregroundLogin(ctx, cmd, config, providedSetupKey, activeProf.ID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("foreground login failed: %v", err)
|
||||
@@ -321,24 +299,6 @@ func runInDaemonMode(ctx context.Context, cmd *cobra.Command, pm *profilemanager
|
||||
return fmt.Errorf("get current user: %v", err)
|
||||
}
|
||||
|
||||
// Enabling SSH root login / disabling SSH auth is a privileged change the
|
||||
// daemon rejects from a non-privileged caller. Offer to elevate (polkit/UAC/
|
||||
// osascript) so an elevated helper applies just those settings.
|
||||
if wantsDangerousSSH(cmd) && !isProcessPrivileged() {
|
||||
cmd.Println("Enabling SSH root login / disabling SSH authentication requires administrator privileges, requesting elevation...")
|
||||
var enableRoot, disableAuth *bool
|
||||
if cmd.Flag(enableSSHRootFlag).Changed {
|
||||
enableRoot = &enableSSHRoot
|
||||
}
|
||||
if cmd.Flag(disableSSHAuthFlag).Changed {
|
||||
disableAuth = &disableSSHAuth
|
||||
}
|
||||
if err := ElevateSSHConfig(activeProf.ID.String(), username.Username, enableRoot, disableAuth); err != nil {
|
||||
return fmt.Errorf("apply privileged SSH settings: %w", err)
|
||||
}
|
||||
sshConfigElevated = true
|
||||
}
|
||||
|
||||
// set the new config
|
||||
req := setupSetConfigReq(customDNSAddressConverted, cmd, activeProf.ID.String(), username.Username)
|
||||
if _, err := client.SetConfig(ctx, req); err != nil {
|
||||
@@ -411,7 +371,6 @@ func doDaemonUp(ctx context.Context, cmd *cobra.Command, client proto.DaemonServ
|
||||
if _, err := client.Up(ctx, &proto.UpRequest{
|
||||
ProfileName: &profileID,
|
||||
Username: &username,
|
||||
ClaimOwner: claimOwner,
|
||||
}); err != nil {
|
||||
return fmt.Errorf("call service up method: %v", err)
|
||||
}
|
||||
@@ -442,7 +401,7 @@ func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, pro
|
||||
if cmd.Flag(serverSSHAllowedFlag).Changed {
|
||||
req.ServerSSHAllowed = &serverSSHAllowed
|
||||
}
|
||||
if cmd.Flag(enableSSHRootFlag).Changed && !sshConfigElevated {
|
||||
if cmd.Flag(enableSSHRootFlag).Changed {
|
||||
req.EnableSSHRoot = &enableSSHRoot
|
||||
}
|
||||
if cmd.Flag(enableSSHSFTPFlag).Changed {
|
||||
@@ -454,7 +413,7 @@ func setupSetConfigReq(customDNSAddressConverted []byte, cmd *cobra.Command, pro
|
||||
if cmd.Flag(enableSSHRemotePortForwardFlag).Changed {
|
||||
req.EnableSSHRemotePortForwarding = &enableSSHRemotePortForward
|
||||
}
|
||||
if cmd.Flag(disableSSHAuthFlag).Changed && !sshConfigElevated {
|
||||
if cmd.Flag(disableSSHAuthFlag).Changed {
|
||||
req.DisableSSHAuth = &disableSSHAuth
|
||||
}
|
||||
if cmd.Flag(sshJWTCacheTTLFlag).Changed {
|
||||
@@ -670,7 +629,7 @@ func setupLoginRequest(providedSetupKey string, customDNSAddressConverted []byte
|
||||
loginRequest.ServerSSHAllowed = &serverSSHAllowed
|
||||
}
|
||||
|
||||
if cmd.Flag(enableSSHRootFlag).Changed && !sshConfigElevated {
|
||||
if cmd.Flag(enableSSHRootFlag).Changed {
|
||||
loginRequest.EnableSSHRoot = &enableSSHRoot
|
||||
}
|
||||
|
||||
@@ -686,7 +645,7 @@ func setupLoginRequest(providedSetupKey string, customDNSAddressConverted []byte
|
||||
loginRequest.EnableSSHRemotePortForwarding = &enableSSHRemotePortForward
|
||||
}
|
||||
|
||||
if cmd.Flag(disableSSHAuthFlag).Changed && !sshConfigElevated {
|
||||
if cmd.Flag(disableSSHAuthFlag).Changed {
|
||||
loginRequest.DisableSSHAuth = &disableSSHAuth
|
||||
}
|
||||
|
||||
|
||||
@@ -29,7 +29,7 @@ func TestUpDaemon(t *testing.T) {
|
||||
}
|
||||
|
||||
sm := profilemanager.ServiceManager{}
|
||||
created, err := sm.AddProfile("test1", currUser.Username, nil)
|
||||
created, err := sm.AddProfile("test1", currUser.Username)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to add profile: %v", err)
|
||||
return
|
||||
|
||||
@@ -121,7 +121,6 @@ type Manager struct {
|
||||
udpTracker *conntrack.UDPTracker
|
||||
icmpTracker *conntrack.ICMPTracker
|
||||
tcpTracker *conntrack.TCPTracker
|
||||
fragments *fragmentTracker
|
||||
forwarder atomic.Pointer[forwarder.Forwarder]
|
||||
pendingCapture atomic.Pointer[forwarder.PacketCapture]
|
||||
logger *nblog.Logger
|
||||
@@ -184,41 +183,6 @@ func (d *decoder) decodePacket(data []byte) error {
|
||||
}
|
||||
}
|
||||
|
||||
// decodeTransport decodes the transport header of a first fragment (which
|
||||
// gopacket leaves undecoded) into the decoder and appends its layer type to
|
||||
// decoded, so the ACL pipeline can evaluate it like a normal packet. It returns
|
||||
// false if the protocol is unsupported or the header is truncated.
|
||||
func (d *decoder) decodeTransport(proto layers.IPProtocol, payload []byte) bool {
|
||||
var l4 gopacket.DecodingLayer
|
||||
var layerType gopacket.LayerType
|
||||
var minLen int
|
||||
switch proto {
|
||||
case layers.IPProtocolTCP:
|
||||
l4, layerType, minLen = &d.tcp, layers.LayerTypeTCP, 20
|
||||
case layers.IPProtocolUDP:
|
||||
l4, layerType, minLen = &d.udp, layers.LayerTypeUDP, 8
|
||||
case layers.IPProtocolICMPv4:
|
||||
l4, layerType, minLen = &d.icmp4, layers.LayerTypeICMPv4, 8
|
||||
case layers.IPProtocolICMPv6:
|
||||
l4, layerType, minLen = &d.icmp6, layers.LayerTypeICMPv6, 8
|
||||
default:
|
||||
return false
|
||||
}
|
||||
|
||||
// Reject a fragment too small to hold the full transport header before
|
||||
// decoding: it can't be ACL-evaluated (tiny-fragment attack), and skipping
|
||||
// the decode avoids gopacket allocating an error on the drop path.
|
||||
if len(payload) < minLen {
|
||||
return false
|
||||
}
|
||||
|
||||
if err := l4.DecodeFromBytes(payload, gopacket.NilDecodeFeedback); err != nil {
|
||||
return false
|
||||
}
|
||||
d.decoded = append(d.decoded, layerType)
|
||||
return true
|
||||
}
|
||||
|
||||
// Create userspace firewall manager constructor
|
||||
func Create(iface common.IFaceMapper, disableServerRoutes bool, flowLogger nftypes.FlowLogger, mtu uint16) (*Manager, error) {
|
||||
return create(iface, nil, disableServerRoutes, flowLogger, mtu)
|
||||
@@ -322,8 +286,6 @@ func create(iface common.IFaceMapper, nativeFirewall firewall.Manager, disableSe
|
||||
if err := m.localipmanager.UpdateLocalIPs(iface); err != nil {
|
||||
return nil, fmt.Errorf("update local IPs: %w", err)
|
||||
}
|
||||
m.fragments = newFragmentTracker(m.logger)
|
||||
|
||||
if disableConntrack {
|
||||
log.Info("conntrack is disabled")
|
||||
} else {
|
||||
@@ -337,7 +299,6 @@ func create(iface common.IFaceMapper, nativeFirewall firewall.Manager, disableSe
|
||||
}
|
||||
}
|
||||
if err := iface.SetFilter(m); err != nil {
|
||||
m.fragments.Close()
|
||||
return nil, fmt.Errorf("set filter: %w", err)
|
||||
}
|
||||
return m, nil
|
||||
@@ -733,10 +694,6 @@ func (m *Manager) resetState() {
|
||||
m.tcpTracker.Close()
|
||||
}
|
||||
|
||||
if m.fragments != nil {
|
||||
m.fragments.Close()
|
||||
}
|
||||
|
||||
if fwder := m.forwarder.Load(); fwder != nil {
|
||||
fwder.SetCapture(nil)
|
||||
fwder.Stop()
|
||||
@@ -1089,20 +1046,19 @@ func (m *Manager) filterInbound(packetData []byte, size int) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// gopacket does not decode the transport header of any IP fragment, so
|
||||
// fragments take a dedicated path: the first fragment's header is decoded
|
||||
// and ACL-evaluated here, and the remaining fragments inherit its verdict.
|
||||
// TODO: pass fragments of routed packets to forwarder
|
||||
if fragment {
|
||||
return m.filterInboundFragment(d, srcIP, dstIP, size)
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
if d.decoded[0] == layers.LayerTypeIPv4 {
|
||||
m.logger.Trace4("packet is a fragment: src=%v dst=%v id=%v flags=%v",
|
||||
srcIP, dstIP, d.ip4.Id, d.ip4.Flags)
|
||||
} else {
|
||||
m.logger.Trace2("packet is an IPv6 fragment: src=%v dst=%v", srcIP, dstIP)
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
return m.filterInboundDecoded(d, srcIP, dstIP, packetData, size)
|
||||
}
|
||||
|
||||
// filterInboundDecoded runs the ACL, DNAT and conntrack pipeline on a fully
|
||||
// decoded (non-fragment) inbound packet. It returns true if the packet should
|
||||
// be dropped.
|
||||
func (m *Manager) filterInboundDecoded(d *decoder, srcIP, dstIP netip.Addr, packetData []byte, size int) bool {
|
||||
// TODO: optimize port DNAT by caching matched rules in conntrack
|
||||
if translated := m.translateInboundPortDNAT(packetData, d, srcIP, dstIP); translated {
|
||||
// Re-decode after port DNAT translation to update port information
|
||||
@@ -1133,226 +1089,33 @@ func (m *Manager) filterInboundDecoded(d *decoder, srcIP, dstIP netip.Addr, pack
|
||||
return m.handleRoutedTraffic(d, srcIP, dstIP, packetData, size)
|
||||
}
|
||||
|
||||
// fragmentMeta holds the reassembly identity and layout of an IP fragment,
|
||||
// extracted uniformly for IPv4 and IPv6.
|
||||
type fragmentMeta struct {
|
||||
key fragmentKey
|
||||
// offset is the fragment offset in 8-byte units (zero for the first
|
||||
// fragment).
|
||||
offset uint16
|
||||
// moreFragments is the More Fragments bit. A first fragment with it unset is
|
||||
// an IPv6 atomic fragment (a complete datagram, RFC 6946): it has no trailing
|
||||
// fragments to inherit a verdict, so it must not be recorded.
|
||||
moreFragments bool
|
||||
proto layers.IPProtocol
|
||||
// l4payload is the fragmentable payload of this fragment. For the first
|
||||
// fragment it starts with the transport header.
|
||||
l4payload []byte
|
||||
// headerEndOctets is the first fragment's payload length in 8-byte units:
|
||||
// the smallest offset a trailing fragment may start at without overlapping
|
||||
// the inspected transport header.
|
||||
headerEndOctets uint16
|
||||
}
|
||||
|
||||
// fragmentMetadata extracts the fragment identity and layout from a decoded IP
|
||||
// fragment. It returns false for fragments it can't interpret (e.g. an IPv6
|
||||
// fragment header shorter than 8 bytes), which are then dropped.
|
||||
func fragmentMetadata(d *decoder, srcIP, dstIP netip.Addr) (fragmentMeta, bool) {
|
||||
switch d.decoded[0] {
|
||||
case layers.LayerTypeIPv4:
|
||||
payload := d.ip4.Payload
|
||||
return fragmentMeta{
|
||||
key: fragmentKey{srcIP: srcIP, dstIP: dstIP, id: uint32(d.ip4.Id), proto: uint8(d.ip4.Protocol)},
|
||||
offset: d.ip4.FragOffset,
|
||||
moreFragments: d.ip4.Flags&layers.IPv4MoreFragments != 0,
|
||||
proto: d.ip4.Protocol,
|
||||
l4payload: payload,
|
||||
headerEndOctets: octets(len(payload)),
|
||||
}, true
|
||||
|
||||
case layers.LayerTypeIPv6:
|
||||
// IPv6 fragment extension header: 8 bytes, followed by the fragmentable
|
||||
// payload. Layout: next header (1), reserved (1), offset+flags (2), id (4).
|
||||
payload := d.ip6.Payload
|
||||
if len(payload) < 8 {
|
||||
return fragmentMeta{}, false
|
||||
}
|
||||
nextHeader := layers.IPProtocol(payload[0])
|
||||
offsetFlags := binary.BigEndian.Uint16(payload[2:4])
|
||||
id := binary.BigEndian.Uint32(payload[4:8])
|
||||
l4 := payload[8:]
|
||||
return fragmentMeta{
|
||||
key: fragmentKey{srcIP: srcIP, dstIP: dstIP, id: id, proto: uint8(nextHeader)},
|
||||
offset: offsetFlags >> 3,
|
||||
moreFragments: offsetFlags&1 != 0,
|
||||
proto: nextHeader,
|
||||
l4payload: l4,
|
||||
headerEndOctets: octets(len(l4)),
|
||||
}, true
|
||||
|
||||
default:
|
||||
return fragmentMeta{}, false
|
||||
}
|
||||
}
|
||||
|
||||
// octets rounds a byte length up to whole 8-byte units, the granularity of the
|
||||
// IP fragment offset field.
|
||||
func octets(nbytes int) uint16 {
|
||||
return uint16((nbytes + 7) / 8)
|
||||
}
|
||||
|
||||
// filterInboundFragment decides the fate of an IP fragment. gopacket stops
|
||||
// decoding at the network layer for every fragment, so the first fragment's
|
||||
// transport header is decoded and ACL-evaluated here and its verdict recorded;
|
||||
// the remaining (headerless) fragments inherit that verdict. Anything that
|
||||
// cannot be tied to an allowed, non-overlapping first fragment is dropped.
|
||||
func (m *Manager) filterInboundFragment(d *decoder, srcIP, dstIP netip.Addr, size int) bool {
|
||||
meta, ok := fragmentMetadata(d, srcIP, dstIP)
|
||||
if !ok {
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
m.logger.Trace2("dropping unsupported fragment: src=%v dst=%v", srcIP, dstIP)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
if meta.offset != 0 {
|
||||
return m.filterTrailingFragment(meta, srcIP, dstIP)
|
||||
}
|
||||
|
||||
// A new first fragment supersedes any recorded verdict for this datagram, so
|
||||
// a re-sent or overlapping offset-zero fragment can't inherit the old one.
|
||||
m.fragments.poison(meta.key)
|
||||
|
||||
// First fragment: decode its transport header so the ACL can evaluate it. A
|
||||
// decode failure means the fragment is too small to hold the full transport
|
||||
// header (RFC 1858 §3 tiny-fragment attack); it can't be evaluated, so drop it.
|
||||
if !d.decodeTransport(meta.proto, meta.l4payload) {
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
m.logger.Trace3("dropping first fragment without full L4 header: src=%v dst=%v id=%v",
|
||||
srcIP, dstIP, meta.key.id)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
return m.filterFirstFragment(d, meta, srcIP, dstIP, size)
|
||||
}
|
||||
|
||||
// filterTrailingFragment applies a recorded first-fragment verdict to a
|
||||
// non-first fragment.
|
||||
func (m *Manager) filterTrailingFragment(meta fragmentMeta, srcIP, dstIP netip.Addr) bool {
|
||||
switch m.fragments.verdict(meta.key, meta.offset) {
|
||||
case fragmentAllow:
|
||||
return false
|
||||
case fragmentOverlap:
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
m.logger.Trace3("dropping overlapping fragment rewriting inspected header: src=%v dst=%v id=%v",
|
||||
srcIP, dstIP, meta.key.id)
|
||||
}
|
||||
return true
|
||||
default:
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
m.logger.Trace3("dropping fragment with no allowed first fragment: src=%v dst=%v id=%v",
|
||||
srcIP, dstIP, meta.key.id)
|
||||
}
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// filterFirstFragment runs the verdict part of the inbound pipeline on a first
|
||||
// fragment with its transport header decoded. It mirrors filterInboundDecoded
|
||||
// but skips DNAT (port rewriting on fragments is unsupported) and forwarder
|
||||
// injection (fragments are left to the stack to reassemble, not forwarded).
|
||||
// Allowed fragments have their verdict recorded so the datagram's trailing
|
||||
// fragments inherit it.
|
||||
func (m *Manager) filterFirstFragment(d *decoder, meta fragmentMeta, srcIP, dstIP netip.Addr, size int) bool {
|
||||
if m.stateful && m.isValidTrackedConnection(d, srcIP, dstIP, size) {
|
||||
m.recordFirstFragment(meta)
|
||||
return false
|
||||
}
|
||||
|
||||
if m.localipmanager.IsLocalIP(dstIP) {
|
||||
ruleID, blocked := m.peerACLsBlock(srcIP, d, nil)
|
||||
if blocked {
|
||||
m.storeDropFlow("Dropping local first fragment (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||
d, srcIP, dstIP, ruleID, size)
|
||||
return true
|
||||
}
|
||||
m.trackInbound(d, srcIP, dstIP, ruleID, size)
|
||||
m.recordFirstFragment(meta)
|
||||
return false
|
||||
}
|
||||
|
||||
if !m.routingEnabled.Load() {
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
m.logger.Trace2("Dropping routed fragment (routing disabled): src=%s dst=%s", srcIP, dstIP)
|
||||
}
|
||||
return true
|
||||
}
|
||||
if m.nativeRouter.Load() {
|
||||
m.trackInbound(d, srcIP, dstIP, nil, size)
|
||||
m.recordFirstFragment(meta)
|
||||
return false
|
||||
}
|
||||
|
||||
// TODO: pass fragments of routed packets to the forwarder; until then
|
||||
// allowed routed fragments go to the native stack.
|
||||
srcPort, dstPort := getPortsFromPacket(d)
|
||||
ruleID, pass := m.routeACLsPass(srcIP, dstIP, d.decoded[1], srcPort, dstPort)
|
||||
if !pass {
|
||||
m.storeDropFlow("Dropping routed first fragment (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||
d, srcIP, dstIP, ruleID, size)
|
||||
return true
|
||||
}
|
||||
|
||||
m.recordFirstFragment(meta)
|
||||
return false
|
||||
}
|
||||
|
||||
// recordFirstFragment caches an allowed first fragment's verdict for its
|
||||
// trailing fragments to inherit. Atomic fragments (no More Fragments bit) are
|
||||
// complete datagrams with no trailing fragments, so they are not cached and
|
||||
// cannot exhaust the verdict table.
|
||||
func (m *Manager) recordFirstFragment(meta fragmentMeta) {
|
||||
if !meta.moreFragments {
|
||||
return
|
||||
}
|
||||
m.fragments.recordAllowed(meta.key, meta.headerEndOctets)
|
||||
}
|
||||
|
||||
// storeDropFlow logs and records a netflow drop event for an inbound packet
|
||||
// denied by the ACLs. msg is the trace format taking rule id, protocol, source
|
||||
// and destination.
|
||||
func (m *Manager) storeDropFlow(msg string, d *decoder, srcIP, dstIP netip.Addr, ruleID []byte, size int) {
|
||||
pnum := getProtocolFromPacket(d)
|
||||
srcPort, dstPort := getPortsFromPacket(d)
|
||||
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
m.logger.Trace6(msg, ruleID, pnum, srcIP, srcPort, dstIP, dstPort)
|
||||
}
|
||||
|
||||
m.flowLogger.StoreEvent(nftypes.EventFields{
|
||||
FlowID: uuid.New(),
|
||||
Type: nftypes.TypeDrop,
|
||||
RuleID: ruleID,
|
||||
Direction: nftypes.Ingress,
|
||||
Protocol: pnum,
|
||||
SourceIP: srcIP,
|
||||
DestIP: dstIP,
|
||||
SourcePort: srcPort,
|
||||
DestPort: dstPort,
|
||||
// TODO: icmp type/code
|
||||
RxPackets: 1,
|
||||
RxBytes: uint64(size),
|
||||
})
|
||||
}
|
||||
|
||||
// handleLocalTraffic handles local traffic.
|
||||
// If it returns true, the packet should be dropped.
|
||||
func (m *Manager) handleLocalTraffic(d *decoder, srcIP, dstIP netip.Addr, packetData []byte, size int) bool {
|
||||
ruleID, blocked := m.peerACLsBlock(srcIP, d, packetData)
|
||||
if blocked {
|
||||
m.storeDropFlow("Dropping local packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||
d, srcIP, dstIP, ruleID, size)
|
||||
pnum := getProtocolFromPacket(d)
|
||||
srcPort, dstPort := getPortsFromPacket(d)
|
||||
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
m.logger.Trace6("Dropping local packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||
ruleID, pnum, srcIP, srcPort, dstIP, dstPort)
|
||||
}
|
||||
|
||||
m.flowLogger.StoreEvent(nftypes.EventFields{
|
||||
FlowID: uuid.New(),
|
||||
Type: nftypes.TypeDrop,
|
||||
RuleID: ruleID,
|
||||
Direction: nftypes.Ingress,
|
||||
Protocol: pnum,
|
||||
SourceIP: srcIP,
|
||||
DestIP: dstIP,
|
||||
SourcePort: srcPort,
|
||||
DestPort: dstPort,
|
||||
// TODO: icmp type/code
|
||||
RxPackets: 1,
|
||||
RxBytes: uint64(size),
|
||||
})
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -1405,8 +1168,27 @@ func (m *Manager) handleRoutedTraffic(d *decoder, srcIP, dstIP netip.Addr, packe
|
||||
|
||||
ruleID, pass := m.routeACLsPass(srcIP, dstIP, protoLayer, srcPort, dstPort)
|
||||
if !pass {
|
||||
m.storeDropFlow("Dropping routed packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||
d, srcIP, dstIP, ruleID, size)
|
||||
proto := getProtocolFromPacket(d)
|
||||
|
||||
if m.logger.Enabled(nblog.LevelTrace) {
|
||||
m.logger.Trace6("Dropping routed packet (ACL denied): rule_id=%s proto=%v src=%s:%d dst=%s:%d",
|
||||
ruleID, proto, srcIP, srcPort, dstIP, dstPort)
|
||||
}
|
||||
|
||||
m.flowLogger.StoreEvent(nftypes.EventFields{
|
||||
FlowID: uuid.New(),
|
||||
Type: nftypes.TypeDrop,
|
||||
RuleID: ruleID,
|
||||
Direction: nftypes.Ingress,
|
||||
Protocol: proto,
|
||||
SourceIP: srcIP,
|
||||
DestIP: dstIP,
|
||||
SourcePort: srcPort,
|
||||
DestPort: dstPort,
|
||||
// TODO: icmp type/code
|
||||
RxPackets: 1,
|
||||
RxBytes: uint64(size),
|
||||
})
|
||||
return true
|
||||
}
|
||||
|
||||
|
||||
@@ -5,9 +5,7 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -33,11 +31,6 @@ const (
|
||||
defaultMaxInFlight = 1024
|
||||
iosReceiveWindow = 16384
|
||||
iosMaxInFlight = 256
|
||||
|
||||
// envForceTCPRACK overrides the platform default for gVisor's RACK loss
|
||||
// detection. Set to a truthy value to force RACK on, or a falsy value to
|
||||
// force it off, on any platform.
|
||||
envForceTCPRACK = "NB_FORCE_TCP_RACK"
|
||||
)
|
||||
|
||||
type Forwarder struct {
|
||||
@@ -159,8 +152,6 @@ func New(iface common.IFaceMapper, logger *nblog.Logger, flowLogger nftypes.Flow
|
||||
maxInFlight = iosMaxInFlight
|
||||
}
|
||||
|
||||
configureTCPRecovery(s)
|
||||
|
||||
tcpForwarder := tcp.NewForwarder(s, receiveWindow, maxInFlight, f.handleTCP)
|
||||
s.SetTransportProtocolHandler(tcp.ProtocolNumber, tcpForwarder.HandlePacket)
|
||||
|
||||
@@ -475,31 +466,3 @@ func probeRawICMP(network, addr string, logger *nblog.Logger) bool {
|
||||
logger.Debug1("forwarder: raw %s socket access available", network)
|
||||
return true
|
||||
}
|
||||
|
||||
// configureTCPRecovery disables gVisor's RACK loss detection on Windows, where
|
||||
// it interacts poorly with the host and collapses throughput on routed TCP
|
||||
// connections (gVisor issue #9778). Other platforms keep the default. The
|
||||
// EnvForceTCPRACK environment variable overrides the platform default.
|
||||
func configureTCPRecovery(s *stack.Stack) {
|
||||
disableRACK := runtime.GOOS == "windows"
|
||||
|
||||
if val := os.Getenv(envForceTCPRACK); val != "" {
|
||||
force, err := strconv.ParseBool(val)
|
||||
if err != nil {
|
||||
log.Warnf("parse %s: %v", envForceTCPRACK, err)
|
||||
} else {
|
||||
disableRACK = !force
|
||||
}
|
||||
}
|
||||
|
||||
if !disableRACK {
|
||||
return
|
||||
}
|
||||
|
||||
opt := tcpip.TCPRecovery(0)
|
||||
if err := s.SetTransportProtocolOption(tcp.ProtocolNumber, &opt); err != nil {
|
||||
log.Warnf("disable TCP RACK loss detection: %v", err)
|
||||
return
|
||||
}
|
||||
log.Info("forwarder: TCP RACK loss detection disabled")
|
||||
}
|
||||
|
||||
@@ -1,204 +0,0 @@
|
||||
package uspfilter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
nblog "github.com/netbirdio/netbird/client/firewall/uspfilter/log"
|
||||
)
|
||||
|
||||
const (
|
||||
// defaultFragmentTimeout bounds how long a first-fragment verdict is kept
|
||||
// while the remaining fragments arrive. It mirrors the Linux IP reassembly
|
||||
// timeout (net.ipv4.ipfrag_time).
|
||||
defaultFragmentTimeout = 30 * time.Second
|
||||
// fragmentCleanupInterval is how often expired verdicts are purged.
|
||||
fragmentCleanupInterval = 10 * time.Second
|
||||
// defaultMaxFragmentEntries caps the number of concurrently tracked
|
||||
// fragmented datagrams. The table stays bounded because each datagram is a
|
||||
// single small entry regardless of how many fragments it is split into, and
|
||||
// the 13-bit IPv4 fragment-offset field limits any datagram to 64 KiB.
|
||||
defaultMaxFragmentEntries = 16384
|
||||
|
||||
// EnvFragmentMaxEntries overrides defaultMaxFragmentEntries.
|
||||
EnvFragmentMaxEntries = "NB_FRAGMENT_MAX_ENTRIES"
|
||||
)
|
||||
|
||||
// fragmentVerdict is the decision for a trailing (headerless) fragment.
|
||||
type fragmentVerdict int
|
||||
|
||||
const (
|
||||
// fragmentDeny drops the fragment: no allowed first fragment is on record.
|
||||
fragmentDeny fragmentVerdict = iota
|
||||
// fragmentAllow passes the fragment: it belongs to an allowed datagram and
|
||||
// does not overlap the already-inspected transport header.
|
||||
fragmentAllow
|
||||
// fragmentOverlap drops the fragment and poisons its datagram: it overlaps
|
||||
// the transport header the ACL inspected (RFC 1858 §4, RFC 3128; RFC 5722
|
||||
// requires discarding the whole datagram on overlap for IPv6).
|
||||
fragmentOverlap
|
||||
)
|
||||
|
||||
// fragmentKey identifies a fragmented datagram. It matches the RFC 791 / RFC
|
||||
// 8200 reassembly key: source, destination, protocol and identification. The id
|
||||
// is 32-bit to hold both the IPv4 (16-bit) and IPv6 (32-bit) identification.
|
||||
type fragmentKey struct {
|
||||
srcIP netip.Addr
|
||||
dstIP netip.Addr
|
||||
id uint32
|
||||
proto uint8
|
||||
}
|
||||
|
||||
// fragmentEntry records the verdict of an allowed first fragment.
|
||||
type fragmentEntry struct {
|
||||
// headerEndOctets is the offset, in 8-byte units, at which the first
|
||||
// fragment's payload ended. A trailing fragment starting before this
|
||||
// overlaps bytes the ACL already inspected and is rejected.
|
||||
headerEndOctets uint16
|
||||
// recordedAt is when the first fragment was accepted. The verdict expires a
|
||||
// fixed timeout later and is not refreshed, mirroring the kernel reassembly
|
||||
// timer so a trailing-fragment flood can't keep a datagram alive.
|
||||
recordedAt time.Time
|
||||
}
|
||||
|
||||
// fragmentTracker records the ACL verdict of a datagram's first fragment so the
|
||||
// remaining fragments, which carry no L4 header, can inherit the decision
|
||||
// without reassembling the datagram. Only allowed first fragments are stored;
|
||||
// anything that cannot be tied to an allowed, non-overlapping first fragment is
|
||||
// dropped (fail closed).
|
||||
type fragmentTracker struct {
|
||||
logger *nblog.Logger
|
||||
mutex sync.Mutex
|
||||
entries map[fragmentKey]fragmentEntry
|
||||
timeout time.Duration
|
||||
// maxEntries caps the table; atCapacity dedups the capacity warning until
|
||||
// the table drains below the cap again.
|
||||
maxEntries int
|
||||
atCapacity bool
|
||||
cleanupTicker *time.Ticker
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
func newFragmentTracker(logger *nblog.Logger) *fragmentTracker {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t := &fragmentTracker{
|
||||
logger: logger,
|
||||
entries: make(map[fragmentKey]fragmentEntry),
|
||||
timeout: defaultFragmentTimeout,
|
||||
maxEntries: fragmentMaxEntries(logger),
|
||||
cleanupTicker: time.NewTicker(fragmentCleanupInterval),
|
||||
cancel: cancel,
|
||||
}
|
||||
go t.cleanupRoutine(ctx)
|
||||
return t
|
||||
}
|
||||
|
||||
func fragmentMaxEntries(logger *nblog.Logger) int {
|
||||
v := os.Getenv(EnvFragmentMaxEntries)
|
||||
if v == "" {
|
||||
return defaultMaxFragmentEntries
|
||||
}
|
||||
n, err := strconv.Atoi(v)
|
||||
if err != nil || n <= 0 {
|
||||
logger.Warn2("invalid %s=%q, using default", EnvFragmentMaxEntries, v)
|
||||
return defaultMaxFragmentEntries
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// recordAllowed stores the verdict of an allowed first fragment. headerEndOctets
|
||||
// is the first fragment's payload length in 8-byte units. When the table is full
|
||||
// the record is dropped, which fails closed: the datagram's trailing fragments
|
||||
// will be denied.
|
||||
func (t *fragmentTracker) recordAllowed(key fragmentKey, headerEndOctets uint16) {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
|
||||
if t.entries == nil {
|
||||
return
|
||||
}
|
||||
if _, ok := t.entries[key]; !ok && len(t.entries) >= t.maxEntries {
|
||||
if !t.atCapacity {
|
||||
t.atCapacity = true
|
||||
t.logger.Warn2("fragment verdict table at capacity (%d/%d): trailing fragments of new datagrams will be dropped",
|
||||
len(t.entries), t.maxEntries)
|
||||
}
|
||||
return
|
||||
}
|
||||
t.entries[key] = fragmentEntry{
|
||||
headerEndOctets: headerEndOctets,
|
||||
recordedAt: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
// poison drops any recorded verdict for a datagram, so its later fragments are
|
||||
// denied until a new allowed first fragment is recorded. Called on every
|
||||
// offset-zero fragment to defeat offset-zero overlap rewrites (RFC 3128).
|
||||
func (t *fragmentTracker) poison(key fragmentKey) {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
delete(t.entries, key)
|
||||
}
|
||||
|
||||
// verdict decides the fate of a trailing fragment at fragOffsetOctets (the IPv4
|
||||
// fragment offset, in 8-byte units). A fragment overlapping the inspected
|
||||
// header poisons the datagram: the entry is removed so all further fragments of
|
||||
// that datagram are denied too.
|
||||
func (t *fragmentTracker) verdict(key fragmentKey, fragOffsetOctets uint16) fragmentVerdict {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
|
||||
entry, ok := t.entries[key]
|
||||
if !ok {
|
||||
return fragmentDeny
|
||||
}
|
||||
if time.Since(entry.recordedAt) > t.timeout {
|
||||
delete(t.entries, key)
|
||||
return fragmentDeny
|
||||
}
|
||||
if fragOffsetOctets < entry.headerEndOctets {
|
||||
delete(t.entries, key)
|
||||
return fragmentOverlap
|
||||
}
|
||||
return fragmentAllow
|
||||
}
|
||||
|
||||
func (t *fragmentTracker) cleanupRoutine(ctx context.Context) {
|
||||
defer t.cleanupTicker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-t.cleanupTicker.C:
|
||||
t.cleanup()
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *fragmentTracker) cleanup() {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
|
||||
for key, entry := range t.entries {
|
||||
if time.Since(entry.recordedAt) > t.timeout {
|
||||
delete(t.entries, key)
|
||||
}
|
||||
}
|
||||
|
||||
if len(t.entries) < t.maxEntries {
|
||||
t.atCapacity = false
|
||||
}
|
||||
}
|
||||
|
||||
// Close stops the cleanup routine and releases resources.
|
||||
func (t *fragmentTracker) Close() {
|
||||
t.cancel()
|
||||
|
||||
t.mutex.Lock()
|
||||
t.entries = nil
|
||||
t.mutex.Unlock()
|
||||
}
|
||||
@@ -1,115 +0,0 @@
|
||||
package uspfilter
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// benchFilterInbound drives filterInbound over a fixed packet in a tight loop.
|
||||
// Packets are built once, outside the timed region, so the benchmark measures
|
||||
// only pipeline cost, which is what an attacker can amplify.
|
||||
func benchFilterInbound(b *testing.B, pkt []byte) {
|
||||
b.Helper()
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkt)))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
m := benchManager
|
||||
m.filterInbound(pkt, len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
// benchManager is a package-level manager reused across fragment benchmarks so
|
||||
// setup cost stays out of the timed region.
|
||||
var benchManager *Manager
|
||||
|
||||
func setupBenchManager(b *testing.B) *Manager {
|
||||
b.Helper()
|
||||
m := newFragmentTestManager(b)
|
||||
allowUDP(b, m, 8080)
|
||||
// Disable conntrack so the allowed-first-fragment path measures transport
|
||||
// decode + ACL every iteration instead of matching the connection tracked
|
||||
// on the first iteration.
|
||||
m.stateful = false
|
||||
benchManager = m
|
||||
return m
|
||||
}
|
||||
|
||||
// BenchmarkInbound_NormalPacket is the baseline: a full, non-fragmented UDP
|
||||
// packet that passes the ACL. Fragment paths should stay comparable to this.
|
||||
func BenchmarkInbound_NormalPacket(b *testing.B) {
|
||||
setupBenchManager(b)
|
||||
pkt := normalUDPPacket(b, 8080, 32)
|
||||
benchFilterInbound(b, pkt)
|
||||
}
|
||||
|
||||
// BenchmarkInbound_FirstFragmentAllowed measures the first-fragment path:
|
||||
// transport decode + ACL evaluation + verdict record.
|
||||
func BenchmarkInbound_FirstFragmentAllowed(b *testing.B) {
|
||||
setupBenchManager(b)
|
||||
pkt := firstFragmentUDP(b, 0x2000, 8080, 32)
|
||||
benchFilterInbound(b, pkt)
|
||||
}
|
||||
|
||||
// BenchmarkInbound_TrailingFragmentAllowed measures the common trailing-fragment
|
||||
// path: a single map lookup after the first fragment is on record.
|
||||
func BenchmarkInbound_TrailingFragmentAllowed(b *testing.B) {
|
||||
m := setupBenchManager(b)
|
||||
first := firstFragmentUDP(b, 0x3000, 8080, 32)
|
||||
m.filterInbound(first, len(first))
|
||||
pkt := trailingFragment(b, 0x3000, 5, false, 24)
|
||||
benchFilterInbound(b, pkt)
|
||||
}
|
||||
|
||||
// BenchmarkInbound_TrailingFragmentNoFirst is the primary DoS vector: an
|
||||
// attacker floods trailing fragments with no first fragment on record. Each is
|
||||
// a map miss and must be cheap.
|
||||
func BenchmarkInbound_TrailingFragmentNoFirst(b *testing.B) {
|
||||
setupBenchManager(b)
|
||||
pkt := trailingFragment(b, 0x4000, 185, false, 40)
|
||||
benchFilterInbound(b, pkt)
|
||||
}
|
||||
|
||||
// BenchmarkInbound_TinyFirstFragment measures the tiny-fragment drop path: a
|
||||
// first fragment too small to decode a transport header.
|
||||
func BenchmarkInbound_TinyFirstFragment(b *testing.B) {
|
||||
setupBenchManager(b)
|
||||
pkt := trailingFragment(b, 0x5000, 0, true, 4)
|
||||
benchFilterInbound(b, pkt)
|
||||
}
|
||||
|
||||
// BenchmarkInbound_TrailingFragmentDistinctIDs is the worst case for the
|
||||
// verdict table: an attacker varies the datagram id on every packet so no first
|
||||
// fragment ever matches. Verdict lookups always miss and nothing is recorded,
|
||||
// so the table cannot grow. Each iteration rewrites the id field in place.
|
||||
func BenchmarkInbound_TrailingFragmentDistinctIDs(b *testing.B) {
|
||||
setupBenchManager(b)
|
||||
pkt := trailingFragment(b, 0x6000, 185, false, 40)
|
||||
m := benchManager
|
||||
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkt)))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
// IPv4 identification field is at bytes 4:6.
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(i))
|
||||
m.filterInbound(pkt, len(pkt))
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkInbound_FirstFragmentDistinctIDs measures sustained first-fragment
|
||||
// pressure with distinct ids: transport decode + ACL + verdict insert until the
|
||||
// table caps, exercising the map growth and capacity guard.
|
||||
func BenchmarkInbound_FirstFragmentDistinctIDs(b *testing.B) {
|
||||
setupBenchManager(b)
|
||||
pkt := firstFragmentUDP(b, 0x7000, 8080, 32)
|
||||
m := benchManager
|
||||
|
||||
b.ReportAllocs()
|
||||
b.SetBytes(int64(len(pkt)))
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(i))
|
||||
m.filterInbound(pkt, len(pkt))
|
||||
}
|
||||
}
|
||||
@@ -1,554 +0,0 @@
|
||||
package uspfilter
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
fw "github.com/netbirdio/netbird/client/firewall/manager"
|
||||
nbiface "github.com/netbirdio/netbird/client/iface"
|
||||
"github.com/netbirdio/netbird/client/iface/device"
|
||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||
)
|
||||
|
||||
const (
|
||||
fragTestSrc = "100.10.0.1"
|
||||
fragTestDst = "100.10.0.100"
|
||||
fragTestSrcV6 = "fd00::1"
|
||||
fragTestDstV6 = "fd00::100"
|
||||
)
|
||||
|
||||
func newFragmentTestManager(tb testing.TB) *Manager {
|
||||
tb.Helper()
|
||||
|
||||
ifaceMock := &IFaceMock{
|
||||
SetFilterFunc: func(device.PacketFilter) error { return nil },
|
||||
AddressFunc: func() wgaddr.Address {
|
||||
return wgaddr.Address{
|
||||
IP: netip.MustParseAddr(fragTestDst),
|
||||
Network: netip.MustParsePrefix("100.10.0.0/16"),
|
||||
IPv6: netip.MustParseAddr(fragTestDstV6),
|
||||
IPv6Net: netip.MustParsePrefix("fd00::/64"),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
m, err := Create(ifaceMock, false, flowLogger, nbiface.DefaultMTU)
|
||||
require.NoError(tb, err)
|
||||
require.NoError(tb, m.UpdateLocalIPs())
|
||||
tb.Cleanup(func() { require.NoError(tb, m.Close(nil)) })
|
||||
return m
|
||||
}
|
||||
|
||||
// firstFragmentUDPTo builds the first fragment of a fragmented UDP datagram to
|
||||
// the given destination: it carries the full UDP header plus payloadLen bytes
|
||||
// of data, with the More Fragments flag set and offset zero.
|
||||
func firstFragmentUDPTo(tb testing.TB, dst string, id uint16, dstPort uint16, payloadLen int) []byte {
|
||||
tb.Helper()
|
||||
|
||||
ip := &layers.IPv4{
|
||||
Version: 4,
|
||||
TTL: 64,
|
||||
Id: id,
|
||||
Protocol: layers.IPProtocolUDP,
|
||||
SrcIP: net.ParseIP(fragTestSrc),
|
||||
DstIP: net.ParseIP(dst),
|
||||
Flags: layers.IPv4MoreFragments,
|
||||
}
|
||||
udp := &layers.UDP{SrcPort: 40000, DstPort: layers.UDPPort(dstPort)}
|
||||
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||
|
||||
buf := gopacket.NewSerializeBuffer()
|
||||
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, payloadLen))))
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
func firstFragmentUDP(tb testing.TB, id uint16, dstPort uint16, payloadLen int) []byte {
|
||||
tb.Helper()
|
||||
return firstFragmentUDPTo(tb, fragTestDst, id, dstPort, payloadLen)
|
||||
}
|
||||
|
||||
// firstFragmentTCP builds the first fragment of a fragmented TCP datagram: the
|
||||
// full 20-byte TCP header plus 12 bytes of data, with the More Fragments flag
|
||||
// set and offset zero.
|
||||
func firstFragmentTCP(tb testing.TB, id uint16, dstPort uint16) []byte {
|
||||
tb.Helper()
|
||||
|
||||
ip := &layers.IPv4{
|
||||
Version: 4,
|
||||
TTL: 64,
|
||||
Id: id,
|
||||
Protocol: layers.IPProtocolTCP,
|
||||
SrcIP: net.ParseIP(fragTestSrc),
|
||||
DstIP: net.ParseIP(fragTestDst),
|
||||
Flags: layers.IPv4MoreFragments,
|
||||
}
|
||||
tcp := &layers.TCP{SrcPort: 40000, DstPort: layers.TCPPort(dstPort), SYN: true, Window: 64240}
|
||||
require.NoError(tb, tcp.SetNetworkLayerForChecksum(ip))
|
||||
|
||||
buf := gopacket.NewSerializeBuffer()
|
||||
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, tcp, gopacket.Payload(make([]byte, 12))))
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
// trailingFragmentTo builds a non-first fragment to the given destination: an
|
||||
// IPv4 header at the given fragment offset (in 8-byte units) carrying raw
|
||||
// payload and no L4 header.
|
||||
func trailingFragmentTo(tb testing.TB, dst string, proto layers.IPProtocol, id uint16, fragOffsetOctets uint16, moreFragments bool, payloadLen int) []byte {
|
||||
tb.Helper()
|
||||
|
||||
ip := &layers.IPv4{
|
||||
Version: 4,
|
||||
TTL: 64,
|
||||
Id: id,
|
||||
Protocol: proto,
|
||||
SrcIP: net.ParseIP(fragTestSrc),
|
||||
DstIP: net.ParseIP(dst),
|
||||
FragOffset: fragOffsetOctets,
|
||||
}
|
||||
if moreFragments {
|
||||
ip.Flags = layers.IPv4MoreFragments
|
||||
}
|
||||
|
||||
buf := gopacket.NewSerializeBuffer()
|
||||
opts := gopacket.SerializeOptions{FixLengths: true}
|
||||
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, gopacket.Payload(make([]byte, payloadLen))))
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
func trailingFragment(tb testing.TB, id uint16, fragOffsetOctets uint16, moreFragments bool, payloadLen int) []byte {
|
||||
tb.Helper()
|
||||
return trailingFragmentTo(tb, fragTestDst, layers.IPProtocolUDP, id, fragOffsetOctets, moreFragments, payloadLen)
|
||||
}
|
||||
|
||||
// outboundUDPPacket builds a complete outbound UDP packet from the local
|
||||
// address, used to establish conntrack state for reply-direction tests.
|
||||
func outboundUDPPacket(tb testing.TB, srcPort, dstPort uint16) []byte {
|
||||
tb.Helper()
|
||||
|
||||
ip := &layers.IPv4{
|
||||
Version: 4,
|
||||
TTL: 64,
|
||||
Id: 1,
|
||||
Protocol: layers.IPProtocolUDP,
|
||||
SrcIP: net.ParseIP(fragTestDst),
|
||||
DstIP: net.ParseIP(fragTestSrc),
|
||||
}
|
||||
udp := &layers.UDP{SrcPort: layers.UDPPort(srcPort), DstPort: layers.UDPPort(dstPort)}
|
||||
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||
|
||||
buf := gopacket.NewSerializeBuffer()
|
||||
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, 16))))
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
// normalUDPPacket builds a complete, non-fragmented UDP packet for baseline
|
||||
// comparisons against the fragment paths.
|
||||
func normalUDPPacket(tb testing.TB, dstPort uint16, payloadLen int) []byte {
|
||||
tb.Helper()
|
||||
|
||||
ip := &layers.IPv4{
|
||||
Version: 4,
|
||||
TTL: 64,
|
||||
Id: 1,
|
||||
Protocol: layers.IPProtocolUDP,
|
||||
SrcIP: net.ParseIP(fragTestSrc),
|
||||
DstIP: net.ParseIP(fragTestDst),
|
||||
}
|
||||
udp := &layers.UDP{SrcPort: 40000, DstPort: layers.UDPPort(dstPort)}
|
||||
require.NoError(tb, udp.SetNetworkLayerForChecksum(ip))
|
||||
|
||||
buf := gopacket.NewSerializeBuffer()
|
||||
opts := gopacket.SerializeOptions{ComputeChecksums: true, FixLengths: true}
|
||||
require.NoError(tb, gopacket.SerializeLayers(buf, opts, ip, udp, gopacket.Payload(make([]byte, payloadLen))))
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
func allowUDP(tb testing.TB, m *Manager, dstPort uint16) {
|
||||
tb.Helper()
|
||||
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolUDP, nil,
|
||||
&fw.Port{Values: []uint16{dstPort}}, fw.ActionAccept, "")
|
||||
require.NoError(tb, err)
|
||||
}
|
||||
|
||||
// TestFragment_TrailingWithoutFirstDropped is the core bypass repro: a trailing
|
||||
// fragment with no allowed first fragment on record must be dropped. Before the
|
||||
// fix, filterInbound returned false (allow) for any fragment.
|
||||
func TestFragment_TrailingWithoutFirstDropped(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
|
||||
frag := trailingFragment(t, 0x1234, 185, false, 40)
|
||||
require.True(t, m.filterInbound(frag, len(frag)),
|
||||
"trailing fragment without an allowed first fragment must be dropped")
|
||||
}
|
||||
|
||||
// TestFragment_AllowedFirstPassesTrailing verifies that once a first fragment
|
||||
// passes the ACL, its trailing fragments inherit the allow verdict.
|
||||
func TestFragment_AllowedFirstPassesTrailing(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
allowUDP(t, m, 8080)
|
||||
|
||||
// First fragment: UDP header (8) + 32 payload = 40 octets -> headerEnd = 5.
|
||||
first := firstFragmentUDP(t, 0x2222, 8080, 32)
|
||||
require.False(t, m.filterInbound(first, len(first)),
|
||||
"allowed first fragment should pass and be recorded")
|
||||
|
||||
trailing := trailingFragment(t, 0x2222, 5, false, 24)
|
||||
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||
"trailing fragment of an allowed datagram should pass")
|
||||
}
|
||||
|
||||
// TestFragment_DeniedFirstDropsTrailing verifies that a first fragment blocked
|
||||
// by the ACL leaves no verdict, so its trailing fragments are dropped.
|
||||
func TestFragment_DeniedFirstDropsTrailing(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
// No accept rule: local traffic defaults to deny.
|
||||
|
||||
first := firstFragmentUDP(t, 0x3333, 9999, 32)
|
||||
require.True(t, m.filterInbound(first, len(first)),
|
||||
"first fragment to a blocked port should be dropped by the ACL")
|
||||
|
||||
trailing := trailingFragment(t, 0x3333, 5, false, 24)
|
||||
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||
"trailing fragment of a denied datagram must be dropped")
|
||||
}
|
||||
|
||||
// TestFragment_OverlappingHeaderDropped covers the RFC 1858 §4 / RFC 3128
|
||||
// overlapping-fragment rewrite: a trailing fragment starting inside the range
|
||||
// the ACL already inspected is dropped and poisons the datagram. TCP is used so
|
||||
// the overlap lands on real header bytes (the flags at byte 13).
|
||||
func TestFragment_OverlappingHeaderDropped(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||
require.NoError(t, err)
|
||||
|
||||
// First fragment: TCP header (20) + 12 data = 32 bytes -> headerEnd = 4 octets.
|
||||
first := firstFragmentTCP(t, 0x4444, 8080)
|
||||
require.False(t, m.filterInbound(first, len(first)))
|
||||
|
||||
// Overlapping fragment at offset 1 (byte 8) falls inside the inspected TCP
|
||||
// header, so it could rewrite the flags or port on reassembly.
|
||||
overlap := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x4444, 1, true, 32)
|
||||
require.True(t, m.filterInbound(overlap, len(overlap)),
|
||||
"fragment overlapping the inspected header must be dropped")
|
||||
|
||||
// The datagram is now poisoned: a later, non-overlapping fragment is also
|
||||
// dropped because the verdict was removed.
|
||||
later := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x4444, 4, false, 24)
|
||||
require.True(t, m.filterInbound(later, len(later)),
|
||||
"fragments after an overlap must be dropped (datagram poisoned)")
|
||||
}
|
||||
|
||||
// TestFragment_OffsetZeroOverlapPoisons covers the RFC 3128 offset-zero rewrite:
|
||||
// an allowed first fragment followed by a denied offset-zero fragment for the
|
||||
// same datagram must not leave the earlier allow verdict in place.
|
||||
func TestFragment_OffsetZeroOverlapPoisons(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
allowUDP(t, m, 8080)
|
||||
|
||||
allowed := firstFragmentUDP(t, 0x5A5A, 8080, 32)
|
||||
require.False(t, m.filterInbound(allowed, len(allowed)),
|
||||
"allowed first fragment should pass and be recorded")
|
||||
|
||||
// A second offset-zero fragment to a denied port supersedes the datagram's
|
||||
// verdict; it is dropped and must not leave the allow in place.
|
||||
denied := firstFragmentUDP(t, 0x5A5A, 9999, 32)
|
||||
require.True(t, m.filterInbound(denied, len(denied)),
|
||||
"denied offset-zero fragment must be dropped")
|
||||
|
||||
trailing := trailingFragment(t, 0x5A5A, 5, false, 24)
|
||||
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||
"trailing fragment must be denied after the datagram was poisoned")
|
||||
}
|
||||
|
||||
// TestFragment_TinyFirstDropped covers the tiny-fragment attack: a first
|
||||
// fragment too small to contain the full transport header can't be
|
||||
// ACL-evaluated and must be dropped.
|
||||
func TestFragment_TinyFirstDropped(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
allowUDP(t, m, 8080)
|
||||
|
||||
// IPv4 header + 4 raw bytes, MF set, offset 0: too small for the 8-byte UDP
|
||||
// header, so it decodes to L3 only.
|
||||
tiny := trailingFragment(t, 0x5555, 0, true, 4)
|
||||
require.True(t, m.filterInbound(tiny, len(tiny)),
|
||||
"tiny first fragment without a full L4 header must be dropped")
|
||||
}
|
||||
|
||||
// TestFragment_TCPFirstFragment verifies the TCP arm of the transport decode: a
|
||||
// first fragment carrying the full 20-byte TCP header is ACL-evaluated and its
|
||||
// trailing fragments inherit the verdict.
|
||||
func TestFragment_TCPFirstFragment(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||
require.NoError(t, err)
|
||||
|
||||
// TCP header (20) + 12 data = 32 bytes -> headerEnd = 4 octets.
|
||||
first := firstFragmentTCP(t, 0x6666, 8080)
|
||||
require.False(t, m.filterInbound(first, len(first)),
|
||||
"allowed TCP first fragment should pass and be recorded")
|
||||
|
||||
trailing := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x6666, 4, false, 24)
|
||||
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||
"trailing fragment of an allowed TCP datagram should pass")
|
||||
}
|
||||
|
||||
// TestFragment_TCPTinyFirstDropped verifies the TCP minimum header length: 12
|
||||
// bytes would satisfy a UDP header but falls short of the 20-byte TCP header.
|
||||
func TestFragment_TCPTinyFirstDropped(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrc), fw.ProtocolTCP, nil,
|
||||
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||
require.NoError(t, err)
|
||||
|
||||
tiny := trailingFragmentTo(t, fragTestDst, layers.IPProtocolTCP, 0x7777, 0, true, 12)
|
||||
require.True(t, m.filterInbound(tiny, len(tiny)),
|
||||
"first fragment shorter than the TCP header must be dropped")
|
||||
}
|
||||
|
||||
// TestFragment_ConntrackAllowsFirstFragment verifies the conntrack branch: reply
|
||||
// fragments of an outbound-established UDP flow pass without any inbound rule.
|
||||
func TestFragment_ConntrackAllowsFirstFragment(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
|
||||
out := outboundUDPPacket(t, 12345, 40000)
|
||||
require.False(t, m.filterOutbound(out, len(out)))
|
||||
|
||||
first := firstFragmentUDP(t, 0x8888, 12345, 32)
|
||||
require.False(t, m.filterInbound(first, len(first)),
|
||||
"reply first fragment should pass via conntrack")
|
||||
|
||||
trailing := trailingFragment(t, 0x8888, 5, false, 24)
|
||||
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||
"trailing fragment of a tracked flow should pass")
|
||||
}
|
||||
|
||||
// TestFragment_RoutingDisabledDropsFragment verifies routed first fragments are
|
||||
// dropped when routing is disabled.
|
||||
func TestFragment_RoutingDisabledDropsFragment(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
m.routingEnabled.Store(false)
|
||||
|
||||
first := firstFragmentUDPTo(t, "198.51.100.10", 0x9999, 8080, 32)
|
||||
require.True(t, m.filterInbound(first, len(first)),
|
||||
"routed first fragment must be dropped when routing is disabled")
|
||||
}
|
||||
|
||||
// TestFragment_RouteACL verifies the route-ACL branch: fragments to a non-local
|
||||
// destination follow the route rules, allowed datagrams pass their trailing
|
||||
// fragments and denied ones don't.
|
||||
func TestFragment_RouteACL(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
m.routingEnabled.Store(true)
|
||||
m.nativeRouter.Store(false)
|
||||
|
||||
_, err := m.AddRouteFiltering(
|
||||
[]byte("rt-1"),
|
||||
[]netip.Prefix{netip.MustParsePrefix("100.10.0.0/16")},
|
||||
fw.Network{Prefix: netip.MustParsePrefix("198.51.100.0/24")},
|
||||
fw.ProtocolUDP,
|
||||
nil,
|
||||
&fw.Port{Values: []uint16{8080}},
|
||||
fw.ActionAccept,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
first := firstFragmentUDPTo(t, "198.51.100.10", 0xAAAA, 8080, 32)
|
||||
require.False(t, m.filterInbound(first, len(first)),
|
||||
"route-ACL-allowed first fragment should pass")
|
||||
trailing := trailingFragmentTo(t, "198.51.100.10", layers.IPProtocolUDP, 0xAAAA, 5, false, 24)
|
||||
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||
"trailing fragment of an allowed routed datagram should pass")
|
||||
|
||||
denied := firstFragmentUDPTo(t, "198.51.100.10", 0xBBBB, 9999, 32)
|
||||
require.True(t, m.filterInbound(denied, len(denied)),
|
||||
"route-ACL-denied first fragment must be dropped")
|
||||
deniedTrailing := trailingFragmentTo(t, "198.51.100.10", layers.IPProtocolUDP, 0xBBBB, 5, false, 24)
|
||||
require.True(t, m.filterInbound(deniedTrailing, len(deniedTrailing)),
|
||||
"trailing fragment of a denied routed datagram must be dropped")
|
||||
}
|
||||
|
||||
// TestFragment_ExpiredVerdictDropsTrailing verifies a verdict older than the
|
||||
// tracker timeout no longer admits trailing fragments.
|
||||
func TestFragment_ExpiredVerdictDropsTrailing(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
allowUDP(t, m, 8080)
|
||||
|
||||
first := firstFragmentUDP(t, 0xCCCC, 8080, 32)
|
||||
require.False(t, m.filterInbound(first, len(first)))
|
||||
|
||||
m.fragments.mutex.Lock()
|
||||
for key, entry := range m.fragments.entries {
|
||||
entry.recordedAt = time.Now().Add(-defaultFragmentTimeout - time.Second)
|
||||
m.fragments.entries[key] = entry
|
||||
}
|
||||
m.fragments.mutex.Unlock()
|
||||
|
||||
trailing := trailingFragment(t, 0xCCCC, 5, false, 24)
|
||||
require.True(t, m.filterInbound(trailing, len(trailing)),
|
||||
"trailing fragment after verdict expiry must be dropped")
|
||||
}
|
||||
|
||||
// TestFragment_CapacityFailsClosed verifies the table cap: at capacity, new
|
||||
// datagram verdicts are not recorded (their trailing fragments are dropped)
|
||||
// while already-recorded datagrams keep working.
|
||||
func TestFragment_CapacityFailsClosed(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
allowUDP(t, m, 8080)
|
||||
|
||||
m.fragments.mutex.Lock()
|
||||
m.fragments.maxEntries = 1
|
||||
m.fragments.mutex.Unlock()
|
||||
|
||||
first1 := firstFragmentUDP(t, 0x0101, 8080, 32)
|
||||
require.False(t, m.filterInbound(first1, len(first1)))
|
||||
|
||||
first2 := firstFragmentUDP(t, 0x0202, 8080, 32)
|
||||
require.False(t, m.filterInbound(first2, len(first2)),
|
||||
"first fragment itself still passes at capacity")
|
||||
|
||||
trailing2 := trailingFragment(t, 0x0202, 5, false, 24)
|
||||
require.True(t, m.filterInbound(trailing2, len(trailing2)),
|
||||
"trailing fragment of an unrecorded datagram must be dropped at capacity")
|
||||
|
||||
trailing1 := trailingFragment(t, 0x0101, 5, false, 24)
|
||||
require.False(t, m.filterInbound(trailing1, len(trailing1)),
|
||||
"already-recorded datagram should keep passing at capacity")
|
||||
}
|
||||
|
||||
// v6FragmentHeader builds the 8-byte IPv6 fragment extension header for the
|
||||
// given inner protocol, offset (8-byte units), More Fragments bit and id.
|
||||
func v6FragmentHeader(proto layers.IPProtocol, offsetOctets uint16, moreFragments bool, id uint32) []byte {
|
||||
offsetFlags := offsetOctets << 3
|
||||
if moreFragments {
|
||||
offsetFlags |= 1
|
||||
}
|
||||
hdr := make([]byte, 8)
|
||||
hdr[0] = uint8(proto)
|
||||
binary.BigEndian.PutUint16(hdr[2:4], offsetFlags)
|
||||
binary.BigEndian.PutUint32(hdr[4:8], id)
|
||||
return hdr
|
||||
}
|
||||
|
||||
func v6UDPHeader(dstPort uint16, dataLen int) []byte {
|
||||
hdr := make([]byte, 8)
|
||||
binary.BigEndian.PutUint16(hdr[0:2], 40000)
|
||||
binary.BigEndian.PutUint16(hdr[2:4], dstPort)
|
||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(8+dataLen))
|
||||
return hdr
|
||||
}
|
||||
|
||||
// firstFragmentUDPv6 builds the first fragment of a fragmented IPv6 UDP
|
||||
// datagram: fragment header (offset 0, More Fragments set) + full UDP header +
|
||||
// data.
|
||||
func firstFragmentUDPv6(tb testing.TB, id uint32, dstPort uint16, dataLen int) []byte {
|
||||
tb.Helper()
|
||||
return fragmentUDPv6(tb, id, dstPort, dataLen, true)
|
||||
}
|
||||
|
||||
// fragmentUDPv6 builds an offset-zero IPv6 UDP fragment. With moreFragments
|
||||
// false it is an atomic fragment (a complete datagram, RFC 6946).
|
||||
func fragmentUDPv6(tb testing.TB, id uint32, dstPort uint16, dataLen int, moreFragments bool) []byte {
|
||||
tb.Helper()
|
||||
|
||||
ip := &layers.IPv6{
|
||||
Version: 6,
|
||||
NextHeader: layers.IPProtocolIPv6Fragment,
|
||||
HopLimit: 64,
|
||||
SrcIP: net.ParseIP(fragTestSrcV6),
|
||||
DstIP: net.ParseIP(fragTestDstV6),
|
||||
}
|
||||
payload := append(v6FragmentHeader(layers.IPProtocolUDP, 0, moreFragments, id), v6UDPHeader(dstPort, dataLen)...)
|
||||
payload = append(payload, make([]byte, dataLen)...)
|
||||
|
||||
buf := gopacket.NewSerializeBuffer()
|
||||
require.NoError(tb, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true}, ip, gopacket.Payload(payload)))
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
// trailingFragmentV6 builds a non-first IPv6 fragment: fragment header at the
|
||||
// given offset carrying raw data and no transport header.
|
||||
func trailingFragmentV6(tb testing.TB, id uint32, offsetOctets uint16, moreFragments bool, dataLen int) []byte {
|
||||
tb.Helper()
|
||||
|
||||
ip := &layers.IPv6{
|
||||
Version: 6,
|
||||
NextHeader: layers.IPProtocolIPv6Fragment,
|
||||
HopLimit: 64,
|
||||
SrcIP: net.ParseIP(fragTestSrcV6),
|
||||
DstIP: net.ParseIP(fragTestDstV6),
|
||||
}
|
||||
payload := append(v6FragmentHeader(layers.IPProtocolUDP, offsetOctets, moreFragments, id), make([]byte, dataLen)...)
|
||||
|
||||
buf := gopacket.NewSerializeBuffer()
|
||||
require.NoError(tb, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true}, ip, gopacket.Payload(payload)))
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
// TestFragmentV6_TrailingWithoutFirstDropped verifies the IPv6 bypass is closed:
|
||||
// a trailing fragment with no allowed first fragment is dropped.
|
||||
func TestFragmentV6_TrailingWithoutFirstDropped(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
|
||||
frag := trailingFragmentV6(t, 0xAABBCCDD, 100, false, 40)
|
||||
require.True(t, m.filterInbound(frag, len(frag)),
|
||||
"IPv6 trailing fragment without an allowed first fragment must be dropped")
|
||||
}
|
||||
|
||||
// TestFragmentV6_AllowedFirstPassesTrailing verifies IPv6 fragments are
|
||||
// evaluated like IPv4: an allowed first fragment lets its trailing fragments
|
||||
// through.
|
||||
func TestFragmentV6_AllowedFirstPassesTrailing(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrcV6), fw.ProtocolUDP, nil,
|
||||
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||
require.NoError(t, err)
|
||||
|
||||
// First fragment: UDP header (8) + 32 data = 40 octets -> headerEnd = 5.
|
||||
first := firstFragmentUDPv6(t, 0xAABBCCDD, 8080, 32)
|
||||
require.False(t, m.filterInbound(first, len(first)),
|
||||
"allowed IPv6 first fragment should pass and be recorded")
|
||||
|
||||
trailing := trailingFragmentV6(t, 0xAABBCCDD, 5, false, 24)
|
||||
require.False(t, m.filterInbound(trailing, len(trailing)),
|
||||
"trailing fragment of an allowed IPv6 datagram should pass")
|
||||
}
|
||||
|
||||
// TestFragmentV6_AtomicNotCached verifies an IPv6 atomic fragment (fragment
|
||||
// header with offset 0 and no More Fragments, a complete datagram per RFC 6946)
|
||||
// is evaluated but not recorded, so a flood of allowed atomic fragments can't
|
||||
// exhaust the verdict table.
|
||||
func TestFragmentV6_AtomicNotCached(t *testing.T) {
|
||||
m := newFragmentTestManager(t)
|
||||
_, err := m.AddPeerFiltering(nil, net.ParseIP(fragTestSrcV6), fw.ProtocolUDP, nil,
|
||||
&fw.Port{Values: []uint16{8080}}, fw.ActionAccept, "")
|
||||
require.NoError(t, err)
|
||||
|
||||
atomic := fragmentUDPv6(t, 0xA70301C, 8080, 16, false)
|
||||
require.False(t, m.filterInbound(atomic, len(atomic)),
|
||||
"allowed IPv6 atomic fragment should pass")
|
||||
|
||||
m.fragments.mutex.Lock()
|
||||
n := len(m.fragments.entries)
|
||||
m.fragments.mutex.Unlock()
|
||||
require.Zero(t, n, "atomic fragment must not create a verdict entry")
|
||||
|
||||
// A genuine fragmented datagram (More Fragments set) is still recorded.
|
||||
first := fragmentUDPv6(t, 0xBEEF, 8080, 32, true)
|
||||
require.False(t, m.filterInbound(first, len(first)))
|
||||
m.fragments.mutex.Lock()
|
||||
n = len(m.fragments.entries)
|
||||
m.fragments.mutex.Unlock()
|
||||
require.Equal(t, 1, n, "genuine first fragment must record a verdict")
|
||||
}
|
||||
@@ -3,31 +3,14 @@
|
||||
package netstack
|
||||
|
||||
import (
|
||||
"net"
|
||||
"fmt"
|
||||
"os"
|
||||
"strconv"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
const (
|
||||
EnvUseNetstackMode = "NB_USE_NETSTACK_MODE"
|
||||
|
||||
// EnvSocks5ListenerPort overrides the port the SOCKS5 proxy listens on.
|
||||
EnvSocks5ListenerPort = "NB_SOCKS5_LISTENER_PORT"
|
||||
|
||||
// EnvSocks5ListenerAddress overrides the host/IP the SOCKS5 proxy binds to.
|
||||
// The proxy is a bridge for local host applications into the userspace
|
||||
// WireGuard netstack, so it binds to loopback by default. Override this only
|
||||
// when the proxy must be reachable from other hosts (e.g. a container
|
||||
// gateway); doing so exposes an unauthenticated SOCKS5 proxy on that
|
||||
// address.
|
||||
EnvSocks5ListenerAddress = "NB_SOCKS5_LISTENER_ADDRESS"
|
||||
|
||||
// defaultSocks5Host is the loopback address the SOCKS5 proxy binds to unless
|
||||
// overridden via EnvSocks5ListenerAddress.
|
||||
defaultSocks5Host = "127.0.0.1"
|
||||
)
|
||||
const EnvUseNetstackMode = "NB_USE_NETSTACK_MODE"
|
||||
|
||||
// IsEnabled todo: move these function to cmd layer
|
||||
func IsEnabled() bool {
|
||||
@@ -35,40 +18,24 @@ func IsEnabled() bool {
|
||||
}
|
||||
|
||||
func ListenAddr() string {
|
||||
return net.JoinHostPort(listenHost(), strconv.Itoa(listenPort()))
|
||||
}
|
||||
|
||||
// listenHost returns the host/IP the SOCKS5 proxy binds to. It defaults to
|
||||
// loopback and only honors EnvSocks5ListenerAddress when it holds a valid IP.
|
||||
func listenHost() string {
|
||||
addr := os.Getenv(EnvSocks5ListenerAddress)
|
||||
if addr == "" {
|
||||
return defaultSocks5Host
|
||||
}
|
||||
if net.ParseIP(addr) == nil {
|
||||
log.Warnf("invalid socks5 listener address %q, falling back to default: %s", addr, defaultSocks5Host)
|
||||
return defaultSocks5Host
|
||||
}
|
||||
return addr
|
||||
}
|
||||
|
||||
// listenPort returns the port the SOCKS5 proxy binds to, defaulting to
|
||||
// DefaultSocks5Port when EnvSocks5ListenerPort is unset or invalid.
|
||||
func listenPort() int {
|
||||
sPort := os.Getenv(EnvSocks5ListenerPort)
|
||||
sPort := os.Getenv("NB_SOCKS5_LISTENER_PORT")
|
||||
if sPort == "" {
|
||||
return DefaultSocks5Port
|
||||
return listenAddr(DefaultSocks5Port)
|
||||
}
|
||||
|
||||
port, err := strconv.Atoi(sPort)
|
||||
if err != nil {
|
||||
log.Warnf("invalid socks5 listener port, unable to convert it to int, falling back to default: %d", DefaultSocks5Port)
|
||||
return DefaultSocks5Port
|
||||
return listenAddr(DefaultSocks5Port)
|
||||
}
|
||||
if port < 1 || port > 65535 {
|
||||
log.Warnf("invalid socks5 listener port, it should be in the range 1-65535, falling back to default: %d", DefaultSocks5Port)
|
||||
return DefaultSocks5Port
|
||||
return listenAddr(DefaultSocks5Port)
|
||||
}
|
||||
|
||||
return port
|
||||
return listenAddr(port)
|
||||
}
|
||||
|
||||
func listenAddr(port int) string {
|
||||
return fmt.Sprintf("0.0.0.0:%d", port)
|
||||
}
|
||||
|
||||
@@ -1,63 +0,0 @@
|
||||
//go:build !js
|
||||
|
||||
package netstack
|
||||
|
||||
import (
|
||||
"net"
|
||||
"strconv"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestListenAddr_DefaultsToLoopback(t *testing.T) {
|
||||
// No env overrides: must bind loopback, never all interfaces.
|
||||
got := ListenAddr()
|
||||
want := net.JoinHostPort("127.0.0.1", strconv.Itoa(DefaultSocks5Port))
|
||||
if got != want {
|
||||
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListenAddr_AddressOverride(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
env string
|
||||
want string
|
||||
}{
|
||||
{name: "valid override honored", env: "0.0.0.0", want: "0.0.0.0"},
|
||||
{name: "valid specific ip honored", env: "10.0.0.5", want: "10.0.0.5"},
|
||||
{name: "ipv6 loopback bracketed", env: "::1", want: "::1"},
|
||||
{name: "invalid falls back to loopback", env: "not-an-ip", want: "127.0.0.1"},
|
||||
{name: "empty falls back to loopback", env: "", want: "127.0.0.1"},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Setenv(EnvSocks5ListenerAddress, tc.env)
|
||||
want := net.JoinHostPort(tc.want, strconv.Itoa(DefaultSocks5Port))
|
||||
if got := ListenAddr(); got != want {
|
||||
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestListenAddr_PortOverride(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
env string
|
||||
want int
|
||||
}{
|
||||
{name: "valid port honored", env: "1081", want: 1081},
|
||||
{name: "non-numeric falls back", env: "abc", want: DefaultSocks5Port},
|
||||
{name: "out of range falls back", env: "70000", want: DefaultSocks5Port},
|
||||
{name: "zero falls back", env: "0", want: DefaultSocks5Port},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Setenv(EnvSocks5ListenerPort, tc.env)
|
||||
want := net.JoinHostPort("127.0.0.1", strconv.Itoa(tc.want))
|
||||
if got := ListenAddr(); got != want {
|
||||
t.Fatalf("ListenAddr() = %q, want %q", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -299,7 +299,7 @@ func (d *DeviceAuthorizationFlow) WaitToken(ctx context.Context, info AuthFlowIn
|
||||
UseIDToken: d.providerConfig.UseIDToken,
|
||||
}
|
||||
|
||||
err = validateTokenAudience(tokenInfo.GetTokenToUse(), d.providerConfig.Audience)
|
||||
err = isValidAccessToken(tokenInfo.GetTokenToUse(), d.providerConfig.Audience)
|
||||
if err != nil {
|
||||
return TokenInfo{}, fmt.Errorf("validate access token failed with error: %v", err)
|
||||
}
|
||||
|
||||
@@ -306,7 +306,7 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo,
|
||||
audience = p.providerConfig.ClientID
|
||||
}
|
||||
|
||||
if err := validateTokenAudience(tokenInfo.GetTokenToUse(), audience); err != nil {
|
||||
if err := isValidAccessToken(tokenInfo.GetTokenToUse(), audience); err != nil {
|
||||
return TokenInfo{}, fmt.Errorf("authentication failed: invalid access token - %w", err)
|
||||
}
|
||||
|
||||
@@ -320,11 +320,6 @@ func (p *PKCEAuthorizationFlow) parseOAuthToken(token *oauth2.Token) (TokenInfo,
|
||||
return tokenInfo, nil
|
||||
}
|
||||
|
||||
// parseEmailFromIDToken extracts the email (or name) claim from an ID token
|
||||
// without verifying its signature. The value is best-effort and used only as a
|
||||
// UX convenience (login hint prefill and display); it never drives an
|
||||
// authorization decision. The authoritative identity is established server-side
|
||||
// from the signature-verified token.
|
||||
func parseEmailFromIDToken(token string) (string, error) {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) < 2 {
|
||||
|
||||
@@ -24,7 +24,11 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
maxPastHorizon = 30 * 24 * time.Hour
|
||||
// Skew tolerates a small clock difference between the management
|
||||
// server and this peer before treating a deadline as "in the past".
|
||||
// Slightly above typical NTP drift; tight enough that the UI doesn't
|
||||
// paint a stale expiry as if it were valid.
|
||||
Skew = 30 * time.Second
|
||||
|
||||
// maxDeadlineHorizon caps how far in the future an accepted deadline
|
||||
// can sit. A timestamp beyond this is almost certainly a protocol
|
||||
@@ -53,7 +57,7 @@ var (
|
||||
ErrDeadlineTooFarFuture = errors.New("session deadline too far in the future")
|
||||
|
||||
// ErrDeadlineInPast is returned by Update when the supplied deadline
|
||||
// is more than maxPastHorizon in the past.
|
||||
// is more than Skew in the past.
|
||||
ErrDeadlineInPast = errors.New("session deadline in the past")
|
||||
)
|
||||
|
||||
@@ -62,14 +66,15 @@ var (
|
||||
// for deadline change/clear, PublishEvent for the two warnings); tests pass
|
||||
// a fake recorder so the same surface is observable without an engine.
|
||||
//
|
||||
// While the watcher runs, it owns the deadline propagated to the recorder:
|
||||
// every set, clear and sanity-check rejection routes the value through
|
||||
// SetSessionExpiresAt, so the SubscribeStatus snapshot the UI reads can
|
||||
// never drift from the watcher's timer state. (SetSessionExpiresAt fans
|
||||
// out its own state-change notification, so no separate notify is needed.)
|
||||
// The recorder is server-scoped and outlives this engine-scoped watcher;
|
||||
// Close deliberately leaves the recorder value in place so transient engine
|
||||
// restarts don't blank it — the client run loop clears it on real teardown.
|
||||
// The watcher is the single owner of the deadline propagated to the
|
||||
// recorder: every set, clear, sanity-check rejection and Close routes the
|
||||
// value through SetSessionExpiresAt, so the SubscribeStatus snapshot the UI
|
||||
// reads can never drift from the watcher's timer state. (SetSessionExpiresAt
|
||||
// fans out its own state-change notification, so no separate notify is
|
||||
// needed.) The recorder is server-scoped and outlives this engine-scoped
|
||||
// watcher — without the Close-time clear a teardown (Down, or the Down+Up of
|
||||
// a profile switch) would leave the next session showing the previous one's
|
||||
// stale "expires in" value.
|
||||
//
|
||||
// PublishEvent's signature mirrors peer.Status.PublishEvent: the watcher
|
||||
// composes the metadata internally so the wire format (MetaSession*) is
|
||||
@@ -130,13 +135,10 @@ func NewWithLeads(lead, final time.Duration, recorder StatusRecorder) *Watcher {
|
||||
// was disabled).
|
||||
//
|
||||
// Same-value updates are no-ops. A different non-zero value cancels any
|
||||
// pending timer, resets the "already fired" guards, and — when the
|
||||
// deadline lies in the future — arms fresh warning timers. A deadline
|
||||
// already in the past (within maxPastHorizon) is recorded as-is with no
|
||||
// timers: the session has expired and consumers render it that way.
|
||||
// pending timer, resets the "already fired" guard, and arms a new one.
|
||||
//
|
||||
// Returns one of the sentinel Err* values when the deadline fails the
|
||||
// sanity checks (pre-epoch, far future, or past beyond maxPastHorizon).
|
||||
// sanity checks (pre-epoch, far future, or in the past beyond Skew).
|
||||
// In every error case the watcher first clears its state so it stays
|
||||
// consistent with what the caller will push into its other sinks (e.g.
|
||||
// applySessionDeadline forces a zero deadline into the status recorder
|
||||
@@ -161,7 +163,7 @@ func (w *Watcher) Update(deadline time.Time) error {
|
||||
case deadline.After(now.Add(maxDeadlineHorizon)):
|
||||
w.clearLocked()
|
||||
return fmt.Errorf("%w: %v", ErrDeadlineTooFarFuture, deadline)
|
||||
case deadline.Before(now.Add(-maxPastHorizon)):
|
||||
case deadline.Before(now.Add(-Skew)):
|
||||
w.clearLocked()
|
||||
return fmt.Errorf("%w: %v (now=%v)", ErrDeadlineInPast, deadline, now)
|
||||
}
|
||||
@@ -181,9 +183,7 @@ func (w *Watcher) Update(deadline time.Time) error {
|
||||
w.finalFiredAt = time.Time{}
|
||||
w.dismissedAt = time.Time{}
|
||||
|
||||
if deadline.After(now) {
|
||||
w.armTimerLocked(deadline)
|
||||
}
|
||||
w.armTimerLocked(deadline)
|
||||
recorder := w.recorder
|
||||
w.mu.Unlock()
|
||||
if recorder != nil {
|
||||
@@ -227,25 +227,30 @@ func (w *Watcher) Dismiss() {
|
||||
log.Infof("auth session final-warning dismissed for deadline %s", w.current.Format(time.RFC3339))
|
||||
}
|
||||
|
||||
// Close stops any pending timer. Update calls after Close are ignored.
|
||||
// The recorder keeps its deadline: the watcher is engine-scoped and closes
|
||||
// on every engine restart (network change, sleep/wake, stream errors)
|
||||
// while the SSO deadline stays valid across those, so clearing here would
|
||||
// blank the UI's "expires in" row on every transient reconnect. The
|
||||
// client run loop clears the server-scoped recorder when it exits for
|
||||
// real (Down, profile switch, permanent login failure).
|
||||
// Close stops any pending timer and drops the deadline on the status
|
||||
// recorder. Update calls after Close are ignored. Clearing the recorder
|
||||
// here is what keeps a teardown (Down, or the Down+Up of a profile switch)
|
||||
// from leaving the next session showing this one's stale "expires in"
|
||||
// value — the recorder is server-scoped and outlives this engine-scoped
|
||||
// watcher, so nothing else drops the anchor on teardown.
|
||||
func (w *Watcher) Close() {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
if w.closed {
|
||||
w.mu.Unlock()
|
||||
return
|
||||
}
|
||||
w.closed = true
|
||||
w.stopTimerLocked()
|
||||
hadDeadline := !w.current.IsZero()
|
||||
w.current = time.Time{}
|
||||
w.firedAt = time.Time{}
|
||||
w.finalFiredAt = time.Time{}
|
||||
w.dismissedAt = time.Time{}
|
||||
recorder := w.recorder
|
||||
w.mu.Unlock()
|
||||
if recorder != nil && hadDeadline {
|
||||
recorder.SetSessionExpiresAt(time.Time{})
|
||||
}
|
||||
}
|
||||
|
||||
// clearLocked drops the tracked deadline and notifies the recorder so
|
||||
|
||||
@@ -224,13 +224,11 @@ func TestNewDeadlineCancelsPriorTimer(t *testing.T) {
|
||||
|
||||
func TestRefreshAfterFireArmsNewWarning(t *testing.T) {
|
||||
r := &fakeRecorder{}
|
||||
lead := 150 * time.Millisecond
|
||||
lead := 30 * time.Millisecond
|
||||
w := newWatcher(lead, r)
|
||||
defer w.Close()
|
||||
|
||||
// Warning fires ~20ms in; the deadline itself stays 150ms away so the
|
||||
// replacement below lands well before it.
|
||||
first := time.Now().Add(170 * time.Millisecond)
|
||||
first := time.Now().Add(50 * time.Millisecond)
|
||||
_ = w.Update(first)
|
||||
|
||||
// Wait for stateChange + warning of the first cycle.
|
||||
@@ -308,29 +306,7 @@ func TestUpdateRejectsTooFarFuture(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateRecentPastRecordedAsExpired(t *testing.T) {
|
||||
r := &fakeRecorder{}
|
||||
w := newWatcher(50*time.Millisecond, r)
|
||||
defer w.Close()
|
||||
|
||||
d := time.Now().Add(-1 * time.Hour)
|
||||
if err := w.Update(d); err != nil {
|
||||
t.Fatalf("recent-past Update should succeed, got %v", err)
|
||||
}
|
||||
if !w.Deadline().Equal(d) {
|
||||
t.Fatalf("expected deadline to be recorded, got %v want %v", w.Deadline(), d)
|
||||
}
|
||||
if got := r.deadline(); !got.Equal(d) {
|
||||
t.Fatalf("recorder deadline = %v, want %v", got, d)
|
||||
}
|
||||
|
||||
time.Sleep(80 * time.Millisecond)
|
||||
if n := countWhere(r.snapshot(), func(e event) bool { return e.kind == publish }); n != 0 {
|
||||
t.Fatalf("no warning events may fire for an already-past deadline, got %+v", r.snapshot())
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateAncientPastRejected(t *testing.T) {
|
||||
func TestUpdateInPastClearsDeadline(t *testing.T) {
|
||||
r := &fakeRecorder{}
|
||||
w := newWatcher(50*time.Millisecond, r)
|
||||
defer w.Close()
|
||||
@@ -342,12 +318,12 @@ func TestUpdateAncientPastRejected(t *testing.T) {
|
||||
// Drain the stateChange from the seed.
|
||||
waitForEvents(t, r, 1)
|
||||
|
||||
err := w.Update(time.Now().Add(-31 * 24 * time.Hour))
|
||||
err := w.Update(time.Now().Add(-1 * time.Hour))
|
||||
if !errors.Is(err, ErrDeadlineInPast) {
|
||||
t.Fatalf("want ErrDeadlineInPast, got %v", err)
|
||||
}
|
||||
if !w.Deadline().IsZero() {
|
||||
t.Fatalf("rejected ancient-past update must clear the deadline, got %v", w.Deadline())
|
||||
t.Fatalf("in-past update must clear the deadline, got %v", w.Deadline())
|
||||
}
|
||||
events := waitForEvents(t, r, 2)
|
||||
if events[1].kind != stateChange {
|
||||
@@ -355,25 +331,39 @@ func TestUpdateAncientPastRejected(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateWithinSkewAccepted(t *testing.T) {
|
||||
r := &fakeRecorder{}
|
||||
w := newWatcher(50*time.Millisecond, r)
|
||||
defer w.Close()
|
||||
|
||||
// 5 seconds in the past is within the 30s Skew tolerance — accept it.
|
||||
d := time.Now().Add(-5 * time.Second)
|
||||
if err := w.Update(d); err != nil {
|
||||
t.Fatalf("within-skew Update should succeed, got %v", err)
|
||||
}
|
||||
if !w.Deadline().Equal(d) {
|
||||
t.Fatalf("expected deadline to be applied, got %v want %v", w.Deadline(), d)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloseSilencesUpdates(t *testing.T) {
|
||||
r := &fakeRecorder{}
|
||||
w := newWatcher(50*time.Millisecond, r)
|
||||
w.Close()
|
||||
|
||||
if err := w.Update(time.Now().Add(time.Hour)); err != nil {
|
||||
t.Fatalf("Update after Close: want nil, got %v", err)
|
||||
}
|
||||
_ = w.Update(time.Now().Add(time.Hour))
|
||||
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
if got := r.snapshot(); len(got) != 0 {
|
||||
t.Fatalf("expected no events after Close, got %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCloseKeepsRecorderDeadline pins the reconnect-flap fix: the watcher
|
||||
// closes on every engine restart (network change, sleep/wake) while the
|
||||
// SSO deadline stays valid across those, so Close must leave the
|
||||
// server-scoped recorder's value in place. The client run loop clears the
|
||||
// recorder when it exits for real.
|
||||
func TestCloseKeepsRecorderDeadline(t *testing.T) {
|
||||
// TestCloseClearsRecorderDeadline pins the profile-switch fix: a watcher
|
||||
// holding a live deadline must zero the recorder on Close so the next
|
||||
// engine's watcher (and the UI reading the shared server-scoped recorder)
|
||||
// doesn't start out showing the previous session's stale "expires in".
|
||||
func TestCloseClearsRecorderDeadline(t *testing.T) {
|
||||
r := &fakeRecorder{}
|
||||
w := newWatcher(time.Hour, r)
|
||||
|
||||
@@ -387,8 +377,8 @@ func TestCloseKeepsRecorderDeadline(t *testing.T) {
|
||||
|
||||
w.Close()
|
||||
|
||||
if got := r.deadline(); !got.Equal(d) {
|
||||
t.Fatalf("recorder deadline after Close = %v, want %v", got, d)
|
||||
if got := r.deadline(); !got.IsZero() {
|
||||
t.Fatalf("recorder deadline after Close = %v, want zero", got)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -20,26 +20,14 @@ func randomBytesInHex(count int) (string, error) {
|
||||
return hex.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
// validateTokenAudience checks that the token is a well-formed JWT whose
|
||||
// audience claim matches the expected audience.
|
||||
//
|
||||
// It does NOT verify the token's cryptographic signature and therefore must not
|
||||
// be treated as an authenticity check. The token is obtained by the client
|
||||
// directly from the IdP token endpoint over TLS, and its signature is verified
|
||||
// server-side by the management server against the IdP's JWKS
|
||||
// (see shared/auth/jwt/validator.go). This function is only a client-side
|
||||
// sanity check that the returned token targets the expected audience.
|
||||
func validateTokenAudience(token string, audience string) error {
|
||||
// isValidAccessToken is a simple validation of the access token
|
||||
func isValidAccessToken(token string, audience string) error {
|
||||
if token == "" {
|
||||
return fmt.Errorf("token received is empty")
|
||||
}
|
||||
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != 3 {
|
||||
return fmt.Errorf("token is not a well-formed JWT")
|
||||
}
|
||||
|
||||
claimsString, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
encodedClaims := strings.Split(token, ".")[1]
|
||||
claimsString, err := base64.RawURLEncoding.DecodeString(encodedClaims)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -1,108 +0,0 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// makeJWT builds an unsigned JWT-shaped string (header.payload.signature) with
|
||||
// the given claims payload. The signature part is arbitrary because
|
||||
// validateTokenAudience intentionally does not verify it.
|
||||
func makeJWT(t *testing.T, claims map[string]interface{}) string {
|
||||
t.Helper()
|
||||
header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"RS256","typ":"JWT"}`))
|
||||
payloadBytes, err := json.Marshal(claims)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal claims: %v", err)
|
||||
}
|
||||
payload := base64.RawURLEncoding.EncodeToString(payloadBytes)
|
||||
return header + "." + payload + ".unverified-signature"
|
||||
}
|
||||
|
||||
func TestValidateTokenAudience(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
token string
|
||||
audience string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "empty token",
|
||||
token: "",
|
||||
audience: "netbird",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "not a JWT - no dots",
|
||||
token: "notajwt",
|
||||
audience: "netbird",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "not a JWT - two parts only",
|
||||
token: "header.payload",
|
||||
audience: "netbird",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "matching string audience",
|
||||
token: makeJWT(t, map[string]interface{}{"aud": "netbird"}),
|
||||
audience: "netbird",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "mismatching string audience",
|
||||
token: makeJWT(t, map[string]interface{}{"aud": "other"}),
|
||||
audience: "netbird",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "matching audience in array",
|
||||
token: makeJWT(t, map[string]interface{}{"aud": []interface{}{"other", "netbird"}}),
|
||||
audience: "netbird",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "mismatching audience array",
|
||||
token: makeJWT(t, map[string]interface{}{"aud": []interface{}{"a", "b"}}),
|
||||
audience: "netbird",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "missing audience claim",
|
||||
token: makeJWT(t, map[string]interface{}{"sub": "user"}),
|
||||
audience: "netbird",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid base64 payload",
|
||||
token: "header.!!!not-base64!!!.sig",
|
||||
audience: "netbird",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := validateTokenAudience(tc.token, tc.audience)
|
||||
if tc.wantErr && err == nil {
|
||||
t.Fatalf("expected error, got nil")
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Fatalf("expected no error, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestValidateTokenAudienceNoPanic guards the regression where a non-empty
|
||||
// token without the JWT dot structure caused an index-out-of-range panic.
|
||||
func TestValidateTokenAudienceNoPanic(t *testing.T) {
|
||||
inputs := []string{"a", ".", "a.", "aaaa", "no-dots-here"}
|
||||
for _, in := range inputs {
|
||||
if err := validateTokenAudience(in, "netbird"); err == nil {
|
||||
t.Fatalf("expected error for malformed token %q, got nil", in)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -257,10 +257,7 @@ func (c *ConnectClient) run(mobileDependency MobileDependency, runningChan chan
|
||||
log.Errorf("failed to clean up temporary installer file: %v", err)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
c.statusRecorder.SetSessionExpiresAt(time.Time{})
|
||||
c.statusRecorder.ClientStop()
|
||||
}()
|
||||
defer c.statusRecorder.ClientStop()
|
||||
operation := func() error {
|
||||
// if context cancelled we not start new backoff cycle
|
||||
if c.ctx.Err() != nil {
|
||||
|
||||
@@ -480,6 +480,7 @@ func (g *BundleGenerator) addStatus() error {
|
||||
|
||||
fullStatus := g.statusRecorder.GetFullStatus()
|
||||
protoFullStatus := nbstatus.ToProtoFullStatus(fullStatus)
|
||||
protoFullStatus.Events = g.statusRecorder.GetEventHistory()
|
||||
overview := nbstatus.ConvertToStatusOutputOverview(protoFullStatus, nbstatus.ConvertOptions{
|
||||
Anonymize: g.anonymize,
|
||||
ProfileName: profName,
|
||||
|
||||
@@ -292,16 +292,18 @@ func (s *serviceViaListener) generateFreePort() (uint16, error) {
|
||||
return customPort, nil
|
||||
}
|
||||
|
||||
probeListener, err := net.ListenUDP("udp4", &net.UDPAddr{})
|
||||
udpAddr := net.UDPAddrFromAddrPort(netip.MustParseAddrPort("0.0.0.0:0"))
|
||||
probeListener, err := net.ListenUDP("udp", udpAddr)
|
||||
if err != nil {
|
||||
log.Debugf("failed to bind random port for DNS: %s", err)
|
||||
return 0, err
|
||||
}
|
||||
|
||||
port := uint16(probeListener.LocalAddr().(*net.UDPAddr).Port)
|
||||
if err = probeListener.Close(); err != nil {
|
||||
addrPort := netip.MustParseAddrPort(probeListener.LocalAddr().String()) // might panic if address is incorrect
|
||||
err = probeListener.Close()
|
||||
if err != nil {
|
||||
log.Debugf("failed to free up DNS port: %s", err)
|
||||
return 0, err
|
||||
}
|
||||
return port, nil
|
||||
return addrPort.Port(), nil
|
||||
}
|
||||
|
||||
@@ -48,6 +48,7 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal/peer"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/guard"
|
||||
icemaker "github.com/netbirdio/netbird/client/internal/peer/ice"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/signaling"
|
||||
"github.com/netbirdio/netbird/client/internal/peerstore"
|
||||
"github.com/netbirdio/netbird/client/internal/portforward"
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
@@ -182,7 +183,7 @@ type EngineServices struct {
|
||||
type Engine struct {
|
||||
// signal is a Signal Service client
|
||||
signal signal.Client
|
||||
signaler *peer.Signaler
|
||||
signaler *signaling.Signaler
|
||||
// mgmClient is a Management Service client
|
||||
mgmClient mgm.Client
|
||||
// peerConns is a map that holds all the peers that are known to this peer
|
||||
@@ -318,7 +319,7 @@ func NewEngine(
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
signal: services.SignalClient,
|
||||
signaler: peer.NewSignaler(services.SignalClient, config.WgPrivateKey),
|
||||
signaler: signaling.NewSignaler(services.SignalClient, config.WgPrivateKey),
|
||||
mgmClient: services.MgmClient,
|
||||
relayManager: services.RelayManager,
|
||||
peerStore: peerstore.NewConnStore(),
|
||||
@@ -2605,14 +2606,13 @@ func (e *Engine) updateForwardRules(rules []*mgmProto.ForwardingRule) ([]firewal
|
||||
|
||||
func (e *Engine) toExcludedLazyPeers(rules []firewallManager.ForwardRule, peers []*mgmProto.RemotePeerConfig) map[string]bool {
|
||||
excludedPeers := make(map[string]bool)
|
||||
|
||||
// Ingress forward targets: inbound forwarded traffic is initiated remotely and
|
||||
// cannot wake a lazy connection, so the peer routing the target must stay
|
||||
// permanently connected. AllowedIPs are already parsed on the peer conn, so
|
||||
// reuse those typed prefixes instead of re-parsing the network map strings.
|
||||
for _, r := range rules {
|
||||
ip := r.TranslatedAddress
|
||||
for _, p := range peers {
|
||||
if e.peerRoutesAddr(p, r.TranslatedAddress) {
|
||||
for _, allowedIP := range p.GetAllowedIps() {
|
||||
if allowedIP != ip.String() {
|
||||
continue
|
||||
}
|
||||
log.Infof("exclude forwarder peer from lazy connection: %s", p.GetWgPubKey())
|
||||
excludedPeers[p.GetWgPubKey()] = true
|
||||
}
|
||||
@@ -2622,27 +2622,6 @@ func (e *Engine) toExcludedLazyPeers(rules []firewallManager.ForwardRule, peers
|
||||
return excludedPeers
|
||||
}
|
||||
|
||||
// peerRoutesAddr reports whether the peer is a router for addr, matched against
|
||||
// the peer's already-parsed AllowedIPs from the store (the same typed value the
|
||||
// lazy manager consumes) rather than re-parsing the network map strings.
|
||||
func (e *Engine) peerRoutesAddr(p *mgmProto.RemotePeerConfig, addr netip.Addr) bool {
|
||||
prefixes, ok := e.peerStore.AllowedIPs(p.GetWgPubKey())
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return prefixesContain(prefixes, addr)
|
||||
}
|
||||
|
||||
// prefixesContain reports whether addr falls within any of the prefixes.
|
||||
func prefixesContain(prefixes []netip.Prefix, addr netip.Addr) bool {
|
||||
for _, prefix := range prefixes {
|
||||
if prefix.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// isChecksEqual checks if two slices of checks are equal.
|
||||
func isChecksEqual(checks1, checks2 []*mgmProto.Checks) bool {
|
||||
normalize := func(checks []*mgmProto.Checks) []string {
|
||||
@@ -2769,7 +2748,7 @@ func createFile(path string) error {
|
||||
return file.Close()
|
||||
}
|
||||
|
||||
func convertToOfferAnswer(msg *sProto.Message) (*peer.OfferAnswer, error) {
|
||||
func convertToOfferAnswer(msg *sProto.Message) (*signaling.OfferAnswer, error) {
|
||||
remoteCred, err := signal.UnMarshalCredential(msg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -2785,9 +2764,9 @@ func convertToOfferAnswer(msg *sProto.Message) (*peer.OfferAnswer, error) {
|
||||
}
|
||||
|
||||
// Handle optional SessionID
|
||||
var sessionID *peer.ICESessionID
|
||||
var sessionID *icemaker.SessionID
|
||||
if sessionBytes := msg.GetBody().GetSessionId(); sessionBytes != nil {
|
||||
if id, err := peer.ICESessionIDFromBytes(sessionBytes); err != nil {
|
||||
if id, err := icemaker.SessionIDFromBytes(sessionBytes); err != nil {
|
||||
log.Warnf("Invalid session ID in message: %v", err)
|
||||
sessionID = nil // Set to nil if conversion fails
|
||||
} else {
|
||||
@@ -2797,8 +2776,8 @@ func convertToOfferAnswer(msg *sProto.Message) (*peer.OfferAnswer, error) {
|
||||
|
||||
relayIP := decodeRelayIP(msg.GetBody().GetRelayServerIP())
|
||||
|
||||
offerAnswer := peer.OfferAnswer{
|
||||
IceCredentials: peer.IceCredentials{
|
||||
offerAnswer := signaling.OfferAnswer{
|
||||
IceCredentials: signaling.IceCredentials{
|
||||
UFrag: remoteCred.UFrag,
|
||||
Pwd: remoteCred.Pwd,
|
||||
},
|
||||
|
||||
@@ -1,87 +0,0 @@
|
||||
package internal
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
firewallManager "github.com/netbirdio/netbird/client/firewall/manager"
|
||||
"github.com/netbirdio/netbird/client/internal/peer"
|
||||
"github.com/netbirdio/netbird/client/internal/peerstore"
|
||||
mgmProto "github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
func TestPrefixesContain(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
prefixes []string
|
||||
addr string
|
||||
want bool
|
||||
}{
|
||||
{name: "own overlay /32 matches", prefixes: []string{"100.110.8.145/32"}, addr: "100.110.8.145", want: true},
|
||||
{name: "addr inside routed subnet", prefixes: []string{"10.121.0.0/16"}, addr: "10.121.208.4", want: true},
|
||||
{name: "addr outside subnet", prefixes: []string{"10.121.0.0/16"}, addr: "10.122.0.1", want: false},
|
||||
{name: "different /32", prefixes: []string{"100.110.8.145/32"}, addr: "100.110.8.146", want: false},
|
||||
{name: "ipv6 /128 matches", prefixes: []string{"fd00::1/128"}, addr: "fd00::1", want: true},
|
||||
{name: "no prefixes", prefixes: nil, addr: "10.121.208.4", want: false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
prefixes := make([]netip.Prefix, 0, len(tt.prefixes))
|
||||
for _, p := range tt.prefixes {
|
||||
prefixes = append(prefixes, netip.MustParsePrefix(p))
|
||||
}
|
||||
require.Equal(t, tt.want, prefixesContain(prefixes, netip.MustParseAddr(tt.addr)))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestToExcludedLazyPeers_ForwardTarget guards a regression: the forward-target
|
||||
// peer (the peer routing a ForwardRule.TranslatedAddress) must be excluded from
|
||||
// lazy connections, matched via the peer's already-parsed AllowedIPs.
|
||||
func TestToExcludedLazyPeers_ForwardTarget(t *testing.T) {
|
||||
const targetPeerKey = "cccccccccccccccccccccccccccccccccccccccccc0="
|
||||
const otherPeerKey = "dddddddddddddddddddddddddddddddddddddddddd0="
|
||||
|
||||
store := peerstore.NewConnStore()
|
||||
store.AddPeerConn(targetPeerKey, newTestConn(t, targetPeerKey, "100.110.8.145/32"))
|
||||
store.AddPeerConn(otherPeerKey, newTestConn(t, otherPeerKey, "100.110.9.10/32"))
|
||||
|
||||
e := &Engine{peerStore: store}
|
||||
|
||||
peers := []*mgmProto.RemotePeerConfig{
|
||||
{WgPubKey: targetPeerKey, AllowedIps: []string{"100.110.8.145/32"}},
|
||||
{WgPubKey: otherPeerKey, AllowedIps: []string{"100.110.9.10/32"}},
|
||||
}
|
||||
rules := []firewallManager.ForwardRule{
|
||||
{TranslatedAddress: netip.MustParseAddr("100.110.8.145")},
|
||||
}
|
||||
|
||||
excluded := e.toExcludedLazyPeers(rules, peers)
|
||||
|
||||
require.True(t, excluded[targetPeerKey], "forward-target peer must be excluded from lazy connections")
|
||||
require.False(t, excluded[otherPeerKey], "non-target peer must not be excluded")
|
||||
require.Len(t, excluded, 1)
|
||||
}
|
||||
|
||||
func TestToExcludedLazyPeers_NoRules(t *testing.T) {
|
||||
e := &Engine{peerStore: peerstore.NewConnStore()}
|
||||
|
||||
peers := []*mgmProto.RemotePeerConfig{
|
||||
{WgPubKey: "peer-a", AllowedIps: []string{"100.110.8.145/32"}},
|
||||
}
|
||||
|
||||
require.Empty(t, e.toExcludedLazyPeers(nil, peers))
|
||||
}
|
||||
|
||||
func newTestConn(t *testing.T, key, allowedIP string) *peer.Conn {
|
||||
t.Helper()
|
||||
conn, err := peer.NewConn(peer.ConnConfig{
|
||||
Key: key,
|
||||
WgConfig: peer.WgConfig{AllowedIps: []netip.Prefix{netip.MustParsePrefix(allowedIP)}},
|
||||
}, peer.ServiceDependencies{})
|
||||
require.NoError(t, err)
|
||||
return conn
|
||||
}
|
||||
@@ -75,14 +75,4 @@ func TestApplySessionDeadline_ThreeState(t *testing.T) {
|
||||
require.True(t, e.statusRecorder.GetSessionExpiresAt().IsZero(),
|
||||
"invalid timestamp must clear the deadline")
|
||||
})
|
||||
|
||||
t.Run("recently expired timestamp stays visible as expired", func(t *testing.T) {
|
||||
e := newEngine()
|
||||
expired := time.Now().Add(-5 * time.Minute).UTC().Truncate(time.Second)
|
||||
|
||||
e.ApplySessionDeadline(timestamppb.New(expired))
|
||||
|
||||
require.True(t, e.statusRecorder.GetSessionExpiresAt().Equal(expired),
|
||||
"recently-expired deadline must stay on the recorder so consumers render it as expired")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,84 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// Ownership is a profile's access policy: the typed owner principals plus the
|
||||
// opt-in shared flag.
|
||||
type Ownership struct {
|
||||
Owners []string
|
||||
Shared bool
|
||||
}
|
||||
|
||||
// GroupResolver resolves a Unix caller's effective group IDs and owner group
|
||||
// names to GIDs. A nil resolver disables group matching.
|
||||
type GroupResolver interface {
|
||||
// CallerGIDs returns the set of group IDs the caller belongs to.
|
||||
CallerGIDs(id Identity) map[uint32]struct{}
|
||||
// GroupNameGID resolves a group name to its GID.
|
||||
GroupNameGID(name string) (uint32, bool)
|
||||
}
|
||||
|
||||
// Authorize reports whether the identity may control a profile with the given
|
||||
// ownership. Privileged callers and shared profiles are always allowed.
|
||||
func Authorize(o Ownership, id Identity, r GroupResolver) bool {
|
||||
if id.IsPrivileged() {
|
||||
return true
|
||||
}
|
||||
if o.Shared {
|
||||
return true
|
||||
}
|
||||
for _, raw := range o.Owners {
|
||||
p, ok := ParsePrincipal(raw)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if principalMatches(p, id, r) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func principalMatches(p Principal, id Identity, r GroupResolver) bool {
|
||||
switch p.Kind {
|
||||
case KindUID:
|
||||
if id.IsWindows() {
|
||||
return false
|
||||
}
|
||||
uid, err := strconv.ParseUint(p.Value, 10, 32)
|
||||
return err == nil && uint32(uid) == id.UID
|
||||
case KindGID:
|
||||
if id.IsWindows() {
|
||||
return false
|
||||
}
|
||||
gid, err := strconv.ParseUint(p.Value, 10, 32)
|
||||
return err == nil && callerHasGID(uint32(gid), id, r)
|
||||
case KindGroup:
|
||||
if id.IsWindows() || r == nil {
|
||||
return false
|
||||
}
|
||||
gid, ok := r.GroupNameGID(p.Value)
|
||||
return ok && callerHasGID(gid, id, r)
|
||||
case KindSID:
|
||||
if !id.IsWindows() {
|
||||
return false
|
||||
}
|
||||
return id.SID == p.Value || slices.Contains(id.Groups, p.Value)
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func callerHasGID(gid uint32, id Identity, r GroupResolver) bool {
|
||||
if id.GID == gid {
|
||||
return true
|
||||
}
|
||||
if r == nil {
|
||||
return false
|
||||
}
|
||||
_, ok := r.CallerGIDs(id)[gid]
|
||||
return ok
|
||||
}
|
||||
@@ -1,22 +0,0 @@
|
||||
//go:build !linux && !darwin && !freebsd && !windows
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"runtime"
|
||||
|
||||
"google.golang.org/grpc/credentials"
|
||||
)
|
||||
|
||||
// NewTransportCredentials returns nil on platforms without a peer-identity
|
||||
// primitive.
|
||||
func NewTransportCredentials() credentials.TransportCredentials {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ConnIdentity is unsupported on platforms without a peer-identity primitive.
|
||||
func ConnIdentity(net.Conn) (Identity, error) {
|
||||
return Identity{}, fmt.Errorf("peer identity not supported on %s", runtime.GOOS)
|
||||
}
|
||||
@@ -1,51 +0,0 @@
|
||||
//go:build linux || darwin || freebsd
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
|
||||
"google.golang.org/grpc/credentials"
|
||||
)
|
||||
|
||||
// NewTransportCredentials returns gRPC transport credentials that extract the
|
||||
// caller's kernel-authenticated identity from a Unix-socket connection and
|
||||
// expose it via IdentityFromContext. Non-nil on platforms with a
|
||||
// peer-credential primitive.
|
||||
func NewTransportCredentials() credentials.TransportCredentials {
|
||||
return unixCreds{}
|
||||
}
|
||||
|
||||
type unixCreds struct{}
|
||||
|
||||
func (unixCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
|
||||
return conn, AuthInfo{}, nil
|
||||
}
|
||||
|
||||
// ConnIdentity extracts the caller's identity from an accepted local IPC
|
||||
// connection. On Unix it reads peer credentials from the socket. It is shared
|
||||
// by the gRPC transport credentials and the JSON gateway (which forwards it).
|
||||
func ConnIdentity(conn net.Conn) (Identity, error) {
|
||||
return PeerIdentity(conn)
|
||||
}
|
||||
|
||||
// ServerHandshake extracts the peer identity and fails closed if it cannot be read.
|
||||
func (unixCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
|
||||
id, err := ConnIdentity(conn)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return conn, AuthInfo{
|
||||
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
|
||||
Identity: id,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (unixCreds) Info() credentials.ProtocolInfo {
|
||||
return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()}
|
||||
}
|
||||
|
||||
func (unixCreds) Clone() credentials.TransportCredentials { return unixCreds{} }
|
||||
|
||||
func (unixCreds) OverrideServerName(string) error { return nil }
|
||||
@@ -1,135 +0,0 @@
|
||||
//go:build windows
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"runtime"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
"google.golang.org/grpc/credentials"
|
||||
)
|
||||
|
||||
var (
|
||||
modadvapi32 = windows.NewLazySystemDLL("advapi32.dll")
|
||||
procImpersonateNamedPipeClient = modadvapi32.NewProc("ImpersonateNamedPipeClient")
|
||||
)
|
||||
|
||||
// DefaultPipeSDDL keeps the daemon control pipe open to any LOCAL caller,
|
||||
// like Unix socket with 0666 permissions.
|
||||
//
|
||||
// D:P protected DACL, no inheritance
|
||||
// (D;;GA;;;NU) deny GENERIC_ALL to NETWORK (remote/SMB)
|
||||
// (A;;GA;;;SY) allow GENERIC_ALL to LocalSystem (the daemon itself)
|
||||
// (A;;GA;;;WD) allow GENERIC_ALL to Everyone (local, per-RPC ACL gates)
|
||||
func DefaultPipeSDDL() string {
|
||||
return "D:P(D;;GA;;;NU)(A;;GA;;;SY)(A;;GA;;;WD)"
|
||||
}
|
||||
|
||||
// NewTransportCredentials returns gRPC transport credentials that derive the
|
||||
// caller's identity from the named-pipe client token.
|
||||
//
|
||||
// This requires the client to dial at SECURITY_IDENTIFICATION (see dialNamedPipe).
|
||||
func NewTransportCredentials() credentials.TransportCredentials {
|
||||
return winpipeCreds{}
|
||||
}
|
||||
|
||||
type winpipeCreds struct{}
|
||||
|
||||
func (winpipeCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
|
||||
return conn, AuthInfo{}, nil
|
||||
}
|
||||
|
||||
// ConnIdentity extracts the caller's identity from an accepted named-pipe
|
||||
// connection by impersonating the pipe client and reading its token. It is
|
||||
// shared by the gRPC transport credentials and the JSON gateway (which forwards
|
||||
// it). Requires the client to have connected at SECURITY_IDENTIFICATION.
|
||||
func ConnIdentity(conn net.Conn) (Identity, error) {
|
||||
// go-winio's pipe connection embeds *win32File, which exposes Fd().
|
||||
fdConn, ok := conn.(interface{ Fd() uintptr })
|
||||
if !ok {
|
||||
return Identity{}, fmt.Errorf("connection %T does not expose a pipe handle", conn)
|
||||
}
|
||||
return pipeClientIdentity(windows.Handle(fdConn.Fd()))
|
||||
}
|
||||
|
||||
// ServerHandshake extracts the connecting client's identity from the pipe. Fails
|
||||
// closed if the handle or token cannot be read.
|
||||
func (winpipeCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
|
||||
id, err := ConnIdentity(conn)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return conn, AuthInfo{
|
||||
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
|
||||
Identity: id,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (winpipeCreds) Info() credentials.ProtocolInfo {
|
||||
return credentials.ProtocolInfo{SecurityProtocol: AuthInfo{}.AuthType()}
|
||||
}
|
||||
|
||||
func (winpipeCreds) Clone() credentials.TransportCredentials { return winpipeCreds{} }
|
||||
|
||||
func (winpipeCreds) OverrideServerName(string) error { return nil }
|
||||
|
||||
// pipeClientIdentity reads the connecting client's user SID, enabled group SIDs,
|
||||
// and elevation by impersonating the pipe client on this thread and reading the
|
||||
// impersonation token.
|
||||
func pipeClientIdentity(handle windows.Handle) (id Identity, err error) {
|
||||
runtime.LockOSThread()
|
||||
defer runtime.UnlockOSThread()
|
||||
|
||||
if err = impersonateNamedPipeClient(handle); err != nil {
|
||||
return Identity{}, fmt.Errorf("impersonate named pipe client: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
// Surface revert error if there are no other errors.
|
||||
revErr := windows.RevertToSelf()
|
||||
if err == nil {
|
||||
err = revErr
|
||||
}
|
||||
}()
|
||||
|
||||
// openAsSelf=true: the token is opened using the daemon's process context
|
||||
// (LocalSystem), not the impersonated client's, so the open always succeeds.
|
||||
var token windows.Token
|
||||
if err = windows.OpenThreadToken(windows.CurrentThread(), windows.TOKEN_QUERY, true, &token); err != nil {
|
||||
return Identity{}, fmt.Errorf("open thread token: %w", err)
|
||||
}
|
||||
defer token.Close()
|
||||
|
||||
tu, err := token.GetTokenUser()
|
||||
if err != nil {
|
||||
return Identity{}, fmt.Errorf("get token user: %w", err)
|
||||
}
|
||||
|
||||
tg, err := token.GetTokenGroups()
|
||||
if err != nil {
|
||||
return Identity{}, fmt.Errorf("get token groups: %w", err)
|
||||
}
|
||||
var groups []string
|
||||
for _, g := range tg.AllGroups() {
|
||||
if g.Attributes&windows.SE_GROUP_ENABLED == 0 || g.Attributes&windows.SE_GROUP_USE_FOR_DENY_ONLY != 0 {
|
||||
continue
|
||||
}
|
||||
groups = append(groups, g.Sid.String())
|
||||
}
|
||||
|
||||
return Identity{
|
||||
SID: tu.User.Sid.String(),
|
||||
Groups: groups,
|
||||
Elevated: token.IsElevated(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func impersonateNamedPipeClient(h windows.Handle) error {
|
||||
r, _, e := procImpersonateNamedPipeClient.Call(uintptr(h))
|
||||
if r == 0 {
|
||||
return e
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -1,77 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
|
||||
"google.golang.org/grpc/metadata"
|
||||
)
|
||||
|
||||
// Metadata keys used by the local JSON gateway to forward the HTTP client's
|
||||
// identity to the daemon.
|
||||
const (
|
||||
mdFwdUID = "x-netbird-fwd-uid" // Unix
|
||||
mdFwdGID = "x-netbird-fwd-gid" // Unix
|
||||
mdFwdSID = "x-netbird-fwd-sid" // Windows user SID
|
||||
mdFwdGroup = "x-netbird-fwd-group" // Windows group SID (repeated)
|
||||
mdFwdElevated = "x-netbird-fwd-elevated" // Windows, "1" if elevated
|
||||
)
|
||||
|
||||
// ForwardIdentityMetadata encodes an identity for the gateway to forward to the
|
||||
// daemon.
|
||||
func ForwardIdentityMetadata(id Identity) metadata.MD {
|
||||
if id.IsWindows() {
|
||||
md := metadata.MD{}
|
||||
md.Set(mdFwdSID, id.SID)
|
||||
if len(id.Groups) > 0 {
|
||||
md.Set(mdFwdGroup, id.Groups...)
|
||||
}
|
||||
if id.Elevated {
|
||||
md.Set(mdFwdElevated, "1")
|
||||
}
|
||||
return md
|
||||
}
|
||||
return metadata.Pairs(
|
||||
mdFwdUID, strconv.FormatUint(uint64(id.UID), 10),
|
||||
mdFwdGID, strconv.FormatUint(uint64(id.GID), 10),
|
||||
)
|
||||
}
|
||||
|
||||
// forwardedIdentity extracts a forwarded identity from incoming gRPC metadata
|
||||
func forwardedIdentity(ctx context.Context) (Identity, bool) {
|
||||
md, ok := metadata.FromIncomingContext(ctx)
|
||||
if !ok {
|
||||
return Identity{}, false
|
||||
}
|
||||
|
||||
if sid := mdFirst(md, mdFwdSID); sid != "" {
|
||||
return Identity{
|
||||
SID: sid,
|
||||
Groups: md.Get(mdFwdGroup),
|
||||
Elevated: mdFirst(md, mdFwdElevated) == "1",
|
||||
}, true
|
||||
}
|
||||
|
||||
uidStr := mdFirst(md, mdFwdUID)
|
||||
if uidStr == "" {
|
||||
return Identity{}, false
|
||||
}
|
||||
uid, err := strconv.ParseUint(uidStr, 10, 32)
|
||||
if err != nil {
|
||||
return Identity{}, false
|
||||
}
|
||||
id := Identity{UID: uint32(uid)}
|
||||
if g := mdFirst(md, mdFwdGID); g != "" {
|
||||
if v, err := strconv.ParseUint(g, 10, 32); err == nil {
|
||||
id.GID = uint32(v)
|
||||
}
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
func mdFirst(md metadata.MD, key string) string {
|
||||
if v := md.Get(key); len(v) > 0 {
|
||||
return v[0]
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -1,41 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"google.golang.org/grpc/metadata"
|
||||
)
|
||||
|
||||
func TestForwardIdentityRoundTrip(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
id Identity
|
||||
}{
|
||||
{"unix uid/gid", Identity{UID: 1000, GID: 1000}},
|
||||
{"windows sid+groups+elevated", Identity{
|
||||
SID: "S-1-5-21-1-2-3-1001",
|
||||
Groups: []string{"S-1-5-32-544", "S-1-1-0"},
|
||||
Elevated: true,
|
||||
}},
|
||||
{"windows sid only", Identity{SID: "S-1-5-21-9"}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ctx := metadata.NewIncomingContext(context.Background(), ForwardIdentityMetadata(tc.id))
|
||||
got, ok := forwardedIdentity(ctx)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, tc.id, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardedIdentity_None(t *testing.T) {
|
||||
_, ok := forwardedIdentity(context.Background())
|
||||
assert.False(t, ok, "no metadata, no forwarded identity")
|
||||
|
||||
ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs("other", "x"))
|
||||
_, ok = forwardedIdentity(ctx)
|
||||
assert.False(t, ok)
|
||||
}
|
||||
@@ -1,89 +0,0 @@
|
||||
// Package ipcauth provides the kernel-authenticated identity of a local IPC
|
||||
// (gRPC) caller and the transport credentials that surface it into the gRPC
|
||||
// context, so the daemon can authorize each RPC by caller identity.
|
||||
//
|
||||
// On Unix the identity is read from the kernel via SO_PEERCRED (Linux) or
|
||||
// LOCAL_PEERCRED (Darwin/FreeBSD). On Windows it is derived from the named-pipe
|
||||
// client token. Platforms without a peer-identity primitive get no credentials
|
||||
// and therefore no enforcement.
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"google.golang.org/grpc/credentials"
|
||||
"google.golang.org/grpc/peer"
|
||||
)
|
||||
|
||||
// sidLocalSystem is the well-known Windows SID for the LocalSystem account.
|
||||
const sidLocalSystem = "S-1-5-18"
|
||||
|
||||
// Identity is the kernel-authenticated identity of a local IPC caller. The zero
|
||||
// value is not a valid identity.
|
||||
type Identity struct {
|
||||
// UID and GID are the caller's Unix user ID and primary group ID.
|
||||
// Zero on Windows, where SID is authoritative instead.
|
||||
UID uint32
|
||||
GID uint32
|
||||
|
||||
// SID is the caller's Windows security identifier (empty on Unix).
|
||||
SID string
|
||||
|
||||
// Groups holds the caller's Windows group SIDs, captured from the client
|
||||
// token at handshake (empty on Unix, where supplementary group membership is
|
||||
// resolved on demand via NSS/getent by the authorizer).
|
||||
Groups []string
|
||||
|
||||
// Elevated reports whether the Windows client token is elevated (run as
|
||||
// administrator). Always false on Unix, where privilege is uid==0.
|
||||
Elevated bool
|
||||
}
|
||||
|
||||
// IsWindows reports whether this identity is a Windows principal (SID-based)
|
||||
// rather than a Unix uid/gid principal.
|
||||
func (i Identity) IsWindows() bool {
|
||||
return i.SID != ""
|
||||
}
|
||||
|
||||
// IsPrivileged reports whether the caller is the platform's administrative
|
||||
// principal.
|
||||
func (i Identity) IsPrivileged() bool {
|
||||
if i.IsWindows() {
|
||||
return i.Elevated || i.SID == sidLocalSystem
|
||||
}
|
||||
return i.UID == 0
|
||||
}
|
||||
|
||||
// String renders the identity for audit logs.
|
||||
func (i Identity) String() string {
|
||||
if i.IsWindows() {
|
||||
return fmt.Sprintf("sid=%s elevated=%t", i.SID, i.Elevated)
|
||||
}
|
||||
return fmt.Sprintf("uid=%d gid=%d", i.UID, i.GID)
|
||||
}
|
||||
|
||||
// AuthInfo carries the peer Identity as a gRPC credentials.AuthInfo so the
|
||||
// interceptor can retrieve it from the request context via IdentityFromContext.
|
||||
type AuthInfo struct {
|
||||
credentials.CommonAuthInfo
|
||||
Identity Identity
|
||||
}
|
||||
|
||||
// AuthType identifies the authentication scheme.
|
||||
func (AuthInfo) AuthType() string { return "netbird-ipc-peercred" }
|
||||
|
||||
// IdentityFromContext extracts the caller's kernel-authenticated identity from
|
||||
// the gRPC peer context. The second return value is false when no IPC transport
|
||||
// credentials were negotiated, callers MUST fail closed in that case.
|
||||
func IdentityFromContext(ctx context.Context) (Identity, bool) {
|
||||
p, ok := peer.FromContext(ctx)
|
||||
if !ok {
|
||||
return Identity{}, false
|
||||
}
|
||||
info, ok := p.AuthInfo.(AuthInfo)
|
||||
if !ok {
|
||||
return Identity{}, false
|
||||
}
|
||||
return info.Identity, true
|
||||
}
|
||||
@@ -1,50 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"google.golang.org/grpc/credentials"
|
||||
"google.golang.org/grpc/peer"
|
||||
)
|
||||
|
||||
func TestIdentityFromContext_NoPeer(t *testing.T) {
|
||||
_, ok := IdentityFromContext(context.Background())
|
||||
assert.False(t, ok, "bare context must report no identity (fail closed)")
|
||||
}
|
||||
|
||||
func TestIdentityFromContext_WrongAuthInfo(t *testing.T) {
|
||||
ctx := peer.NewContext(context.Background(), &peer.Peer{})
|
||||
_, ok := IdentityFromContext(ctx)
|
||||
assert.False(t, ok, "peer without our AuthInfo must report no identity")
|
||||
}
|
||||
|
||||
func TestIdentityFromContext_Present(t *testing.T) {
|
||||
want := Identity{UID: 1000, GID: 1000}
|
||||
ctx := peer.NewContext(context.Background(), &peer.Peer{
|
||||
AuthInfo: AuthInfo{
|
||||
CommonAuthInfo: credentials.CommonAuthInfo{SecurityLevel: credentials.NoSecurity},
|
||||
Identity: want,
|
||||
},
|
||||
})
|
||||
|
||||
got, ok := IdentityFromContext(ctx)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, want, got)
|
||||
}
|
||||
|
||||
func TestIdentity_IsPrivileged(t *testing.T) {
|
||||
// Unix
|
||||
assert.True(t, Identity{UID: 0}.IsPrivileged(), "root is privileged")
|
||||
assert.False(t, Identity{UID: 1000}.IsPrivileged(), "non-root is not privileged")
|
||||
// Windows
|
||||
assert.True(t, Identity{SID: "S-1-5-21-1-2-3-1001", Elevated: true}.IsPrivileged(), "elevated admin is privileged")
|
||||
assert.True(t, Identity{SID: "S-1-5-18"}.IsPrivileged(), "LocalSystem is privileged")
|
||||
assert.False(t, Identity{SID: "S-1-5-21-1-2-3-1001"}.IsPrivileged(), "non-elevated admin is NOT privileged")
|
||||
}
|
||||
|
||||
func TestIdentity_IsWindows(t *testing.T) {
|
||||
assert.True(t, Identity{SID: "S-1-5-18"}.IsWindows())
|
||||
assert.False(t, Identity{UID: 0}.IsWindows())
|
||||
}
|
||||
@@ -1,152 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
// Interceptor enforces per-RPC authorization on the daemon IPC, keyed to
|
||||
// the caller's kernel-authenticated identity. It is safe-by-default:
|
||||
// any RPC without a matching bypass is gated by the active profile's ownership,
|
||||
// and a caller without a readable identity is denied.
|
||||
type Interceptor struct {
|
||||
policy ProfilePolicy
|
||||
resolver GroupResolver
|
||||
// selfUID is the daemon's own effective UID. -1 on Windows.
|
||||
selfUID int
|
||||
}
|
||||
|
||||
// NewInterceptor builds an interceptor over the given policy and group resolver.
|
||||
func NewInterceptor(policy ProfilePolicy, resolver GroupResolver) *Interceptor {
|
||||
return &Interceptor{policy: policy, resolver: resolver, selfUID: os.Geteuid()}
|
||||
}
|
||||
|
||||
// UnaryServerInterceptor authorizes each unary RPC before the handler runs.
|
||||
func (i *Interceptor) UnaryServerInterceptor() grpc.UnaryServerInterceptor {
|
||||
return func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||
if err := i.authorize(ctx, info.FullMethod); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return handler(ctx, req)
|
||||
}
|
||||
}
|
||||
|
||||
// StreamServerInterceptor authorizes each streaming RPC before the handler runs.
|
||||
func (i *Interceptor) StreamServerInterceptor() grpc.StreamServerInterceptor {
|
||||
return func(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
||||
if err := i.authorize(ss.Context(), info.FullMethod); err != nil {
|
||||
return err
|
||||
}
|
||||
return handler(srv, ss)
|
||||
}
|
||||
}
|
||||
|
||||
// authorize decides each RPC from the caller's identity, first match wins:
|
||||
//
|
||||
// 0. no identity DENY
|
||||
// 1. self / privileged / forwarded ALLOW (root, elevated, daemon-self, JSON gateway)
|
||||
// 2. in ownersAuthorizedMethods owner tier (gate on DaemonOwnership)
|
||||
// 3. in handlerAuthorizedMethods profile tier (bypass, handler decides)
|
||||
// 4. everything else default tier (gate on ActiveProfileOwnership)
|
||||
func (i *Interceptor) authorize(ctx context.Context, fullMethod string) error {
|
||||
id, ok := IdentityFromContext(ctx)
|
||||
if !ok {
|
||||
log.Warnf("ipc authz: DENY %s. caller identity unavailable", fullMethod)
|
||||
return status.Error(codes.PermissionDenied, "caller identity could not be verified on the daemon control channel")
|
||||
}
|
||||
|
||||
if i.isSelfOrPrivileged(id) {
|
||||
// The local JSON gateway connects as the daemon itself (self/privileged)
|
||||
// and forwards the real HTTP client's identity. Trust it here, where
|
||||
// the transport peer is already the daemon, then authorize as the
|
||||
// forwarded client. A direct non-privileged caller never reaches this
|
||||
// branch, so it cannot forge the forwarding metadata.
|
||||
fwd, hasFwd := forwardedIdentity(ctx)
|
||||
if !hasFwd {
|
||||
i.auditAllow(id, fullMethod)
|
||||
return nil
|
||||
}
|
||||
log.Infof("ipc authz: gateway-forwarded identity %s", fwd)
|
||||
id = fwd
|
||||
if i.isSelfOrPrivileged(id) {
|
||||
i.auditAllow(id, fullMethod)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// Owner tier: daemon-level RPCs (Add, Down, Status) require a daemon-wide owner.
|
||||
if ownersAuthorizedMethods[fullMethod] {
|
||||
allowed, err := i.authorizeOwnership(id, i.policy.DaemonOwnership, i.policy.ClaimDaemonOwnerIfUnowned)
|
||||
if err != nil {
|
||||
log.Errorf("ipc authz: claim daemon owner for %s: %v", id, err)
|
||||
return status.Error(codes.Internal, "failed to claim daemon ownership")
|
||||
}
|
||||
if allowed {
|
||||
i.auditAllow(id, fullMethod)
|
||||
return nil
|
||||
}
|
||||
log.Warnf("ipc authz: DENY %s for %s. not a daemon owner", fullMethod, id)
|
||||
return status.Errorf(codes.PermissionDenied,
|
||||
"not authorized (caller %s is not a daemon owner). ask an owner or run as root/administrator", id)
|
||||
}
|
||||
|
||||
// Profile tier: the handler self-authorizes against the target profile.
|
||||
if handlerAuthorizedMethods[fullMethod] {
|
||||
i.auditAllow(id, fullMethod)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Default: gated on the active profile's ownership.
|
||||
allowed, err := i.authorizeOwnership(id, i.policy.ActiveProfileOwnership, i.policy.ClaimActiveProfileOwnerIfUnowned)
|
||||
if err != nil {
|
||||
log.Errorf("ipc authz: claim active profile for %s: %v", id, err)
|
||||
return status.Error(codes.Internal, "failed to claim profile ownership")
|
||||
}
|
||||
if allowed {
|
||||
i.auditAllow(id, fullMethod)
|
||||
return nil
|
||||
}
|
||||
|
||||
log.Warnf("ipc authz: DENY %s for %s. active profile owned by another principal", fullMethod, id)
|
||||
return status.Errorf(codes.PermissionDenied,
|
||||
"not authorized to control the active profile (caller %s). ask an owner or run as root/administrator", id)
|
||||
}
|
||||
|
||||
// authorizeOwnership authorizes id against an ownership set, claiming it via
|
||||
// trust-on-first-use when it is unowned and unshared. The claim is atomic, so on
|
||||
// a lost race it re-reads and authorizes normally.
|
||||
func (i *Interceptor) authorizeOwnership(id Identity, get func() Ownership, claim func(Identity) (bool, error)) (bool, error) {
|
||||
o := get()
|
||||
if len(o.Owners) == 0 && !o.Shared {
|
||||
claimed, err := claim(id)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if claimed {
|
||||
return true, nil
|
||||
}
|
||||
o = get()
|
||||
}
|
||||
return Authorize(o, id, i.resolver), nil
|
||||
}
|
||||
|
||||
// isSelfOrPrivileged reports whether the caller is the platform administrator
|
||||
// (root / elevated-admin / LocalSystem) or the daemon's own user.
|
||||
func (i *Interceptor) isSelfOrPrivileged(id Identity) bool {
|
||||
if id.IsPrivileged() {
|
||||
return true
|
||||
}
|
||||
// Daemon-self: only meaningful on Unix (Windows privilege is covered above).
|
||||
return !id.IsWindows() && i.selfUID >= 0 && int(id.UID) == i.selfUID
|
||||
}
|
||||
|
||||
func (i *Interceptor) auditAllow(id Identity, fullMethod string) {
|
||||
if auditMethods[fullMethod] {
|
||||
log.Infof("ipc authz: allow %s for %s", fullMethod, id)
|
||||
}
|
||||
}
|
||||
@@ -1,176 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/grpc/peer"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
type mockPolicy struct {
|
||||
o Ownership // active profile ownership
|
||||
daemon Ownership // daemon-wide ownership
|
||||
claimed bool
|
||||
daemonClaimed bool
|
||||
}
|
||||
|
||||
func (m *mockPolicy) ActiveProfileOwnership() Ownership { return m.o }
|
||||
|
||||
// ClaimActiveProfileOwnerIfUnowned records a claim and marks the profile owned.
|
||||
func (m *mockPolicy) ClaimActiveProfileOwnerIfUnowned(id Identity) (bool, error) {
|
||||
if len(m.o.Owners) == 0 && !m.o.Shared {
|
||||
m.o.Owners = []string{OwnerPrincipalForIdentity(id)}
|
||||
m.claimed = true
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (m *mockPolicy) DaemonOwnership() Ownership { return m.daemon }
|
||||
|
||||
// ClaimDaemonOwnerIfUnowned records a daemon claim and marks the daemon owned.
|
||||
func (m *mockPolicy) ClaimDaemonOwnerIfUnowned(id Identity) (bool, error) {
|
||||
if len(m.daemon.Owners) == 0 && !m.daemon.Shared {
|
||||
m.daemon.Owners = []string{OwnerPrincipalForIdentity(id)}
|
||||
m.daemonClaimed = true
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
type mockResolver struct {
|
||||
gids map[uint32]struct{}
|
||||
names map[string]uint32
|
||||
}
|
||||
|
||||
func (m mockResolver) CallerGIDs(Identity) map[uint32]struct{} { return m.gids }
|
||||
func (m mockResolver) GroupNameGID(n string) (uint32, bool) { g, ok := m.names[n]; return g, ok }
|
||||
|
||||
func ctxWith(id Identity) context.Context {
|
||||
return peer.NewContext(context.Background(), &peer.Peer{AuthInfo: AuthInfo{Identity: id}})
|
||||
}
|
||||
|
||||
const (
|
||||
up = servicePath + "Up"
|
||||
list = servicePath + "ListProfiles"
|
||||
unkwn = servicePath + "SomeFutureMethod"
|
||||
down = servicePath + "Down"
|
||||
statusm = servicePath + "Status"
|
||||
addp = servicePath + "AddProfile"
|
||||
switchp = servicePath + "SwitchProfile"
|
||||
addowner = servicePath + "AddOwner"
|
||||
sharep = servicePath + "ShareProfile"
|
||||
)
|
||||
|
||||
func TestInterceptorAuthorize(t *testing.T) {
|
||||
const selfUID = 4000
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
own Ownership // active profile ownership
|
||||
daemon Ownership // daemon-wide ownership
|
||||
resolver GroupResolver
|
||||
ctx context.Context
|
||||
method string
|
||||
wantErr bool
|
||||
}{
|
||||
// Default gate (active profile ownership).
|
||||
{"no identity denies", Ownership{}, Ownership{}, nil, context.Background(), up, true},
|
||||
{"root allowed", Ownership{}, Ownership{}, nil, ctxWith(Identity{UID: 0}), up, false},
|
||||
{"daemon-self allowed", Ownership{}, Ownership{}, nil, ctxWith(Identity{UID: selfUID}), up, false},
|
||||
{"shared allows any", Ownership{Shared: true}, Ownership{}, nil, ctxWith(Identity{UID: 1234}), up, false},
|
||||
{"uid owner allowed", Ownership{Owners: []string{"uid:1000"}}, Ownership{}, nil, ctxWith(Identity{UID: 1000}), up, false},
|
||||
{"non-owner denied", Ownership{Owners: []string{"uid:1000"}}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 2000}), up, true},
|
||||
{"unknown method gated", Ownership{Owners: []string{"uid:1000"}}, Ownership{}, nil, ctxWith(Identity{UID: 2000}), unkwn, true},
|
||||
{"primary gid owner", Ownership{Owners: []string{"gid:5000"}}, Ownership{}, nil, ctxWith(Identity{UID: 2000, GID: 5000}), up, false},
|
||||
{"group-name owner via resolver", Ownership{Owners: []string{"group:admins"}}, Ownership{},
|
||||
mockResolver{names: map[string]uint32{"admins": 5000}, gids: map[uint32]struct{}{5000: {}}},
|
||||
ctxWith(Identity{UID: 2000, GID: 42}), up, false},
|
||||
{"windows sid owner", Ownership{Owners: []string{"sid:S-1-5-21-9"}}, Ownership{}, nil,
|
||||
ctxWith(Identity{SID: "S-1-5-21-9"}), up, false},
|
||||
{"windows group-sid owner", Ownership{Owners: []string{"sid:S-1-5-32-544"}}, Ownership{}, nil,
|
||||
ctxWith(Identity{SID: "S-1-5-21-1", Groups: []string{"S-1-5-32-544"}}), up, false},
|
||||
{"windows elevated privileged", Ownership{}, Ownership{}, nil,
|
||||
ctxWith(Identity{SID: "S-1-5-21-1", Elevated: true}), up, false},
|
||||
|
||||
// Profile tier (handler self-authorizes, bypass).
|
||||
{"list bypasses gate", Ownership{Owners: []string{"uid:1000"}}, Ownership{}, nil, ctxWith(Identity{UID: 2000}), list, false},
|
||||
{"switch-profile bypasses gate", Ownership{Owners: []string{"uid:1000"}}, Ownership{}, nil, ctxWith(Identity{UID: 2000}), switchp, false},
|
||||
|
||||
// Owner tier (daemon-wide ownership), independent of the active profile.
|
||||
{"down by daemon owner allowed", Ownership{Owners: []string{"uid:9"}}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 1000}), down, false},
|
||||
{"down by non-owner denied", Ownership{}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 2000}), down, true},
|
||||
{"status by non-owner denied", Ownership{}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 2000}), statusm, true},
|
||||
{"add by daemon owner allowed", Ownership{}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 1000}), addp, false},
|
||||
{"add by non-owner denied", Ownership{}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 2000}), addp, true},
|
||||
{"owner-tier TOFU claims unowned daemon", Ownership{}, Ownership{}, nil, ctxWith(Identity{UID: 2000}), down, false},
|
||||
|
||||
// Owner-set mutations gate on daemon ownership, not the active profile.
|
||||
{"add-owner by daemon owner allowed", Ownership{Owners: []string{"uid:9"}}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 1000}), addowner, false},
|
||||
{"add-owner by active-profile owner (non daemon owner) denied", Ownership{Owners: []string{"uid:2000"}}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 2000}), addowner, true},
|
||||
{"share by active-profile owner (non daemon owner) denied", Ownership{Owners: []string{"uid:2000"}}, Ownership{Owners: []string{"uid:1000"}}, nil, ctxWith(Identity{UID: 2000}), sharep, true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
i := &Interceptor{policy: &mockPolicy{o: tt.own, daemon: tt.daemon}, resolver: tt.resolver, selfUID: selfUID}
|
||||
err := i.authorize(tt.ctx, tt.method)
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
assert.Equal(t, codes.PermissionDenied, status.Code(err))
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestInterceptorForwardedIdentity verifies the JSON-gateway trust model: a
|
||||
// self/privileged transport peer (the loopback gateway) may forward a real
|
||||
// client identity, but a non-privileged caller cannot forge it.
|
||||
func TestInterceptorForwardedIdentity(t *testing.T) {
|
||||
const selfUID = 4000
|
||||
owners := Ownership{Owners: []string{"uid:1000"}}
|
||||
|
||||
withFwd := func(peerUID, fwdUID uint32) context.Context {
|
||||
ctx := ctxWith(Identity{UID: peerUID})
|
||||
return metadata.NewIncomingContext(ctx, metadata.Pairs(mdFwdUID, itoa(fwdUID)))
|
||||
}
|
||||
|
||||
// Gateway forwards a non-owner client: denied as that client.
|
||||
i := &Interceptor{policy: &mockPolicy{o: owners}, selfUID: selfUID}
|
||||
assert.Error(t, i.authorize(withFwd(selfUID, 2000), up))
|
||||
|
||||
// Gateway forwards the owner: allowed.
|
||||
assert.NoError(t, i.authorize(withFwd(selfUID, 1000), up))
|
||||
|
||||
// A non-privileged direct caller's forwarded metadata: denied
|
||||
assert.Error(t, i.authorize(withFwd(2000, 1000), up))
|
||||
}
|
||||
|
||||
func itoa(u uint32) string {
|
||||
return strconv.FormatUint(uint64(u), 10)
|
||||
}
|
||||
|
||||
// TestInterceptorTOFU verifies an unowned, non-shared profile is claimed by the
|
||||
// first non-privileged caller, and a different caller is then denied.
|
||||
func TestInterceptorTOFU(t *testing.T) {
|
||||
policy := &mockPolicy{o: Ownership{}} // unowned
|
||||
i := &Interceptor{policy: policy, resolver: nil, selfUID: 4000}
|
||||
|
||||
// First caller (uid 1000) claims via TOFU.
|
||||
err := i.authorize(ctxWith(Identity{UID: 1000}), up)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, policy.claimed, "first caller should claim ownership")
|
||||
assert.Equal(t, []string{"uid:1000"}, policy.o.Owners)
|
||||
|
||||
// A different caller is now denied (profile owned by uid 1000).
|
||||
err = i.authorize(ctxWith(Identity{UID: 2000}), up)
|
||||
assert.Error(t, err)
|
||||
assert.Equal(t, codes.PermissionDenied, status.Code(err))
|
||||
}
|
||||
@@ -1,42 +0,0 @@
|
||||
//go:build darwin || freebsd
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// PeerIdentity reads the kernel-authenticated identity of the process on the
|
||||
// other end of a Unix socket connection via LOCAL_PEERCRED (xucred). xucred
|
||||
// carries the uid and primary group.
|
||||
func PeerIdentity(c net.Conn) (Identity, error) {
|
||||
uc, ok := c.(*net.UnixConn)
|
||||
if !ok {
|
||||
return Identity{}, fmt.Errorf("connection is not a unix socket: %T", c)
|
||||
}
|
||||
raw, err := uc.SyscallConn()
|
||||
if err != nil {
|
||||
return Identity{}, fmt.Errorf("raw conn: %w", err)
|
||||
}
|
||||
|
||||
var cred *unix.Xucred
|
||||
var credErr error
|
||||
if err := raw.Control(func(fd uintptr) {
|
||||
cred, credErr = unix.GetsockoptXucred(int(fd), unix.SOL_LOCAL, unix.LOCAL_PEERCRED)
|
||||
}); err != nil {
|
||||
return Identity{}, fmt.Errorf("getsockopt control: %w", err)
|
||||
}
|
||||
if credErr != nil {
|
||||
return Identity{}, fmt.Errorf("LOCAL_PEERCRED: %w", credErr)
|
||||
}
|
||||
|
||||
id := Identity{UID: cred.Uid}
|
||||
// Groups[0] is the effective (primary) GID; guard against an empty list.
|
||||
if cred.Ngroups > 0 {
|
||||
id.GID = cred.Groups[0]
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
@@ -1,41 +0,0 @@
|
||||
//go:build linux
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// PeerIdentity reads the kernel-authenticated identity of the process on the
|
||||
// other end of a Unix socket connection via SO_PEERCRED. The credentials are
|
||||
// captured by the kernel at connect() time and cannot be spoofed or changed for
|
||||
// the life of the connection.
|
||||
func PeerIdentity(c net.Conn) (Identity, error) {
|
||||
uc, ok := c.(*net.UnixConn)
|
||||
if !ok {
|
||||
return Identity{}, fmt.Errorf("connection is not a unix socket: %T", c)
|
||||
}
|
||||
raw, err := uc.SyscallConn()
|
||||
if err != nil {
|
||||
return Identity{}, fmt.Errorf("raw conn: %w", err)
|
||||
}
|
||||
|
||||
var cred *unix.Ucred
|
||||
var credErr error
|
||||
if err := raw.Control(func(fd uintptr) {
|
||||
cred, credErr = unix.GetsockoptUcred(int(fd), unix.SOL_SOCKET, unix.SO_PEERCRED)
|
||||
}); err != nil {
|
||||
return Identity{}, fmt.Errorf("getsockopt control: %w", err)
|
||||
}
|
||||
if credErr != nil {
|
||||
return Identity{}, fmt.Errorf("SO_PEERCRED: %w", credErr)
|
||||
}
|
||||
|
||||
return Identity{
|
||||
UID: cred.Uid,
|
||||
GID: cred.Gid,
|
||||
}, nil
|
||||
}
|
||||
@@ -1,16 +0,0 @@
|
||||
//go:build !linux && !darwin && !freebsd
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"runtime"
|
||||
)
|
||||
|
||||
// PeerIdentity is unimplemented on platforms without a Unix-socket peer-credential
|
||||
// primitive. Windows derives identity from the named-pipe client token instead
|
||||
// (see the Windows transport credentials), so it never calls this.
|
||||
func PeerIdentity(net.Conn) (Identity, error) {
|
||||
return Identity{}, fmt.Errorf("peer credential check not supported on %s", runtime.GOOS)
|
||||
}
|
||||
@@ -1,114 +0,0 @@
|
||||
//go:build linux || darwin || freebsd
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
"google.golang.org/grpc/health"
|
||||
healthpb "google.golang.org/grpc/health/grpc_health_v1"
|
||||
)
|
||||
|
||||
// TestPeerIdentity_MatchesCurrentProcess connects to a real Unix socket and
|
||||
// verifies the extracted UID/GID match the running process (both ends are us).
|
||||
func TestPeerIdentity_MatchesCurrentProcess(t *testing.T) {
|
||||
sock := filepath.Join(t.TempDir(), "peer.sock")
|
||||
ln, err := net.Listen("unix", sock)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = ln.Close() })
|
||||
|
||||
type result struct {
|
||||
id Identity
|
||||
err error
|
||||
}
|
||||
done := make(chan result, 1)
|
||||
go func() {
|
||||
c, aerr := ln.Accept()
|
||||
if aerr != nil {
|
||||
done <- result{err: aerr}
|
||||
return
|
||||
}
|
||||
defer func() { _ = c.Close() }()
|
||||
id, ierr := PeerIdentity(c)
|
||||
done <- result{id: id, err: ierr}
|
||||
}()
|
||||
|
||||
client, err := net.Dial("unix", sock)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = client.Close() })
|
||||
|
||||
res := <-done
|
||||
require.NoError(t, res.err)
|
||||
assert.Equal(t, uint32(os.Getuid()), res.id.UID, "UID should match current process")
|
||||
assert.Equal(t, uint32(os.Getgid()), res.id.GID, "primary GID should match current process")
|
||||
}
|
||||
|
||||
// TestPeerIdentity_NonUnixConn rejects non-Unix connections (fail closed).
|
||||
func TestPeerIdentity_NonUnixConn(t *testing.T) {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = ln.Close() })
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
c, aerr := ln.Accept()
|
||||
if aerr != nil {
|
||||
done <- aerr
|
||||
return
|
||||
}
|
||||
defer func() { _ = c.Close() }()
|
||||
_, ierr := PeerIdentity(c)
|
||||
done <- ierr
|
||||
}()
|
||||
|
||||
client, err := net.Dial("tcp", ln.Addr().String())
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = client.Close() })
|
||||
|
||||
assert.Error(t, <-done, "PeerIdentity must reject a non-Unix connection")
|
||||
}
|
||||
|
||||
// TestGRPCRoundTrip_ServerCredsClientInsecure proves the transport contract end
|
||||
// to end: a gRPC server using the peercred transport credentials still serves a
|
||||
// plain insecure client (the CLI never changed), and the caller's kernel
|
||||
// identity reaches the handler via IdentityFromContext.
|
||||
func TestGRPCRoundTrip_ServerCredsClientInsecure(t *testing.T) {
|
||||
sock := filepath.Join(t.TempDir(), "rt.sock")
|
||||
ln, err := net.Listen("unix", sock)
|
||||
require.NoError(t, err)
|
||||
|
||||
var gotUID uint32
|
||||
var gotOK bool
|
||||
srv := grpc.NewServer(
|
||||
grpc.Creds(NewTransportCredentials()),
|
||||
grpc.UnaryInterceptor(func(ctx context.Context, req any, _ *grpc.UnaryServerInfo, h grpc.UnaryHandler) (any, error) {
|
||||
id, ok := IdentityFromContext(ctx)
|
||||
gotUID, gotOK = id.UID, ok
|
||||
return h(ctx, req)
|
||||
}),
|
||||
)
|
||||
healthpb.RegisterHealthServer(srv, health.NewServer())
|
||||
go func() { _ = srv.Serve(ln) }()
|
||||
t.Cleanup(srv.Stop)
|
||||
|
||||
conn, err := grpc.NewClient("unix://"+sock, grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
_, err = healthpb.NewHealthClient(conn).Check(ctx, &healthpb.HealthCheckRequest{})
|
||||
require.NoError(t, err, "insecure client must reach the peercred server")
|
||||
|
||||
assert.True(t, gotOK, "handler must see a peer identity")
|
||||
assert.Equal(t, uint32(os.Getuid()), gotUID, "handler must see the caller's UID")
|
||||
}
|
||||
@@ -1,135 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import "sync"
|
||||
|
||||
const servicePath = "/daemon.DaemonService/"
|
||||
|
||||
// ProfilePolicy exposes ownership to the interceptor. The daemon server
|
||||
// implements it. ConfigAdapter bridges the gap because the gRPC server (and its
|
||||
// interceptor) is constructed before the server instance exists.
|
||||
type ProfilePolicy interface {
|
||||
// ActiveProfileOwnership returns the active profile's ownership policy.
|
||||
ActiveProfileOwnership() Ownership
|
||||
|
||||
// ClaimActiveProfileOwnerIfUnowned atomically claims the active profile for
|
||||
// id when it has no owners and is not shared (trust-on-first-use), and
|
||||
// reports whether id is now an owner. A false return means the profile was
|
||||
// already owned or shared or another caller won the claim.
|
||||
ClaimActiveProfileOwnerIfUnowned(id Identity) (bool, error)
|
||||
|
||||
// DaemonOwnership returns the daemon-wide ownership policy that governs the
|
||||
// owner-tier RPCs and the default profile.
|
||||
DaemonOwnership() Ownership
|
||||
|
||||
// ClaimDaemonOwnerIfUnowned atomically claims daemon-wide ownership for id
|
||||
// when the daemon is unowned and not shared (trust-on-first-use).
|
||||
ClaimDaemonOwnerIfUnowned(id Identity) (bool, error)
|
||||
}
|
||||
|
||||
// ownersAuthorizedMethods gate on the daemon-wide owner set (or root): daemon-level
|
||||
// ops independent of any profile. The owner-set mutations (AddOwner, ShareProfile,
|
||||
// ResetOwner) must gate here, not on the active profile, else a per-profile owner
|
||||
// could escalate via `owner add`. ResetOwner also requires root in its handler.
|
||||
var ownersAuthorizedMethods = map[string]bool{
|
||||
servicePath + "AddProfile": true,
|
||||
servicePath + "Down": true,
|
||||
servicePath + "Status": true,
|
||||
servicePath + "AddOwner": true,
|
||||
servicePath + "ShareProfile": true,
|
||||
servicePath + "ResetOwner": true,
|
||||
}
|
||||
|
||||
// handlerAuthorizedMethods bypass the ownership gate (identity still required)
|
||||
// and let the handler authorize. GetActiveProfile bypasses only to return
|
||||
// public metadata any local user may read.
|
||||
var handlerAuthorizedMethods = map[string]bool{
|
||||
servicePath + "ListProfiles": true,
|
||||
servicePath + "RemoveProfile": true,
|
||||
servicePath + "RenameProfile": true,
|
||||
servicePath + "SwitchProfile": true,
|
||||
servicePath + "GetActiveProfile": true,
|
||||
}
|
||||
|
||||
// auditMethods are worth an audit log line. Denials are always logged.
|
||||
var auditMethods = map[string]bool{
|
||||
servicePath + "GetConfig": true,
|
||||
servicePath + "SetConfig": true,
|
||||
servicePath + "Login": true,
|
||||
servicePath + "WaitSSOLogin": true,
|
||||
servicePath + "RequestJWTAuth": true,
|
||||
servicePath + "WaitJWTToken": true,
|
||||
servicePath + "StartCapture": true,
|
||||
servicePath + "StartBundleCapture": true,
|
||||
servicePath + "DebugBundle": true,
|
||||
servicePath + "ExposeService": true,
|
||||
servicePath + "Up": true,
|
||||
servicePath + "Down": true,
|
||||
servicePath + "SelectNetworks": true,
|
||||
servicePath + "DeselectNetworks": true,
|
||||
servicePath + "SwitchProfile": true,
|
||||
servicePath + "TriggerUpdate": true,
|
||||
servicePath + "Logout": true,
|
||||
servicePath + "CleanState": true,
|
||||
servicePath + "DeleteState": true,
|
||||
servicePath + "AddOwner": true,
|
||||
servicePath + "ResetOwner": true,
|
||||
servicePath + "ShareProfile": true,
|
||||
}
|
||||
|
||||
// ConfigAdapter is a ProfilePolicy whose backend is set lazily, once the daemon
|
||||
// server instance is created. Until then it reports an unowned profile
|
||||
// (Ownership zero value), so non-privileged callers are denied.
|
||||
type ConfigAdapter struct {
|
||||
mu sync.RWMutex
|
||||
backend ProfilePolicy
|
||||
}
|
||||
|
||||
// SetBackend installs the real policy. Must be called before serving RPCs.
|
||||
func (a *ConfigAdapter) SetBackend(backend ProfilePolicy) {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
a.backend = backend
|
||||
}
|
||||
|
||||
// ActiveProfileOwnership delegates to the backend, or reports an unowned profile
|
||||
// when no backend is set yet.
|
||||
func (a *ConfigAdapter) ActiveProfileOwnership() Ownership {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
if a.backend == nil {
|
||||
return Ownership{}
|
||||
}
|
||||
return a.backend.ActiveProfileOwnership()
|
||||
}
|
||||
|
||||
// ClaimActiveProfileOwnerIfUnowned delegates to the backend. Before the backend
|
||||
// is set it cannot claim, so it reports not-owned (fail closed).
|
||||
func (a *ConfigAdapter) ClaimActiveProfileOwnerIfUnowned(id Identity) (bool, error) {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
if a.backend == nil {
|
||||
return false, nil
|
||||
}
|
||||
return a.backend.ClaimActiveProfileOwnerIfUnowned(id)
|
||||
}
|
||||
|
||||
// DaemonOwnership delegates to the backend, reporting unowned when none is set.
|
||||
func (a *ConfigAdapter) DaemonOwnership() Ownership {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
if a.backend == nil {
|
||||
return Ownership{}
|
||||
}
|
||||
return a.backend.DaemonOwnership()
|
||||
}
|
||||
|
||||
// ClaimDaemonOwnerIfUnowned delegates to the backend. Before the backend is set
|
||||
// it cannot claim, so it reports not-owned (fail closed).
|
||||
func (a *ConfigAdapter) ClaimDaemonOwnerIfUnowned(id Identity) (bool, error) {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
if a.backend == nil {
|
||||
return false, nil
|
||||
}
|
||||
return a.backend.ClaimDaemonOwnerIfUnowned(id)
|
||||
}
|
||||
@@ -1,54 +0,0 @@
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// PrincipalKind is the type of an owner principal.
|
||||
type PrincipalKind string
|
||||
|
||||
const (
|
||||
KindUID PrincipalKind = "uid" // Unix user ID
|
||||
KindGID PrincipalKind = "gid" // Unix group ID
|
||||
KindGroup PrincipalKind = "group" // Unix group name (NSS-resolved)
|
||||
KindSID PrincipalKind = "sid" // Windows user or group SID
|
||||
)
|
||||
|
||||
// Principal is a parsed owner entry from a profile's Owners list.
|
||||
type Principal struct {
|
||||
Kind PrincipalKind
|
||||
Value string
|
||||
}
|
||||
|
||||
// ParsePrincipal parses a "kind:value" owner string. Returns false for empty
|
||||
// values or unknown kinds so malformed entries are ignored rather than trusted.
|
||||
func ParsePrincipal(s string) (Principal, bool) {
|
||||
kind, value, ok := strings.Cut(s, ":")
|
||||
if !ok || value == "" {
|
||||
return Principal{}, false
|
||||
}
|
||||
switch PrincipalKind(kind) {
|
||||
case KindUID, KindGID, KindGroup, KindSID:
|
||||
return Principal{Kind: PrincipalKind(kind), Value: value}, true
|
||||
default:
|
||||
return Principal{}, false
|
||||
}
|
||||
}
|
||||
|
||||
// UIDPrincipal builds the owner string for a Unix user ID.
|
||||
func UIDPrincipal(uid uint32) string {
|
||||
return string(KindUID) + ":" + strconv.FormatUint(uint64(uid), 10)
|
||||
}
|
||||
|
||||
// SIDPrincipal builds the owner string for a Windows SID.
|
||||
func SIDPrincipal(sid string) string { return string(KindSID) + ":" + sid }
|
||||
|
||||
// OwnerPrincipalForIdentity returns the self-ownership principal for an identity:
|
||||
// the user's UID on Unix, or the user's SID on Windows.
|
||||
func OwnerPrincipalForIdentity(id Identity) string {
|
||||
if id.IsWindows() {
|
||||
return SIDPrincipal(id.SID)
|
||||
}
|
||||
return UIDPrincipal(id.UID)
|
||||
}
|
||||
@@ -1,71 +0,0 @@
|
||||
//go:build !windows
|
||||
|
||||
package ipcauth
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/shell"
|
||||
)
|
||||
|
||||
const groupCacheTTL = 30 * time.Second
|
||||
|
||||
// NewDefaultGroupResolver returns an NSS-aware group resolver, owners resolve
|
||||
// correctly for LDAP/AD users under CGO_ENABLED=0. Results are cached briefly.
|
||||
func NewDefaultGroupResolver() GroupResolver {
|
||||
return &nssResolver{byUID: make(map[uint32]gidCacheEntry)}
|
||||
}
|
||||
|
||||
type gidCacheEntry struct {
|
||||
gids map[uint32]struct{}
|
||||
at time.Time
|
||||
}
|
||||
|
||||
type nssResolver struct {
|
||||
mu sync.Mutex
|
||||
byUID map[uint32]gidCacheEntry
|
||||
}
|
||||
|
||||
func (r *nssResolver) CallerGIDs(id Identity) map[uint32]struct{} {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if e, ok := r.byUID[id.UID]; ok && time.Since(e.at) < groupCacheTTL {
|
||||
return e.gids
|
||||
}
|
||||
gids := resolveGIDs(id.UID)
|
||||
r.byUID[id.UID] = gidCacheEntry{gids: gids, at: time.Now()}
|
||||
return gids
|
||||
}
|
||||
|
||||
func resolveGIDs(uid uint32) map[uint32]struct{} {
|
||||
out := make(map[uint32]struct{})
|
||||
u, err := shell.GetUserFromGetent(strconv.FormatUint(uint64(uid), 10))
|
||||
if err != nil {
|
||||
return out
|
||||
}
|
||||
ids, err := shell.GroupIdsWithFallback(u)
|
||||
if err != nil {
|
||||
return out
|
||||
}
|
||||
for _, s := range ids {
|
||||
if g, err := strconv.ParseUint(s, 10, 32); err == nil {
|
||||
out[uint32(g)] = struct{}{}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (r *nssResolver) GroupNameGID(name string) (uint32, bool) {
|
||||
g, err := shell.LookupGroupWithGetent(name)
|
||||
if err != nil {
|
||||
return 0, false
|
||||
}
|
||||
gid, err := strconv.ParseUint(g.Gid, 10, 32)
|
||||
if err != nil {
|
||||
return 0, false
|
||||
}
|
||||
return uint32(gid), true
|
||||
}
|
||||
@@ -1,10 +0,0 @@
|
||||
//go:build windows
|
||||
|
||||
package ipcauth
|
||||
|
||||
// NewDefaultGroupResolver returns nil on Windows: group authorization uses the
|
||||
// group SIDs carried in the client token (see the Windows transport
|
||||
// credentials).
|
||||
func NewDefaultGroupResolver() GroupResolver {
|
||||
return nil
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,18 +1,5 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
const (
|
||||
// StatusIdle indicate the peer is in disconnected state
|
||||
StatusIdle ConnStatus = iota
|
||||
// StatusConnecting indicate the peer is in connecting state
|
||||
StatusConnecting
|
||||
// StatusConnected indicate the peer is in connected state
|
||||
StatusConnected
|
||||
)
|
||||
|
||||
// connStatusInputs is the primitive-valued snapshot of the state that drives the
|
||||
// tri-state connection classification. Extracted so the decision logic can be unit-tested
|
||||
// without constructing full Worker/Handshaker objects.
|
||||
@@ -21,24 +8,7 @@ type connStatusInputs struct {
|
||||
peerUsesRelay bool // remote peer advertises relay support AND local has relay
|
||||
relayConnected bool // statusRelay reports Connected (independent of whether peer uses relay)
|
||||
remoteSupportsICE bool // remote peer sent ICE credentials
|
||||
iceWorkerCreated bool // local WorkerICE exists (false in force-relay mode)
|
||||
iceWorkerCreated bool // local ICE worker exists (false in force-relay mode)
|
||||
iceStatusConnecting bool // statusICE is anything other than Disconnected
|
||||
iceInProgress bool // a negotiation is currently in flight
|
||||
}
|
||||
|
||||
// ConnStatus describe the status of a peer's connection
|
||||
type ConnStatus int32
|
||||
|
||||
func (s ConnStatus) String() string {
|
||||
switch s {
|
||||
case StatusConnecting:
|
||||
return "Connecting"
|
||||
case StatusConnected:
|
||||
return "Connected"
|
||||
case StatusIdle:
|
||||
return "Idle"
|
||||
default:
|
||||
log.Errorf("unknown status: %d", s)
|
||||
return "INVALID_PEER_CONNECTION_STATUS"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,28 +3,33 @@ package peer
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/dispatcher"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/guard"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/ice"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/metricsstages"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/signaling"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/status"
|
||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
"github.com/netbirdio/netbird/util"
|
||||
)
|
||||
|
||||
var testDispatcher = dispatcher.NewConnectionDispatcher()
|
||||
|
||||
var connConf = ConnConfig{
|
||||
Key: "LLHf3Ma6z6mdLbriAJbqhX7+nM/B71lgw2+91q3LfhU=",
|
||||
LocalKey: "RRHf3Ma6z6mdLbriAJbqhX7+nM/B71lgw2+91q3LfhU=",
|
||||
Timeout: time.Second,
|
||||
LocalWgPort: 51820,
|
||||
WgConfig: WgConfig{
|
||||
AllowedIps: []netip.Prefix{netip.MustParsePrefix("100.64.0.1/32")},
|
||||
},
|
||||
ICEConfig: ice.Config{
|
||||
InterfaceBlackList: nil,
|
||||
},
|
||||
@@ -52,92 +57,37 @@ func TestConn_GetKey(t *testing.T) {
|
||||
swWatcher := guard.NewSRWatcher(nil, nil, nil, connConf.ICEConfig)
|
||||
|
||||
sd := ServiceDependencies{
|
||||
SrWatcher: swWatcher,
|
||||
PeerConnDispatcher: testDispatcher,
|
||||
SrWatcher: swWatcher,
|
||||
}
|
||||
conn, err := NewConn(connConf, sd)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
|
||||
got := conn.GetKey()
|
||||
|
||||
assert.Equal(t, got, connConf.Key, "they should be equal")
|
||||
}
|
||||
|
||||
func TestConn_OnRemoteOffer(t *testing.T) {
|
||||
// TestConn_DiscardMessagesWhenNotOpened: signal messages posted to a not yet
|
||||
// opened connection must be discarded without blocking or panicking.
|
||||
func TestConn_DiscardMessagesWhenNotOpened(t *testing.T) {
|
||||
swWatcher := guard.NewSRWatcher(nil, nil, nil, connConf.ICEConfig)
|
||||
sd := ServiceDependencies{
|
||||
StatusRecorder: NewRecorder("https://mgm"),
|
||||
SrWatcher: swWatcher,
|
||||
PeerConnDispatcher: testDispatcher,
|
||||
StatusRecorder: status.NewRecorder("https://mgm"),
|
||||
SrWatcher: swWatcher,
|
||||
}
|
||||
conn, err := NewConn(connConf, sd)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
|
||||
onNewOfferChan := make(chan struct{})
|
||||
|
||||
conn.handshaker.AddRelayListener(func(remoteOfferAnswer *OfferAnswer) {
|
||||
onNewOfferChan <- struct{}{}
|
||||
})
|
||||
|
||||
conn.OnRemoteOffer(OfferAnswer{
|
||||
IceCredentials: IceCredentials{
|
||||
offerAnswer := signaling.OfferAnswer{
|
||||
IceCredentials: signaling.IceCredentials{
|
||||
UFrag: "test",
|
||||
Pwd: "test",
|
||||
},
|
||||
WgListenPort: 0,
|
||||
Version: "",
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
select {
|
||||
case <-onNewOfferChan:
|
||||
// success
|
||||
case <-ctx.Done():
|
||||
t.Error("expected to receive a new offer notification, but timed out")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConn_OnRemoteAnswer(t *testing.T) {
|
||||
swWatcher := guard.NewSRWatcher(nil, nil, nil, connConf.ICEConfig)
|
||||
sd := ServiceDependencies{
|
||||
StatusRecorder: NewRecorder("https://mgm"),
|
||||
SrWatcher: swWatcher,
|
||||
PeerConnDispatcher: testDispatcher,
|
||||
}
|
||||
conn, err := NewConn(connConf, sd)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
onNewOfferChan := make(chan struct{})
|
||||
|
||||
conn.handshaker.AddRelayListener(func(remoteOfferAnswer *OfferAnswer) {
|
||||
onNewOfferChan <- struct{}{}
|
||||
})
|
||||
|
||||
conn.OnRemoteAnswer(OfferAnswer{
|
||||
IceCredentials: IceCredentials{
|
||||
UFrag: "test",
|
||||
Pwd: "test",
|
||||
},
|
||||
WgListenPort: 0,
|
||||
Version: "",
|
||||
})
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
select {
|
||||
case <-onNewOfferChan:
|
||||
// success
|
||||
case <-ctx.Done():
|
||||
t.Error("expected to receive a new offer notification, but timed out")
|
||||
}
|
||||
conn.OnRemoteOffer(offerAnswer)
|
||||
conn.OnRemoteAnswer(offerAnswer)
|
||||
conn.OnRemoteCandidate(nil, nil)
|
||||
conn.Close(false)
|
||||
}
|
||||
|
||||
func TestConn_presharedKey(t *testing.T) {
|
||||
@@ -320,7 +270,7 @@ func newWGTimeoutTestConn(rosenpassEnabled bool, disconnected *[]string) *Conn {
|
||||
ctx: context.Background(),
|
||||
config: cfg,
|
||||
Log: log.WithField("peer", cfg.Key),
|
||||
metricsStages: &MetricsStages{},
|
||||
metricsStages: &metricsstages.MetricsStages{},
|
||||
}
|
||||
conn.SetOnDisconnected(func(remotePeer string) {
|
||||
*disconnected = append(*disconnected, remotePeer)
|
||||
@@ -339,20 +289,20 @@ func TestConn_onWGDisconnected_EscalatesToRosenpassReset(t *testing.T) {
|
||||
conn := newWGTimeoutTestConn(true, &disconnected)
|
||||
|
||||
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
|
||||
conn.onWGDisconnected(conn.ctx)
|
||||
conn.handleWGTimeout()
|
||||
}
|
||||
assert.Empty(t, disconnected, "escalation must not fire below the threshold")
|
||||
|
||||
conn.onWGDisconnected(conn.ctx)
|
||||
conn.handleWGTimeout()
|
||||
assert.Equal(t, []string{conn.config.WgConfig.RemoteKey}, disconnected,
|
||||
"reaching the threshold must report the peer disconnected once")
|
||||
|
||||
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
|
||||
conn.onWGDisconnected(conn.ctx)
|
||||
conn.handleWGTimeout()
|
||||
}
|
||||
assert.Len(t, disconnected, 1, "escalation must restart counting after firing")
|
||||
|
||||
conn.onWGDisconnected(conn.ctx)
|
||||
conn.handleWGTimeout()
|
||||
assert.Len(t, disconnected, 2, "continued timeouts must escalate again")
|
||||
}
|
||||
|
||||
@@ -364,12 +314,12 @@ func TestConn_onWGDisconnected_CheckSuccessResetsEscalation(t *testing.T) {
|
||||
conn := newWGTimeoutTestConn(true, &disconnected)
|
||||
|
||||
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
|
||||
conn.onWGDisconnected(conn.ctx)
|
||||
conn.handleWGTimeout()
|
||||
}
|
||||
conn.onWGCheckSuccess()
|
||||
conn.handleWGCheckSuccess()
|
||||
|
||||
for i := 0; i < wgTimeoutEscalationThreshold-1; i++ {
|
||||
conn.onWGDisconnected(conn.ctx)
|
||||
conn.handleWGTimeout()
|
||||
}
|
||||
assert.Empty(t, disconnected, "handshake success must reset the timeout count")
|
||||
}
|
||||
@@ -382,7 +332,7 @@ func TestConn_onWGDisconnected_NoEscalationWithoutRosenpass(t *testing.T) {
|
||||
conn := newWGTimeoutTestConn(false, &disconnected)
|
||||
|
||||
for i := 0; i < wgTimeoutEscalationThreshold*3; i++ {
|
||||
conn.onWGDisconnected(conn.ctx)
|
||||
conn.handleWGTimeout()
|
||||
}
|
||||
assert.Empty(t, disconnected, "escalation must be limited to rosenpass connections")
|
||||
}
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
package dispatcher
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/peer/id"
|
||||
)
|
||||
|
||||
type ConnectionListener struct {
|
||||
OnConnected func(peerID id.ConnID)
|
||||
OnDisconnected func(peerID id.ConnID)
|
||||
}
|
||||
|
||||
type ConnectionDispatcher struct {
|
||||
listeners map[*ConnectionListener]struct{}
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewConnectionDispatcher() *ConnectionDispatcher {
|
||||
return &ConnectionDispatcher{
|
||||
listeners: make(map[*ConnectionListener]struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (e *ConnectionDispatcher) AddListener(listener *ConnectionListener) {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
e.listeners[listener] = struct{}{}
|
||||
}
|
||||
|
||||
func (e *ConnectionDispatcher) RemoveListener(listener *ConnectionListener) {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
|
||||
delete(e.listeners, listener)
|
||||
}
|
||||
|
||||
func (e *ConnectionDispatcher) NotifyConnected(peerConnID id.ConnID) {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
for listener := range e.listeners {
|
||||
listener.OnConnected(peerConnID)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *ConnectionDispatcher) NotifyDisconnected(peerConnID id.ConnID) {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
for listener := range e.listeners {
|
||||
listener.OnDisconnected(peerConnID)
|
||||
}
|
||||
}
|
||||
69
client/internal/peer/event.go
Normal file
69
client/internal/peer/event.go
Normal file
@@ -0,0 +1,69 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/pion/ice/v4"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/peer/signaling"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/worker"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
)
|
||||
|
||||
// event is a message processed by the Conn event loop. All mutable Conn state
|
||||
// is owned by that loop; producers deliver events through the mailbox and
|
||||
// never mutate Conn state directly.
|
||||
type event any
|
||||
|
||||
// evClose asks the event loop to tear down the connection. done is closed
|
||||
// once the teardown finished.
|
||||
type evClose struct {
|
||||
signalToRemote bool
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
type evRemoteOffer struct {
|
||||
offer signaling.OfferAnswer
|
||||
}
|
||||
|
||||
type evRemoteAnswer struct {
|
||||
answer signaling.OfferAnswer
|
||||
}
|
||||
|
||||
type evRemoteCandidate struct {
|
||||
candidate ice.Candidate
|
||||
haRoutes route.HAMap
|
||||
}
|
||||
|
||||
type evICEReady struct {
|
||||
priority worker.ConnPriority
|
||||
info worker.ICEConnInfo
|
||||
}
|
||||
|
||||
type evICEDown struct {
|
||||
sessionChanged bool
|
||||
}
|
||||
|
||||
type evRelayReady struct {
|
||||
info worker.RelayConnInfo
|
||||
}
|
||||
|
||||
type evRelayDown struct{}
|
||||
|
||||
// evRelayDialDone reports that the relay dial helper goroutine finished,
|
||||
// successfully or not, so the loop may dispatch a pending offer.
|
||||
type evRelayDialDone struct{}
|
||||
|
||||
type evWGTimeout struct{}
|
||||
|
||||
// evWGHandshake reports the first WireGuard handshake of the current watcher run.
|
||||
type evWGHandshake struct {
|
||||
when time.Time
|
||||
}
|
||||
|
||||
// evWGCheckOK reports a watcher check that observed a fresh handshake,
|
||||
// including handshakes of connections that were already up.
|
||||
type evWGCheckOK struct{}
|
||||
|
||||
// evGuardTick asks the loop to send a new offer to restore connectivity.
|
||||
type evGuardTick struct{}
|
||||
@@ -21,8 +21,6 @@ const (
|
||||
)
|
||||
|
||||
type ICEMonitor struct {
|
||||
ReconnectCh chan struct{}
|
||||
|
||||
iFaceDiscover stdnet.ExternalIFaceDiscover
|
||||
iceConfig icemaker.Config
|
||||
tickerPeriod time.Duration
|
||||
@@ -34,7 +32,6 @@ type ICEMonitor struct {
|
||||
func NewICEMonitor(iFaceDiscover stdnet.ExternalIFaceDiscover, config icemaker.Config, period time.Duration) *ICEMonitor {
|
||||
log.Debugf("prepare ICE monitor with period: %s", period)
|
||||
cm := &ICEMonitor{
|
||||
ReconnectCh: make(chan struct{}, 1),
|
||||
iFaceDiscover: iFaceDiscover,
|
||||
iceConfig: config,
|
||||
tickerPeriod: period,
|
||||
|
||||
@@ -1,246 +0,0 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrSignalIsNotReady = errors.New("signal is not ready")
|
||||
)
|
||||
|
||||
// IceCredentials ICE protocol credentials struct
|
||||
type IceCredentials struct {
|
||||
UFrag string
|
||||
Pwd string
|
||||
}
|
||||
|
||||
// OfferAnswer represents a session establishment offer or answer
|
||||
type OfferAnswer struct {
|
||||
IceCredentials IceCredentials
|
||||
// WgListenPort is a remote WireGuard listen port.
|
||||
// This field is used when establishing a direct WireGuard connection without any proxy.
|
||||
// We can set the remote peer's endpoint with this port.
|
||||
WgListenPort int
|
||||
|
||||
// Version of NetBird Agent
|
||||
Version string
|
||||
// RosenpassPubKey is the Rosenpass public key of the remote peer when receiving this message
|
||||
// This value is the local Rosenpass server public key when sending the message
|
||||
RosenpassPubKey []byte
|
||||
// RosenpassAddr is the Rosenpass server address (IP:port) of the remote peer when receiving this message
|
||||
// This value is the local Rosenpass server address when sending the message
|
||||
RosenpassAddr string
|
||||
|
||||
// relay server address
|
||||
RelaySrvAddress string
|
||||
// RelaySrvIP is the IP the remote peer is connected to on its
|
||||
// relay server. Used as a dial target if DNS for RelaySrvAddress
|
||||
// fails. Zero value if the peer did not advertise an IP.
|
||||
RelaySrvIP netip.Addr
|
||||
// SessionID is the unique identifier of the session, used to discard old messages
|
||||
SessionID *ICESessionID
|
||||
}
|
||||
|
||||
func (o *OfferAnswer) hasICECredentials() bool {
|
||||
return o.IceCredentials.UFrag != "" && o.IceCredentials.Pwd != ""
|
||||
}
|
||||
|
||||
type Handshaker struct {
|
||||
mu sync.Mutex
|
||||
log *log.Entry
|
||||
config ConnConfig
|
||||
signaler *Signaler
|
||||
ice *WorkerICE
|
||||
relay *WorkerRelay
|
||||
metricsStages *MetricsStages
|
||||
// relayListener is not blocking because the listener is using a goroutine to process the messages
|
||||
// and it will only keep the latest message if multiple offers are received in a short time
|
||||
// this is to avoid blocking the handshaker if the listener is doing some heavy processing
|
||||
// and also to avoid processing old offers if multiple offers are received in a short time
|
||||
// the listener will always process the latest offer
|
||||
relayListener *AsyncOfferListener
|
||||
iceListener func(remoteOfferAnswer *OfferAnswer)
|
||||
|
||||
// remoteICESupported tracks whether the remote peer includes ICE credentials in its offers/answers.
|
||||
// When false, the local side skips ICE listener dispatch and suppresses ICE credentials in responses.
|
||||
remoteICESupported atomic.Bool
|
||||
|
||||
// remoteOffersCh is a channel used to wait for remote credentials to proceed with the connection
|
||||
remoteOffersCh chan OfferAnswer
|
||||
// remoteAnswerCh is a channel used to wait for remote credentials answer (confirmation of our offer) to proceed with the connection
|
||||
remoteAnswerCh chan OfferAnswer
|
||||
}
|
||||
|
||||
func NewHandshaker(log *log.Entry, config ConnConfig, signaler *Signaler, ice *WorkerICE, relay *WorkerRelay, metricsStages *MetricsStages) *Handshaker {
|
||||
h := &Handshaker{
|
||||
log: log,
|
||||
config: config,
|
||||
signaler: signaler,
|
||||
ice: ice,
|
||||
relay: relay,
|
||||
metricsStages: metricsStages,
|
||||
remoteOffersCh: make(chan OfferAnswer),
|
||||
remoteAnswerCh: make(chan OfferAnswer),
|
||||
}
|
||||
// assume remote supports ICE until we learn otherwise from received offers
|
||||
h.remoteICESupported.Store(ice != nil)
|
||||
return h
|
||||
}
|
||||
|
||||
func (h *Handshaker) RemoteICESupported() bool {
|
||||
return h.remoteICESupported.Load()
|
||||
}
|
||||
|
||||
func (h *Handshaker) AddRelayListener(offer func(remoteOfferAnswer *OfferAnswer)) {
|
||||
h.relayListener = NewAsyncOfferListener(offer)
|
||||
}
|
||||
|
||||
func (h *Handshaker) AddICEListener(offer func(remoteOfferAnswer *OfferAnswer)) {
|
||||
h.iceListener = offer
|
||||
}
|
||||
|
||||
func (h *Handshaker) Listen(ctx context.Context) {
|
||||
for {
|
||||
select {
|
||||
case remoteOfferAnswer := <-h.remoteOffersCh:
|
||||
h.log.Infof("received offer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials())
|
||||
|
||||
// Record signaling received for reconnection attempts
|
||||
if h.metricsStages != nil {
|
||||
h.metricsStages.RecordSignalingReceived()
|
||||
}
|
||||
|
||||
h.updateRemoteICEState(&remoteOfferAnswer)
|
||||
|
||||
if h.relayListener != nil {
|
||||
h.relayListener.Notify(&remoteOfferAnswer)
|
||||
}
|
||||
|
||||
if h.iceListener != nil && h.RemoteICESupported() {
|
||||
h.iceListener(&remoteOfferAnswer)
|
||||
}
|
||||
|
||||
if err := h.sendAnswer(); err != nil {
|
||||
h.log.Errorf("failed to send remote offer confirmation: %s", err)
|
||||
continue
|
||||
}
|
||||
case remoteOfferAnswer := <-h.remoteAnswerCh:
|
||||
h.log.Infof("received answer, running version %s, remote WireGuard listen port %d, session id: %s, remote ICE supported: %t", remoteOfferAnswer.Version, remoteOfferAnswer.WgListenPort, remoteOfferAnswer.SessionIDString(), remoteOfferAnswer.hasICECredentials())
|
||||
|
||||
// Record signaling received for reconnection attempts
|
||||
if h.metricsStages != nil {
|
||||
h.metricsStages.RecordSignalingReceived()
|
||||
}
|
||||
|
||||
h.updateRemoteICEState(&remoteOfferAnswer)
|
||||
|
||||
if h.relayListener != nil {
|
||||
h.relayListener.Notify(&remoteOfferAnswer)
|
||||
}
|
||||
|
||||
if h.iceListener != nil && h.RemoteICESupported() {
|
||||
h.iceListener(&remoteOfferAnswer)
|
||||
}
|
||||
case <-ctx.Done():
|
||||
h.log.Infof("stop listening for remote offers and answers")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handshaker) SendOffer() error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
return h.sendOffer()
|
||||
}
|
||||
|
||||
// OnRemoteOffer handles an offer from the remote peer and returns true if the message was accepted, false otherwise
|
||||
// doesn't block, discards the message if connection wasn't ready
|
||||
func (h *Handshaker) OnRemoteOffer(offer OfferAnswer) {
|
||||
select {
|
||||
case h.remoteOffersCh <- offer:
|
||||
return
|
||||
default:
|
||||
h.log.Warnf("skipping remote offer message because receiver not ready")
|
||||
// connection might not be ready yet to receive so we ignore the message
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// OnRemoteAnswer handles an offer from the remote peer and returns true if the message was accepted, false otherwise
|
||||
// doesn't block, discards the message if connection wasn't ready
|
||||
func (h *Handshaker) OnRemoteAnswer(answer OfferAnswer) {
|
||||
select {
|
||||
case h.remoteAnswerCh <- answer:
|
||||
return
|
||||
default:
|
||||
// connection might not be ready yet to receive so we ignore the message
|
||||
h.log.Warnf("skipping remote answer message because receiver not ready")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// sendOffer prepares local user credentials and signals them to the remote peer
|
||||
func (h *Handshaker) sendOffer() error {
|
||||
if !h.signaler.Ready() {
|
||||
return ErrSignalIsNotReady
|
||||
}
|
||||
|
||||
offer := h.buildOfferAnswer()
|
||||
h.log.Debugf("sending offer with serial: %s", offer.SessionIDString())
|
||||
|
||||
return h.signaler.SignalOffer(offer, h.config.Key)
|
||||
}
|
||||
|
||||
func (h *Handshaker) sendAnswer() error {
|
||||
answer := h.buildOfferAnswer()
|
||||
h.log.Debugf("sending answer with serial: %s", answer.SessionIDString())
|
||||
|
||||
return h.signaler.SignalAnswer(answer, h.config.Key)
|
||||
}
|
||||
|
||||
func (h *Handshaker) buildOfferAnswer() OfferAnswer {
|
||||
answer := OfferAnswer{
|
||||
WgListenPort: h.config.LocalWgPort,
|
||||
Version: version.NetbirdVersion(),
|
||||
RosenpassPubKey: h.config.RosenpassConfig.PubKey,
|
||||
RosenpassAddr: h.config.RosenpassConfig.Addr,
|
||||
}
|
||||
|
||||
if h.ice != nil && h.RemoteICESupported() {
|
||||
uFrag, pwd := h.ice.GetLocalUserCredentials()
|
||||
sid := h.ice.SessionID()
|
||||
answer.IceCredentials = IceCredentials{uFrag, pwd}
|
||||
answer.SessionID = &sid
|
||||
}
|
||||
|
||||
if addr, ip, err := h.relay.RelayInstanceAddress(); err == nil {
|
||||
answer.RelaySrvAddress = addr
|
||||
answer.RelaySrvIP = ip
|
||||
}
|
||||
|
||||
return answer
|
||||
}
|
||||
|
||||
func (h *Handshaker) updateRemoteICEState(offer *OfferAnswer) {
|
||||
hasICE := offer.hasICECredentials()
|
||||
prev := h.remoteICESupported.Swap(hasICE)
|
||||
if prev != hasICE {
|
||||
if hasICE {
|
||||
h.log.Infof("remote peer started sending ICE credentials")
|
||||
} else {
|
||||
h.log.Infof("remote peer stopped sending ICE credentials")
|
||||
if h.ice != nil {
|
||||
h.ice.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,62 +0,0 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
"sync"
|
||||
)
|
||||
|
||||
type callbackFunc func(remoteOfferAnswer *OfferAnswer)
|
||||
|
||||
func (oa *OfferAnswer) SessionIDString() string {
|
||||
if oa.SessionID == nil {
|
||||
return "unknown"
|
||||
}
|
||||
return oa.SessionID.String()
|
||||
}
|
||||
|
||||
type AsyncOfferListener struct {
|
||||
fn callbackFunc
|
||||
running bool
|
||||
latest *OfferAnswer
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewAsyncOfferListener(fn callbackFunc) *AsyncOfferListener {
|
||||
return &AsyncOfferListener{
|
||||
fn: fn,
|
||||
}
|
||||
}
|
||||
|
||||
func (o *AsyncOfferListener) Notify(remoteOfferAnswer *OfferAnswer) {
|
||||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
|
||||
// Store the latest offer
|
||||
o.latest = remoteOfferAnswer
|
||||
|
||||
// If already running, the running goroutine will pick up this latest value
|
||||
if o.running {
|
||||
return
|
||||
}
|
||||
|
||||
// Start processing
|
||||
o.running = true
|
||||
|
||||
// Process in a goroutine to avoid blocking the caller
|
||||
go func(remoteOfferAnswer *OfferAnswer) {
|
||||
for {
|
||||
o.fn(remoteOfferAnswer)
|
||||
|
||||
o.mu.Lock()
|
||||
if o.latest == nil {
|
||||
// No more work to do
|
||||
o.running = false
|
||||
o.mu.Unlock()
|
||||
return
|
||||
}
|
||||
remoteOfferAnswer = o.latest
|
||||
// Clear the latest to mark it as being processed
|
||||
o.latest = nil
|
||||
o.mu.Unlock()
|
||||
}
|
||||
}(remoteOfferAnswer)
|
||||
}
|
||||
@@ -1,39 +0,0 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func Test_newOfferListener(t *testing.T) {
|
||||
dummyOfferAnswer := &OfferAnswer{}
|
||||
runChan := make(chan struct{}, 10)
|
||||
|
||||
longRunningFn := func(remoteOfferAnswer *OfferAnswer) {
|
||||
time.Sleep(1 * time.Second)
|
||||
runChan <- struct{}{}
|
||||
}
|
||||
|
||||
hl := NewAsyncOfferListener(longRunningFn)
|
||||
|
||||
hl.Notify(dummyOfferAnswer)
|
||||
hl.Notify(dummyOfferAnswer)
|
||||
hl.Notify(dummyOfferAnswer)
|
||||
|
||||
// Wait for exactly 2 callbacks
|
||||
for i := 0; i < 2; i++ {
|
||||
select {
|
||||
case <-runChan:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("Timeout waiting for callback")
|
||||
}
|
||||
}
|
||||
|
||||
// Verify no additional callbacks happen
|
||||
select {
|
||||
case <-runChan:
|
||||
t.Fatal("Unexpected additional callback")
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
t.Log("Correctly received exactly 2 callbacks")
|
||||
}
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package ice
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
@@ -9,26 +9,26 @@ import (
|
||||
|
||||
const sessionIDSize = 5
|
||||
|
||||
type ICESessionID string
|
||||
type SessionID string
|
||||
|
||||
// NewICESessionID generates a new session ID for distinguishing sessions
|
||||
func NewICESessionID() (ICESessionID, error) {
|
||||
// NewSessionID generates a new session ID for distinguishing sessions
|
||||
func NewSessionID() (SessionID, error) {
|
||||
b := make([]byte, sessionIDSize)
|
||||
if _, err := io.ReadFull(rand.Reader, b); err != nil {
|
||||
return "", fmt.Errorf("failed to generate session ID: %w", err)
|
||||
}
|
||||
return ICESessionID(hex.EncodeToString(b)), nil
|
||||
return SessionID(hex.EncodeToString(b)), nil
|
||||
}
|
||||
|
||||
func ICESessionIDFromBytes(b []byte) (ICESessionID, error) {
|
||||
func SessionIDFromBytes(b []byte) (SessionID, error) {
|
||||
if len(b) != sessionIDSize {
|
||||
return "", fmt.Errorf("invalid session ID length: %d", len(b))
|
||||
}
|
||||
return ICESessionID(hex.EncodeToString(b)), nil
|
||||
return SessionID(hex.EncodeToString(b)), nil
|
||||
}
|
||||
|
||||
// Bytes returns the raw bytes of the session ID for protobuf serialization
|
||||
func (id ICESessionID) Bytes() ([]byte, error) {
|
||||
func (id SessionID) Bytes() ([]byte, error) {
|
||||
if len(id) == 0 {
|
||||
return nil, fmt.Errorf("ICE session ID is empty")
|
||||
}
|
||||
@@ -42,6 +42,6 @@ func (id ICESessionID) Bytes() ([]byte, error) {
|
||||
return b, nil
|
||||
}
|
||||
|
||||
func (id ICESessionID) String() string {
|
||||
func (id SessionID) String() string {
|
||||
return string(id)
|
||||
}
|
||||
@@ -1,22 +0,0 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"time"
|
||||
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface/configurer"
|
||||
"github.com/netbirdio/netbird/client/iface/wgaddr"
|
||||
"github.com/netbirdio/netbird/client/iface/wgproxy"
|
||||
)
|
||||
|
||||
type WGIface interface {
|
||||
UpdatePeer(peerKey string, allowedIps []netip.Prefix, keepAlive time.Duration, endpoint *net.UDPAddr, preSharedKey *wgtypes.Key) error
|
||||
RemovePeer(peerKey string) error
|
||||
GetStats() (map[string]configurer.WGStats, error)
|
||||
GetProxy() wgproxy.Proxy
|
||||
Address() wgaddr.Address
|
||||
RemoveEndpointAddress(key string) error
|
||||
}
|
||||
@@ -1,11 +0,0 @@
|
||||
package peer
|
||||
|
||||
// Listener is a callback type about the NetBird network connection state
|
||||
type Listener interface {
|
||||
OnConnected()
|
||||
OnDisconnected()
|
||||
OnConnecting()
|
||||
OnDisconnecting()
|
||||
OnAddressChanged(string, string)
|
||||
OnPeersListChanged(int)
|
||||
}
|
||||
116
client/internal/peer/mailbox.go
Normal file
116
client/internal/peer/mailbox.go
Normal file
@@ -0,0 +1,116 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
"sync"
|
||||
)
|
||||
|
||||
// maxQueuedCandidates bounds the remote candidate queue; on overflow the
|
||||
// oldest candidate is dropped. Lost candidates are recovered by the next
|
||||
// offer exchange triggered by the guard.
|
||||
const maxQueuedCandidates = 128
|
||||
|
||||
// mailbox is the coalescing inbox of the Conn event loop. Posting never
|
||||
// blocks. Per message kind either the latest value wins (offer, answer,
|
||||
// guard tick), the values queue in bounded FIFO order (candidates) or in
|
||||
// unbounded FIFO order (lifecycle and transport state changes, which are
|
||||
// low-volume and must not be lost). A new offer flushes the queued
|
||||
// candidates because they belong to the superseded session.
|
||||
type mailbox struct {
|
||||
mu sync.Mutex
|
||||
closed bool
|
||||
|
||||
lifecycle []event
|
||||
transport []event
|
||||
offer *evRemoteOffer
|
||||
answer *evRemoteAnswer
|
||||
candidates []evRemoteCandidate
|
||||
guardTick bool
|
||||
|
||||
wake chan struct{}
|
||||
}
|
||||
|
||||
func newMailbox() *mailbox {
|
||||
return &mailbox{
|
||||
wake: make(chan struct{}, 1),
|
||||
}
|
||||
}
|
||||
|
||||
// post stores the event and wakes the loop. It reports false if the mailbox
|
||||
// is already closed and the event was not accepted.
|
||||
func (m *mailbox) post(ev event) bool {
|
||||
m.mu.Lock()
|
||||
if m.closed {
|
||||
m.mu.Unlock()
|
||||
return false
|
||||
}
|
||||
|
||||
switch e := ev.(type) {
|
||||
case evClose:
|
||||
m.lifecycle = append(m.lifecycle, e)
|
||||
case evRemoteOffer:
|
||||
m.offer = &e
|
||||
m.candidates = nil
|
||||
case evRemoteAnswer:
|
||||
m.answer = &e
|
||||
case evRemoteCandidate:
|
||||
if len(m.candidates) >= maxQueuedCandidates {
|
||||
m.candidates = m.candidates[1:]
|
||||
}
|
||||
m.candidates = append(m.candidates, e)
|
||||
case evGuardTick:
|
||||
m.guardTick = true
|
||||
default:
|
||||
m.transport = append(m.transport, ev)
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
select {
|
||||
case m.wake <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// drain returns the pending events in processing order: lifecycle first,
|
||||
// then transport state changes, the coalesced offer and answer, the queued
|
||||
// candidates and finally the guard tick.
|
||||
func (m *mailbox) drain() []event {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
return m.drainLocked()
|
||||
}
|
||||
|
||||
// closeAndDrain marks the mailbox closed so further posts are rejected and
|
||||
// returns the events that were still pending.
|
||||
func (m *mailbox) closeAndDrain() []event {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.closed = true
|
||||
return m.drainLocked()
|
||||
}
|
||||
|
||||
func (m *mailbox) drainLocked() []event {
|
||||
evs := make([]event, 0, len(m.lifecycle)+len(m.transport)+len(m.candidates)+3)
|
||||
evs = append(evs, m.lifecycle...)
|
||||
evs = append(evs, m.transport...)
|
||||
if m.offer != nil {
|
||||
evs = append(evs, *m.offer)
|
||||
}
|
||||
if m.answer != nil {
|
||||
evs = append(evs, *m.answer)
|
||||
}
|
||||
for _, c := range m.candidates {
|
||||
evs = append(evs, c)
|
||||
}
|
||||
if m.guardTick {
|
||||
evs = append(evs, evGuardTick{})
|
||||
}
|
||||
|
||||
m.lifecycle = nil
|
||||
m.transport = nil
|
||||
m.offer = nil
|
||||
m.answer = nil
|
||||
m.candidates = nil
|
||||
m.guardTick = false
|
||||
return evs
|
||||
}
|
||||
128
client/internal/peer/mailbox_test.go
Normal file
128
client/internal/peer/mailbox_test.go
Normal file
@@ -0,0 +1,128 @@
|
||||
package peer
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/peer/signaling"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestMailbox_OfferCoalescing(t *testing.T) {
|
||||
mb := newMailbox()
|
||||
|
||||
require.True(t, mb.post(evRemoteOffer{offer: signaling.OfferAnswer{WgListenPort: 1}}))
|
||||
require.True(t, mb.post(evRemoteOffer{offer: signaling.OfferAnswer{WgListenPort: 2}}))
|
||||
require.True(t, mb.post(evRemoteOffer{offer: signaling.OfferAnswer{WgListenPort: 3}}))
|
||||
|
||||
evs := mb.drain()
|
||||
require.Len(t, evs, 1, "consecutive offers must coalesce to a single event")
|
||||
offer, ok := evs[0].(evRemoteOffer)
|
||||
require.True(t, ok, "coalesced event must be an offer")
|
||||
assert.Equal(t, 3, offer.offer.WgListenPort, "the newest offer must win")
|
||||
}
|
||||
|
||||
func TestMailbox_OfferFlushesCandidates(t *testing.T) {
|
||||
mb := newMailbox()
|
||||
|
||||
require.True(t, mb.post(evRemoteCandidate{}))
|
||||
require.True(t, mb.post(evRemoteCandidate{}))
|
||||
require.True(t, mb.post(evRemoteOffer{offer: signaling.OfferAnswer{}}))
|
||||
|
||||
evs := mb.drain()
|
||||
require.Len(t, evs, 1, "candidates of the superseded session must be flushed")
|
||||
_, ok := evs[0].(evRemoteOffer)
|
||||
assert.True(t, ok, "only the offer must remain after the flush")
|
||||
}
|
||||
|
||||
func TestMailbox_CandidatesKeepOrderAfterOffer(t *testing.T) {
|
||||
mb := newMailbox()
|
||||
|
||||
require.True(t, mb.post(evRemoteOffer{offer: signaling.OfferAnswer{}}))
|
||||
require.True(t, mb.post(evRemoteCandidate{haRoutes: nil}))
|
||||
require.True(t, mb.post(evRemoteCandidate{haRoutes: nil}))
|
||||
|
||||
evs := mb.drain()
|
||||
require.Len(t, evs, 3)
|
||||
_, ok := evs[0].(evRemoteOffer)
|
||||
assert.True(t, ok, "offer must be processed before the candidates")
|
||||
for _, ev := range evs[1:] {
|
||||
_, ok := ev.(evRemoteCandidate)
|
||||
assert.True(t, ok, "candidates posted after the offer must survive")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMailbox_CandidateQueueBounded(t *testing.T) {
|
||||
mb := newMailbox()
|
||||
|
||||
for i := 0; i < maxQueuedCandidates+10; i++ {
|
||||
require.True(t, mb.post(evRemoteCandidate{}))
|
||||
}
|
||||
|
||||
evs := mb.drain()
|
||||
assert.Len(t, evs, maxQueuedCandidates, "candidate queue must stay bounded")
|
||||
}
|
||||
|
||||
func TestMailbox_DrainOrder(t *testing.T) {
|
||||
mb := newMailbox()
|
||||
|
||||
require.True(t, mb.post(evGuardTick{}))
|
||||
require.True(t, mb.post(evRemoteAnswer{answer: signaling.OfferAnswer{}}))
|
||||
require.True(t, mb.post(evRemoteOffer{offer: signaling.OfferAnswer{}}))
|
||||
require.True(t, mb.post(evRelayDown{}))
|
||||
require.True(t, mb.post(evICEDown{sessionChanged: true}))
|
||||
require.True(t, mb.post(evClose{}))
|
||||
|
||||
evs := mb.drain()
|
||||
require.Len(t, evs, 6)
|
||||
|
||||
_, ok := evs[0].(evClose)
|
||||
assert.True(t, ok, "lifecycle events must come first")
|
||||
_, ok = evs[1].(evRelayDown)
|
||||
assert.True(t, ok, "transport events must keep FIFO order")
|
||||
_, ok = evs[2].(evICEDown)
|
||||
assert.True(t, ok, "transport events must keep FIFO order")
|
||||
_, ok = evs[3].(evRemoteOffer)
|
||||
assert.True(t, ok, "offer must come after transport events")
|
||||
_, ok = evs[4].(evRemoteAnswer)
|
||||
assert.True(t, ok, "answer must come after the offer")
|
||||
_, ok = evs[5].(evGuardTick)
|
||||
assert.True(t, ok, "guard tick must come last")
|
||||
}
|
||||
|
||||
func TestMailbox_GuardTickCoalesced(t *testing.T) {
|
||||
mb := newMailbox()
|
||||
|
||||
require.True(t, mb.post(evGuardTick{}))
|
||||
require.True(t, mb.post(evGuardTick{}))
|
||||
require.True(t, mb.post(evGuardTick{}))
|
||||
|
||||
evs := mb.drain()
|
||||
assert.Len(t, evs, 1, "guard ticks must coalesce to a single event")
|
||||
}
|
||||
|
||||
func TestMailbox_PostAfterCloseRejected(t *testing.T) {
|
||||
mb := newMailbox()
|
||||
|
||||
require.True(t, mb.post(evRelayDown{}))
|
||||
leftovers := mb.closeAndDrain()
|
||||
assert.Len(t, leftovers, 1, "pending events must be returned on close")
|
||||
|
||||
assert.False(t, mb.post(evRelayDown{}), "posts must be rejected after close")
|
||||
assert.Empty(t, mb.drain(), "no events must remain after close")
|
||||
}
|
||||
|
||||
func TestMailbox_WakeSignal(t *testing.T) {
|
||||
mb := newMailbox()
|
||||
|
||||
require.True(t, mb.post(evRelayDown{}))
|
||||
require.True(t, mb.post(evGuardTick{}))
|
||||
|
||||
select {
|
||||
case <-mb.wake:
|
||||
default:
|
||||
t.Fatal("wake signal must be pending after posts")
|
||||
}
|
||||
|
||||
assert.Len(t, mb.drain(), 2, "a single wake must deliver all pending events")
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package metricsstages
|
||||
|
||||
import (
|
||||
"sync"
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package metricsstages
|
||||
|
||||
import (
|
||||
"testing"
|
||||
189
client/internal/peer/signaling/handshaker.go
Normal file
189
client/internal/peer/signaling/handshaker.go
Normal file
@@ -0,0 +1,189 @@
|
||||
package signaling
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
icemaker "github.com/netbirdio/netbird/client/internal/peer/ice"
|
||||
relayClient "github.com/netbirdio/netbird/shared/relay/client"
|
||||
"github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrSignalIsNotReady = errors.New("signal is not ready")
|
||||
)
|
||||
|
||||
// IceCredentials ICE protocol credentials struct
|
||||
type IceCredentials struct {
|
||||
UFrag string
|
||||
Pwd string
|
||||
}
|
||||
|
||||
// OfferAnswer represents a session establishment offer or answer
|
||||
type OfferAnswer struct {
|
||||
IceCredentials IceCredentials
|
||||
// WgListenPort is a remote WireGuard listen port.
|
||||
// This field is used when establishing a direct WireGuard connection without any proxy.
|
||||
// We can set the remote peer's endpoint with this port.
|
||||
WgListenPort int
|
||||
|
||||
// Version of NetBird Agent
|
||||
Version string
|
||||
// RosenpassPubKey is the Rosenpass public key of the remote peer when receiving this message
|
||||
// This value is the local Rosenpass server public key when sending the message
|
||||
RosenpassPubKey []byte
|
||||
// RosenpassAddr is the Rosenpass server address (IP:port) of the remote peer when receiving this message
|
||||
// This value is the local Rosenpass server address when sending the message
|
||||
RosenpassAddr string
|
||||
|
||||
// relay server address
|
||||
RelaySrvAddress string
|
||||
// RelaySrvIP is the IP the remote peer is connected to on its
|
||||
// relay server. Used as a dial target if DNS for RelaySrvAddress
|
||||
// fails. Zero value if the peer did not advertise an IP.
|
||||
RelaySrvIP netip.Addr
|
||||
// SessionID is the unique identifier of the session, used to discard old messages
|
||||
SessionID *icemaker.SessionID
|
||||
}
|
||||
|
||||
func (o *OfferAnswer) HasICECredentials() bool {
|
||||
return o.IceCredentials.UFrag != "" && o.IceCredentials.Pwd != ""
|
||||
}
|
||||
|
||||
func (o *OfferAnswer) SessionIDString() string {
|
||||
if o.SessionID == nil {
|
||||
return "unknown"
|
||||
}
|
||||
return o.SessionID.String()
|
||||
}
|
||||
|
||||
// Config carries the peer-specific values the Handshaker embeds into offers
|
||||
// and answers.
|
||||
type Config struct {
|
||||
Key string
|
||||
LocalWgPort int
|
||||
RosenpassPubKey []byte
|
||||
RosenpassAddr string
|
||||
}
|
||||
|
||||
// Credentials are the local ICE credentials and session id the Handshaker embeds in offers.
|
||||
type Credentials struct {
|
||||
UFrag string
|
||||
Pwd string
|
||||
SessionID icemaker.SessionID
|
||||
}
|
||||
|
||||
// ICEWorker is the subset of the ICE worker the Handshaker needs to build offers.
|
||||
type ICEWorker interface {
|
||||
Credentials() Credentials
|
||||
Close()
|
||||
}
|
||||
|
||||
// Handshaker keeps the signaling protocol logic: building and sending offers
|
||||
// and answers and tracking whether the remote peer supports ICE. Incoming
|
||||
// message processing is driven by the Conn event loop.
|
||||
type Handshaker struct {
|
||||
mu sync.Mutex
|
||||
log *log.Entry
|
||||
config Config
|
||||
signaler *Signaler
|
||||
ice ICEWorker
|
||||
relayManager *relayClient.Manager
|
||||
|
||||
// remoteICESupported tracks whether the remote peer includes ICE credentials in its offers/answers.
|
||||
// When false, the local side skips ICE dispatch and suppresses ICE credentials in responses.
|
||||
remoteICESupported atomic.Bool
|
||||
}
|
||||
|
||||
func NewHandshaker(log *log.Entry, config Config, signaler *Signaler, ice ICEWorker, relayManager *relayClient.Manager) *Handshaker {
|
||||
h := &Handshaker{
|
||||
log: log,
|
||||
config: config,
|
||||
signaler: signaler,
|
||||
ice: ice,
|
||||
relayManager: relayManager,
|
||||
}
|
||||
// assume remote supports ICE until we learn otherwise from received offers
|
||||
h.remoteICESupported.Store(ice != nil)
|
||||
return h
|
||||
}
|
||||
|
||||
func (h *Handshaker) RemoteICESupported() bool {
|
||||
return h.remoteICESupported.Load()
|
||||
}
|
||||
|
||||
func (h *Handshaker) SendOffer() error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
return h.sendOffer()
|
||||
}
|
||||
|
||||
func (h *Handshaker) SendAnswer() error {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
return h.sendAnswer()
|
||||
}
|
||||
|
||||
// sendOffer prepares local user credentials and signals them to the remote peer
|
||||
func (h *Handshaker) sendOffer() error {
|
||||
if !h.signaler.Ready() {
|
||||
return ErrSignalIsNotReady
|
||||
}
|
||||
|
||||
offer := h.buildOfferAnswer()
|
||||
h.log.Debugf("sending offer with serial: %s", offer.SessionIDString())
|
||||
|
||||
return h.signaler.SignalOffer(offer, h.config.Key)
|
||||
}
|
||||
|
||||
func (h *Handshaker) sendAnswer() error {
|
||||
answer := h.buildOfferAnswer()
|
||||
h.log.Debugf("sending answer with serial: %s", answer.SessionIDString())
|
||||
|
||||
return h.signaler.SignalAnswer(answer, h.config.Key)
|
||||
}
|
||||
|
||||
func (h *Handshaker) buildOfferAnswer() OfferAnswer {
|
||||
answer := OfferAnswer{
|
||||
WgListenPort: h.config.LocalWgPort,
|
||||
Version: version.NetbirdVersion(),
|
||||
RosenpassPubKey: h.config.RosenpassPubKey,
|
||||
RosenpassAddr: h.config.RosenpassAddr,
|
||||
}
|
||||
|
||||
if h.ice != nil && h.RemoteICESupported() {
|
||||
creds := h.ice.Credentials()
|
||||
answer.IceCredentials = IceCredentials{creds.UFrag, creds.Pwd}
|
||||
sid := creds.SessionID
|
||||
answer.SessionID = &sid
|
||||
}
|
||||
|
||||
if addr, ip, err := h.relayManager.RelayInstanceAddress(); err == nil {
|
||||
answer.RelaySrvAddress = addr
|
||||
answer.RelaySrvIP = ip
|
||||
}
|
||||
|
||||
return answer
|
||||
}
|
||||
|
||||
// UpdateRemoteICEState refreshes the remote ICE support flag from a received
|
||||
// offer or answer and closes the ICE worker when the remote peer stopped
|
||||
// sending ICE credentials. Runs on the Conn event loop.
|
||||
func (h *Handshaker) UpdateRemoteICEState(offer *OfferAnswer) {
|
||||
hasICE := offer.HasICECredentials()
|
||||
prev := h.remoteICESupported.Swap(hasICE)
|
||||
if prev != hasICE {
|
||||
if hasICE {
|
||||
h.log.Infof("remote peer started sending ICE credentials")
|
||||
} else {
|
||||
h.log.Infof("remote peer stopped sending ICE credentials")
|
||||
if h.ice != nil {
|
||||
h.ice.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package signaling
|
||||
|
||||
import (
|
||||
"github.com/pion/ice/v4"
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package state_dump
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -6,11 +6,13 @@ import (
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/peer/status"
|
||||
)
|
||||
|
||||
type stateDump struct {
|
||||
type StateDump struct {
|
||||
log *log.Entry
|
||||
status *Status
|
||||
status *status.Recorder
|
||||
key string
|
||||
|
||||
sentOffer int
|
||||
@@ -26,15 +28,15 @@ type stateDump struct {
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func newStateDump(key string, log *log.Entry, statusRecorder *Status) *stateDump {
|
||||
return &stateDump{
|
||||
func NewStateDump(key string, log *log.Entry, statusRecorder *status.Recorder) *StateDump {
|
||||
return &StateDump{
|
||||
log: log,
|
||||
status: statusRecorder,
|
||||
key: key,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stateDump) Start(ctx context.Context) {
|
||||
func (s *StateDump) Start(ctx context.Context) {
|
||||
ticker := time.NewTicker(10 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
|
||||
@@ -48,25 +50,25 @@ func (s *stateDump) Start(ctx context.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stateDump) RemoteOffer() {
|
||||
func (s *StateDump) RemoteOffer() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.remoteOffer++
|
||||
}
|
||||
|
||||
func (s *stateDump) RemoteCandidate() {
|
||||
func (s *StateDump) RemoteCandidate() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.remoteCandidate++
|
||||
}
|
||||
|
||||
func (s *stateDump) SendOffer() {
|
||||
func (s *StateDump) SendOffer() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.sentOffer++
|
||||
}
|
||||
|
||||
func (s *stateDump) dumpState() {
|
||||
func (s *StateDump) dumpState() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
@@ -80,41 +82,41 @@ func (s *stateDump) dumpState() {
|
||||
status, s.sentOffer, s.remoteOffer, s.remoteAnswer, s.remoteCandidate, s.p2pConnected, s.switchToRelay, s.wgCheckSuccess, s.relayConnected, s.localProxies)
|
||||
}
|
||||
|
||||
func (s *stateDump) RemoteAnswer() {
|
||||
func (s *StateDump) RemoteAnswer() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.remoteAnswer++
|
||||
}
|
||||
|
||||
func (s *stateDump) P2PConnected() {
|
||||
func (s *StateDump) P2PConnected() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.p2pConnected++
|
||||
}
|
||||
|
||||
func (s *stateDump) SwitchToRelay() {
|
||||
func (s *StateDump) SwitchToRelay() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.switchToRelay++
|
||||
}
|
||||
|
||||
func (s *stateDump) WGcheckSuccess() {
|
||||
func (s *StateDump) WGcheckSuccess() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.wgCheckSuccess++
|
||||
}
|
||||
|
||||
func (s *stateDump) RelayConnected() {
|
||||
func (s *StateDump) RelayConnected() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.relayConnected++
|
||||
}
|
||||
|
||||
func (s *stateDump) NewLocalProxy() {
|
||||
func (s *StateDump) NewLocalProxy() {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
31
client/internal/peer/status/conn_status.go
Normal file
31
client/internal/peer/status/conn_status.go
Normal file
@@ -0,0 +1,31 @@
|
||||
package status
|
||||
|
||||
import (
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
const (
|
||||
// StatusIdle indicate the peer is in disconnected state
|
||||
StatusIdle ConnStatus = iota
|
||||
// StatusConnecting indicate the peer is in connecting state
|
||||
StatusConnecting
|
||||
// StatusConnected indicate the peer is in connected state
|
||||
StatusConnected
|
||||
)
|
||||
|
||||
// ConnStatus describe the status of a peer's connection
|
||||
type ConnStatus int32
|
||||
|
||||
func (s ConnStatus) String() string {
|
||||
switch s {
|
||||
case StatusConnecting:
|
||||
return "Connecting"
|
||||
case StatusConnected:
|
||||
return "Connected"
|
||||
case StatusIdle:
|
||||
return "Idle"
|
||||
default:
|
||||
log.Errorf("unknown status: %d", s)
|
||||
return "INVALID_PEER_CONNECTION_STATUS"
|
||||
}
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package status
|
||||
|
||||
import (
|
||||
"testing"
|
||||
48
client/internal/peer/status/events.go
Normal file
48
client/internal/peer/status/events.go
Normal file
@@ -0,0 +1,48 @@
|
||||
package status
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"sync"
|
||||
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
)
|
||||
|
||||
type EventQueue struct {
|
||||
maxSize int
|
||||
events []*proto.SystemEvent
|
||||
mutex sync.RWMutex
|
||||
}
|
||||
|
||||
func NewEventQueue(size int) *EventQueue {
|
||||
return &EventQueue{
|
||||
maxSize: size,
|
||||
events: make([]*proto.SystemEvent, 0, size),
|
||||
}
|
||||
}
|
||||
|
||||
func (q *EventQueue) Add(event *proto.SystemEvent) {
|
||||
q.mutex.Lock()
|
||||
defer q.mutex.Unlock()
|
||||
|
||||
q.events = append(q.events, event)
|
||||
|
||||
if len(q.events) > q.maxSize {
|
||||
q.events = q.events[len(q.events)-q.maxSize:]
|
||||
}
|
||||
}
|
||||
|
||||
func (q *EventQueue) GetAll() []*proto.SystemEvent {
|
||||
q.mutex.RLock()
|
||||
defer q.mutex.RUnlock()
|
||||
|
||||
return slices.Clone(q.events)
|
||||
}
|
||||
|
||||
type EventSubscription struct {
|
||||
id string
|
||||
events chan *proto.SystemEvent
|
||||
}
|
||||
|
||||
func (s *EventSubscription) Events() <-chan *proto.SystemEvent {
|
||||
return s.events
|
||||
}
|
||||
122
client/internal/peer/status/full_status.go
Normal file
122
client/internal/peer/status/full_status.go
Normal file
@@ -0,0 +1,122 @@
|
||||
package status
|
||||
|
||||
import (
|
||||
"golang.org/x/exp/maps"
|
||||
"google.golang.org/protobuf/types/known/durationpb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/relay"
|
||||
"github.com/netbirdio/netbird/client/proto"
|
||||
)
|
||||
|
||||
// FullStatus contains the full state held by the Recorder instance
|
||||
type FullStatus struct {
|
||||
Peers []State
|
||||
ManagementState ManagementState
|
||||
SignalState SignalState
|
||||
LocalPeerState LocalPeerState
|
||||
RosenpassState RosenpassState
|
||||
Relays []relay.ProbeResult
|
||||
NSGroupStates []NSGroupState
|
||||
NumOfForwardingRules int
|
||||
LazyConnectionEnabled bool
|
||||
Events []*proto.SystemEvent
|
||||
}
|
||||
|
||||
// ToProto converts FullStatus to proto.FullStatus.
|
||||
func (fs FullStatus) ToProto() *proto.FullStatus {
|
||||
pbFullStatus := proto.FullStatus{
|
||||
ManagementState: &proto.ManagementState{},
|
||||
SignalState: &proto.SignalState{},
|
||||
LocalPeerState: &proto.LocalPeerState{},
|
||||
Peers: []*proto.PeerState{},
|
||||
}
|
||||
|
||||
pbFullStatus.ManagementState.URL = fs.ManagementState.URL
|
||||
pbFullStatus.ManagementState.Connected = fs.ManagementState.Connected
|
||||
if err := fs.ManagementState.Error; err != nil {
|
||||
pbFullStatus.ManagementState.Error = err.Error()
|
||||
}
|
||||
|
||||
pbFullStatus.SignalState.URL = fs.SignalState.URL
|
||||
pbFullStatus.SignalState.Connected = fs.SignalState.Connected
|
||||
if err := fs.SignalState.Error; err != nil {
|
||||
pbFullStatus.SignalState.Error = err.Error()
|
||||
}
|
||||
|
||||
pbFullStatus.LocalPeerState.IP = fs.LocalPeerState.IP
|
||||
pbFullStatus.LocalPeerState.Ipv6 = fs.LocalPeerState.IPv6
|
||||
pbFullStatus.LocalPeerState.PubKey = fs.LocalPeerState.PubKey
|
||||
pbFullStatus.LocalPeerState.KernelInterface = fs.LocalPeerState.KernelInterface
|
||||
pbFullStatus.LocalPeerState.Fqdn = fs.LocalPeerState.FQDN
|
||||
pbFullStatus.LocalPeerState.WgPort = int32(fs.LocalPeerState.WgPort)
|
||||
pbFullStatus.LocalPeerState.RosenpassPermissive = fs.RosenpassState.Permissive
|
||||
pbFullStatus.LocalPeerState.RosenpassEnabled = fs.RosenpassState.Enabled
|
||||
pbFullStatus.NumberOfForwardingRules = int32(fs.NumOfForwardingRules)
|
||||
pbFullStatus.LazyConnectionEnabled = fs.LazyConnectionEnabled
|
||||
|
||||
pbFullStatus.LocalPeerState.Networks = maps.Keys(fs.LocalPeerState.Routes)
|
||||
|
||||
for _, peerState := range fs.Peers {
|
||||
networks := maps.Keys(peerState.GetRoutes())
|
||||
|
||||
pbPeerState := &proto.PeerState{
|
||||
IP: peerState.IP,
|
||||
Ipv6: peerState.IPv6,
|
||||
PubKey: peerState.PubKey,
|
||||
ConnStatus: peerState.ConnStatus.String(),
|
||||
ConnStatusUpdate: timestamppb.New(peerState.ConnStatusUpdate),
|
||||
Relayed: peerState.Relayed,
|
||||
LocalIceCandidateType: peerState.LocalIceCandidateType,
|
||||
RemoteIceCandidateType: peerState.RemoteIceCandidateType,
|
||||
LocalIceCandidateEndpoint: peerState.LocalIceCandidateEndpoint,
|
||||
RemoteIceCandidateEndpoint: peerState.RemoteIceCandidateEndpoint,
|
||||
RelayAddress: peerState.RelayServerAddress,
|
||||
Fqdn: peerState.FQDN,
|
||||
LastWireguardHandshake: timestamppb.New(peerState.LastWireguardHandshake),
|
||||
BytesRx: peerState.BytesRx,
|
||||
BytesTx: peerState.BytesTx,
|
||||
RosenpassEnabled: peerState.RosenpassEnabled,
|
||||
Networks: networks,
|
||||
Latency: durationpb.New(peerState.Latency),
|
||||
SshHostKey: peerState.SSHHostKey,
|
||||
}
|
||||
pbFullStatus.Peers = append(pbFullStatus.Peers, pbPeerState)
|
||||
}
|
||||
|
||||
for _, relayState := range fs.Relays {
|
||||
pbRelayState := &proto.RelayState{
|
||||
URI: relayState.URI,
|
||||
Available: relayState.Err == nil,
|
||||
Transport: relayState.Transport,
|
||||
}
|
||||
if err := relayState.Err; err != nil {
|
||||
pbRelayState.Error = err.Error()
|
||||
}
|
||||
pbFullStatus.Relays = append(pbFullStatus.Relays, pbRelayState)
|
||||
}
|
||||
|
||||
for _, dnsState := range fs.NSGroupStates {
|
||||
var err string
|
||||
if dnsState.Error != nil {
|
||||
err = dnsState.Error.Error()
|
||||
}
|
||||
|
||||
var servers []string
|
||||
for _, server := range dnsState.Servers {
|
||||
servers = append(servers, server.String())
|
||||
}
|
||||
|
||||
pbDnsState := &proto.NSGroupState{
|
||||
Servers: servers,
|
||||
Domains: dnsState.Domains,
|
||||
Enabled: dnsState.Enabled,
|
||||
Error: err,
|
||||
}
|
||||
pbFullStatus.DnsServers = append(pbFullStatus.DnsServers, pbDnsState)
|
||||
}
|
||||
|
||||
pbFullStatus.Events = fs.Events
|
||||
|
||||
return &pbFullStatus
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package status
|
||||
|
||||
import (
|
||||
"sync"
|
||||
@@ -11,6 +11,16 @@ const (
|
||||
stateDisconnecting
|
||||
)
|
||||
|
||||
// Listener is a callback type about the NetBird network connection state
|
||||
type Listener interface {
|
||||
OnConnected()
|
||||
OnDisconnected()
|
||||
OnConnecting()
|
||||
OnDisconnecting()
|
||||
OnAddressChanged(string, string)
|
||||
OnPeersListChanged(int)
|
||||
}
|
||||
|
||||
type notifier struct {
|
||||
serverStateLock sync.Mutex
|
||||
listenersLock sync.Mutex
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package status
|
||||
|
||||
import (
|
||||
"sync"
|
||||
63
client/internal/peer/status/peer_state.go
Normal file
63
client/internal/peer/status/peer_state.go
Normal file
@@ -0,0 +1,63 @@
|
||||
package status
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/exp/maps"
|
||||
)
|
||||
|
||||
// State contains the latest state of a peer
|
||||
type State struct {
|
||||
Mux *sync.RWMutex
|
||||
IP string
|
||||
IPv6 string
|
||||
PubKey string
|
||||
FQDN string
|
||||
ConnStatus ConnStatus
|
||||
ConnStatusUpdate time.Time
|
||||
Relayed bool
|
||||
LocalIceCandidateType string
|
||||
RemoteIceCandidateType string
|
||||
LocalIceCandidateEndpoint string
|
||||
RemoteIceCandidateEndpoint string
|
||||
RelayServerAddress string
|
||||
LastWireguardHandshake time.Time
|
||||
BytesTx int64
|
||||
BytesRx int64
|
||||
Latency time.Duration
|
||||
RosenpassEnabled bool
|
||||
SSHHostKey []byte
|
||||
routes map[string]struct{}
|
||||
}
|
||||
|
||||
// AddRoute add a single route to routes map
|
||||
func (s *State) AddRoute(network string) {
|
||||
s.Mux.Lock()
|
||||
defer s.Mux.Unlock()
|
||||
if s.routes == nil {
|
||||
s.routes = make(map[string]struct{})
|
||||
}
|
||||
s.routes[network] = struct{}{}
|
||||
}
|
||||
|
||||
// SetRoutes set state routes
|
||||
func (s *State) SetRoutes(routes map[string]struct{}) {
|
||||
s.Mux.Lock()
|
||||
defer s.Mux.Unlock()
|
||||
s.routes = routes
|
||||
}
|
||||
|
||||
// DeleteRoute removes a route from the network amp
|
||||
func (s *State) DeleteRoute(network string) {
|
||||
s.Mux.Lock()
|
||||
defer s.Mux.Unlock()
|
||||
delete(s.routes, network)
|
||||
}
|
||||
|
||||
// GetRoutes return routes map
|
||||
func (s *State) GetRoutes() map[string]struct{} {
|
||||
s.Mux.RLock()
|
||||
defer s.Mux.RUnlock()
|
||||
return maps.Clone(s.routes)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package status
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package status
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
36
client/internal/peer/status_alias.go
Normal file
36
client/internal/peer/status_alias.go
Normal file
@@ -0,0 +1,36 @@
|
||||
package peer
|
||||
|
||||
import "github.com/netbirdio/netbird/client/internal/peer/status"
|
||||
|
||||
// Transitional aliases re-exporting the peer status recorder from its own
|
||||
// package. Callers are being migrated to reference the status package
|
||||
// directly; these aliases will be removed once the migration completes.
|
||||
type (
|
||||
Status = status.Recorder
|
||||
State = status.State
|
||||
ConnStatus = status.ConnStatus
|
||||
FullStatus = status.FullStatus
|
||||
RouterState = status.RouterState
|
||||
LocalPeerState = status.LocalPeerState
|
||||
SignalState = status.SignalState
|
||||
ManagementState = status.ManagementState
|
||||
RosenpassState = status.RosenpassState
|
||||
NSGroupState = status.NSGroupState
|
||||
ResolvedDomainInfo = status.ResolvedDomainInfo
|
||||
StatusChangeSubscription = status.StatusChangeSubscription
|
||||
EventQueue = status.EventQueue
|
||||
EventSubscription = status.EventSubscription
|
||||
WGIfaceStatus = status.WGIfaceStatus
|
||||
Listener = status.Listener
|
||||
EventListener = status.EventListener
|
||||
)
|
||||
|
||||
const (
|
||||
StatusIdle = status.StatusIdle
|
||||
StatusConnecting = status.StatusConnecting
|
||||
StatusConnected = status.StatusConnected
|
||||
)
|
||||
|
||||
var (
|
||||
NewRecorder = status.NewRecorder
|
||||
)
|
||||
@@ -1,13 +1,15 @@
|
||||
package peer
|
||||
package wg_watcher
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface/configurer"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/state_dump"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -23,21 +25,21 @@ type WGInterfaceStater interface {
|
||||
GetStats() (map[string]configurer.WGStats, error)
|
||||
}
|
||||
|
||||
// WGWatcher is single-shot: one instance per connection attempt, run once, then discarded.
|
||||
// Lifecycle is owned by Conn under conn.mu, so it keeps no "enabled" state to go stale.
|
||||
type WGWatcher struct {
|
||||
log *log.Entry
|
||||
wgIfaceStater WGInterfaceStater
|
||||
peerKey string
|
||||
stateDump *stateDump
|
||||
stateDump *state_dump.StateDump
|
||||
|
||||
enabled bool
|
||||
muEnabled sync.Mutex
|
||||
// initialHandshake is not thread-safe; never call PrepareInitialHandshake and EnableWgWatcher concurrently.
|
||||
initialHandshake time.Time
|
||||
|
||||
resetCh chan struct{}
|
||||
}
|
||||
|
||||
func NewWGWatcher(log *log.Entry, wgIfaceStater WGInterfaceStater, peerKey string, stateDump *stateDump) *WGWatcher {
|
||||
func NewWGWatcher(log *log.Entry, wgIfaceStater WGInterfaceStater, peerKey string, stateDump *state_dump.StateDump) *WGWatcher {
|
||||
return &WGWatcher{
|
||||
log: log,
|
||||
wgIfaceStater: wgIfaceStater,
|
||||
@@ -47,14 +49,25 @@ func NewWGWatcher(log *log.Entry, wgIfaceStater WGInterfaceStater, peerKey strin
|
||||
}
|
||||
}
|
||||
|
||||
// PrepareInitialHandshake reads the peer's current WireGuard handshake time. It must be
|
||||
// called before the peer is (re)configured on the WireGuard interface, so the captured
|
||||
// baseline reflects the state prior to this connection attempt instead of racing with
|
||||
// that configuration.
|
||||
func (w *WGWatcher) PrepareInitialHandshake() {
|
||||
// PrepareInitialHandshake reserves the watcher and reads the peer's current WireGuard
|
||||
// handshake time. It must be called before the peer is (re)configured on the WireGuard
|
||||
// interface, so the captured baseline reflects the state prior to this connection attempt
|
||||
// instead of racing with that configuration. Returns ok=false if the watcher is already
|
||||
// running, in which case EnableWgWatcher must not be called.
|
||||
func (w *WGWatcher) PrepareInitialHandshake() (ok bool) {
|
||||
w.muEnabled.Lock()
|
||||
if w.enabled {
|
||||
w.muEnabled.Unlock()
|
||||
return false
|
||||
}
|
||||
|
||||
w.log.Debugf("enable WireGuard watcher")
|
||||
w.enabled = true
|
||||
w.muEnabled.Unlock()
|
||||
|
||||
handshake, _ := w.wgState()
|
||||
w.initialHandshake = handshake
|
||||
return true
|
||||
}
|
||||
|
||||
// EnableWgWatcher runs the WireGuard watcher loop using the handshake baseline captured by
|
||||
@@ -64,6 +77,10 @@ func (w *WGWatcher) PrepareInitialHandshake() {
|
||||
// handshake, including the first.
|
||||
func (w *WGWatcher) EnableWgWatcher(ctx context.Context, enabledTime time.Time, onDisconnectedFn func(), onHandshakeSuccessFn func(when time.Time), onCheckSuccessFn func()) {
|
||||
w.periodicHandshakeCheck(ctx, onDisconnectedFn, onHandshakeSuccessFn, onCheckSuccessFn, enabledTime, w.initialHandshake)
|
||||
|
||||
w.muEnabled.Lock()
|
||||
w.enabled = false
|
||||
w.muEnabled.Unlock()
|
||||
}
|
||||
|
||||
// Reset signals the watcher that the WireGuard peer has been reset and a new
|
||||
@@ -89,7 +106,6 @@ func (w *WGWatcher) periodicHandshakeCheck(ctx context.Context, onDisconnectedFn
|
||||
case <-timer.C:
|
||||
handshake, ok := w.handshakeCheck(lastHandshake)
|
||||
if !ok {
|
||||
// early ctx cancel check return
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
@@ -138,9 +154,9 @@ func (w *WGWatcher) handshakeCheck(lastHandshake time.Time) (*time.Time, bool) {
|
||||
|
||||
w.log.Tracef("previous handshake, handshake: %v, %v", lastHandshake, handshake)
|
||||
|
||||
// the current known handshake did not change
|
||||
// the current know handshake did not change
|
||||
if handshake.Equal(lastHandshake) {
|
||||
w.log.Warnf("WireGuard handshake not updated: %v", handshake)
|
||||
w.log.Warnf("WireGuard handshake timed out: %v", handshake)
|
||||
return nil, false
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package wg_watcher
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -7,8 +7,11 @@ import (
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface/configurer"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/state_dump"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/status"
|
||||
)
|
||||
|
||||
type MocWgIface struct {
|
||||
@@ -56,23 +59,27 @@ func TestWGWatcher_CheckSuccessCallback(t *testing.T) {
|
||||
// platforms with coarse clock resolution (Windows), where two time.Now() calls
|
||||
// microseconds apart can return the same instant and read as a timed-out handshake.
|
||||
stats := &mockHandshakeStats{handshake: time.Now().Add(-time.Hour)}
|
||||
watcher := NewWGWatcher(mlog, stats, "", newStateDump("peer", mlog, &Status{}))
|
||||
watcher := NewWGWatcher(mlog, stats, "", state_dump.NewStateDump("peer", mlog, &status.Recorder{}))
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
watcher.PrepareInitialHandshake()
|
||||
require.True(t, watcher.PrepareInitialHandshake())
|
||||
|
||||
firstHandshake := make(chan struct{}, 1)
|
||||
checkSuccess := make(chan struct{}, 1)
|
||||
go watcher.EnableWgWatcher(ctx, time.Now(), func() {}, func(when time.Time) {
|
||||
firstHandshake <- struct{}{}
|
||||
}, func() {
|
||||
select {
|
||||
case checkSuccess <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
})
|
||||
watcherDone := make(chan struct{})
|
||||
go func() {
|
||||
defer close(watcherDone)
|
||||
watcher.EnableWgWatcher(ctx, time.Now(), func() {}, func(when time.Time) {
|
||||
firstHandshake <- struct{}{}
|
||||
}, func() {
|
||||
select {
|
||||
case checkSuccess <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
})
|
||||
}()
|
||||
|
||||
stats.advance()
|
||||
|
||||
@@ -87,6 +94,11 @@ func TestWGWatcher_CheckSuccessCallback(t *testing.T) {
|
||||
t.Errorf("first-handshake callback must not fire for a non-zero baseline")
|
||||
default:
|
||||
}
|
||||
|
||||
// Wait for the watcher goroutine to exit so it cannot race with other
|
||||
// tests mutating the package-level check timing variables.
|
||||
cancel()
|
||||
<-watcherDone
|
||||
}
|
||||
|
||||
func TestWGWatcher_EnableWgWatcher(t *testing.T) {
|
||||
@@ -95,12 +107,13 @@ func TestWGWatcher_EnableWgWatcher(t *testing.T) {
|
||||
|
||||
mlog := log.WithField("peer", "tet")
|
||||
mocWgIface := &MocWgIface{}
|
||||
watcher := NewWGWatcher(mlog, mocWgIface, "", newStateDump("peer", mlog, &Status{}))
|
||||
watcher := NewWGWatcher(mlog, mocWgIface, "", state_dump.NewStateDump("peer", mlog, &status.Recorder{}))
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
watcher.PrepareInitialHandshake()
|
||||
ok := watcher.PrepareInitialHandshake()
|
||||
require.True(t, ok, "watcher should not be enabled yet")
|
||||
|
||||
onDisconnected := make(chan struct{}, 1)
|
||||
go watcher.EnableWgWatcher(ctx, time.Now(), func() {
|
||||
@@ -127,10 +140,11 @@ func TestWGWatcher_ReEnable(t *testing.T) {
|
||||
|
||||
mlog := log.WithField("peer", "tet")
|
||||
mocWgIface := &MocWgIface{}
|
||||
watcher := NewWGWatcher(mlog, mocWgIface, "", newStateDump("peer", mlog, &Status{}))
|
||||
watcher := NewWGWatcher(mlog, mocWgIface, "", state_dump.NewStateDump("peer", mlog, &status.Recorder{}))
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
watcher.PrepareInitialHandshake()
|
||||
ok := watcher.PrepareInitialHandshake()
|
||||
require.True(t, ok, "watcher should not be enabled yet")
|
||||
|
||||
wg := &sync.WaitGroup{}
|
||||
wg.Add(1)
|
||||
@@ -146,7 +160,8 @@ func TestWGWatcher_ReEnable(t *testing.T) {
|
||||
ctx, cancel = context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
watcher.PrepareInitialHandshake()
|
||||
ok = watcher.PrepareInitialHandshake()
|
||||
require.True(t, ok, "watcher should be re-enabled after the previous run stopped")
|
||||
|
||||
onDisconnected := make(chan struct{}, 1)
|
||||
go watcher.EnableWgWatcher(ctx, time.Now(), func() {
|
||||
@@ -1,4 +1,4 @@
|
||||
package conntype
|
||||
package worker
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package worker
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -13,8 +13,9 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface"
|
||||
"github.com/netbirdio/netbird/client/iface/udpmux"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/conntype"
|
||||
icemaker "github.com/netbirdio/netbird/client/internal/peer/ice"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/signaling"
|
||||
"github.com/netbirdio/netbird/client/internal/peer/status"
|
||||
"github.com/netbirdio/netbird/client/internal/portforward"
|
||||
"github.com/netbirdio/netbird/client/internal/stdnet"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
@@ -32,57 +33,68 @@ type ICEConnInfo struct {
|
||||
RelayedOnLocal bool
|
||||
}
|
||||
|
||||
type WorkerICE struct {
|
||||
ctx context.Context
|
||||
log *log.Entry
|
||||
config ConnConfig
|
||||
conn *Conn
|
||||
signaler *Signaler
|
||||
iFaceDiscover stdnet.ExternalIFaceDiscover
|
||||
statusRecorder *Status
|
||||
hasRelayOnLocally bool
|
||||
type ICEDependencies struct {
|
||||
Signaler *signaling.Signaler
|
||||
IFaceDiscover stdnet.ExternalIFaceDiscover
|
||||
StatusRecorder *status.Recorder
|
||||
PortForwardManager *portforward.Manager
|
||||
}
|
||||
|
||||
type ICE struct {
|
||||
log *log.Entry
|
||||
key string
|
||||
iceConfig icemaker.Config
|
||||
isController bool
|
||||
onConnReady func(priority ConnPriority, iceConnInfo ICEConnInfo)
|
||||
onStatusDisconnect func(sessionChanged bool)
|
||||
signaler *signaling.Signaler
|
||||
iFaceDiscover stdnet.ExternalIFaceDiscover
|
||||
statusRecorder *status.Recorder
|
||||
portForwardManager *portforward.Manager
|
||||
hasRelayOnLocally bool
|
||||
|
||||
agent *icemaker.ThreadSafeAgent
|
||||
agentDialerCancel context.CancelFunc
|
||||
agentConnecting bool // while it is true, drop all incoming offers
|
||||
lastSuccess time.Time // with this avoid the too frequent ICE agent recreation
|
||||
// connectedAgent is the agent whose connection was last reported ready; guarded by muxAgent
|
||||
connectedAgent *icemaker.ThreadSafeAgent
|
||||
// remoteSessionID represents the peer's session identifier from the latest remote offer.
|
||||
remoteSessionID ICESessionID
|
||||
remoteSessionID icemaker.SessionID
|
||||
// sessionID is used to track the current session ID of the ICE agent
|
||||
// increase by one when disconnecting the agent
|
||||
// with it the remote peer can discard the already deprecated offer/answer
|
||||
// Without it the remote peer may recreate a workable ICE connection
|
||||
sessionID ICESessionID
|
||||
sessionID icemaker.SessionID
|
||||
remoteSessionChanged bool
|
||||
muxAgent sync.Mutex
|
||||
|
||||
localUfrag string
|
||||
localPwd string
|
||||
|
||||
// we record the last known state of the ICE agent to avoid duplicate on disconnected events
|
||||
lastKnownState ice.ConnectionState
|
||||
|
||||
// portForwardAttempted tracks if we've already tried port forwarding this session
|
||||
portForwardAttempted bool
|
||||
}
|
||||
|
||||
func NewWorkerICE(ctx context.Context, log *log.Entry, config ConnConfig, conn *Conn, signaler *Signaler, ifaceDiscover stdnet.ExternalIFaceDiscover, statusRecorder *Status, hasRelayOnLocally bool) (*WorkerICE, error) {
|
||||
sessionID, err := NewICESessionID()
|
||||
func NewICE(log *log.Entry, key string, iceConfig icemaker.Config, isController bool, onConnReady func(ConnPriority, ICEConnInfo), onStatusDisconnect func(bool), services ICEDependencies, hasRelayOnLocally bool) (*ICE, error) {
|
||||
sessionID, err := icemaker.NewSessionID()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
w := &WorkerICE{
|
||||
ctx: ctx,
|
||||
log: log,
|
||||
config: config,
|
||||
conn: conn,
|
||||
signaler: signaler,
|
||||
iFaceDiscover: ifaceDiscover,
|
||||
statusRecorder: statusRecorder,
|
||||
hasRelayOnLocally: hasRelayOnLocally,
|
||||
lastKnownState: ice.ConnectionStateDisconnected,
|
||||
sessionID: sessionID,
|
||||
w := &ICE{
|
||||
log: log,
|
||||
key: key,
|
||||
iceConfig: iceConfig,
|
||||
isController: isController,
|
||||
onConnReady: onConnReady,
|
||||
onStatusDisconnect: onStatusDisconnect,
|
||||
signaler: services.Signaler,
|
||||
iFaceDiscover: services.IFaceDiscover,
|
||||
statusRecorder: services.StatusRecorder,
|
||||
portForwardManager: services.PortForwardManager,
|
||||
hasRelayOnLocally: hasRelayOnLocally,
|
||||
sessionID: sessionID,
|
||||
}
|
||||
|
||||
localUfrag, localPwd, err := icemaker.GenerateICECredentials()
|
||||
@@ -94,7 +106,7 @@ func NewWorkerICE(ctx context.Context, log *log.Entry, config ConnConfig, conn *
|
||||
return w, nil
|
||||
}
|
||||
|
||||
func (w *WorkerICE) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
|
||||
func (w *ICE) OnNewOffer(ctx context.Context, remoteOfferAnswer *signaling.OfferAnswer) {
|
||||
w.log.Debugf("OnNewOffer for ICE, serial: %s", remoteOfferAnswer.SessionIDString())
|
||||
w.muxAgent.Lock()
|
||||
defer w.muxAgent.Unlock()
|
||||
@@ -118,7 +130,7 @@ func (w *WorkerICE) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
|
||||
}
|
||||
}
|
||||
|
||||
sessionID, err := NewICESessionID()
|
||||
sessionID, err := icemaker.NewSessionID()
|
||||
if err != nil {
|
||||
w.log.Errorf("failed to create new session ID: %s", err)
|
||||
}
|
||||
@@ -136,8 +148,8 @@ func (w *WorkerICE) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
|
||||
if remoteOfferAnswer.SessionID != nil {
|
||||
w.log.Debugf("recreate ICE agent: %s / %s", w.sessionID, *remoteOfferAnswer.SessionID)
|
||||
}
|
||||
dialerCtx, dialerCancel := context.WithCancel(w.ctx)
|
||||
agent, err := w.reCreateAgent(dialerCancel, preferredCandidateTypes)
|
||||
dialerCtx, dialerCancel := context.WithCancel(ctx)
|
||||
agent, err := w.reCreateAgent(ctx, dialerCancel, preferredCandidateTypes)
|
||||
if err != nil {
|
||||
w.log.Errorf("failed to recreate ICE Agent: %s", err)
|
||||
return
|
||||
@@ -151,14 +163,14 @@ func (w *WorkerICE) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
|
||||
w.remoteSessionID = ""
|
||||
}
|
||||
|
||||
go w.connect(dialerCtx, agent, remoteOfferAnswer)
|
||||
go w.connect(dialerCtx, dialerCancel, agent, remoteOfferAnswer)
|
||||
}
|
||||
|
||||
// OnRemoteCandidate Handles ICE connection Candidate provided by the remote peer.
|
||||
func (w *WorkerICE) OnRemoteCandidate(candidate ice.Candidate, haRoutes route.HAMap) {
|
||||
func (w *ICE) OnRemoteCandidate(candidate ice.Candidate, haRoutes route.HAMap) {
|
||||
w.muxAgent.Lock()
|
||||
defer w.muxAgent.Unlock()
|
||||
w.log.Debugf("OnRemoteCandidate from peer %s -> %s", w.config.Key, candidate.String())
|
||||
w.log.Debugf("OnRemoteCandidate from peer %s -> %s", w.key, candidate.String())
|
||||
if w.agent == nil {
|
||||
w.log.Warnf("ICE Agent is not initialized yet")
|
||||
return
|
||||
@@ -185,18 +197,24 @@ func (w *WorkerICE) OnRemoteCandidate(candidate ice.Candidate, haRoutes route.HA
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WorkerICE) GetLocalUserCredentials() (frag string, pwd string) {
|
||||
return w.localUfrag, w.localPwd
|
||||
func (w *ICE) Credentials() signaling.Credentials {
|
||||
w.muxAgent.Lock()
|
||||
defer w.muxAgent.Unlock()
|
||||
return signaling.Credentials{
|
||||
UFrag: w.localUfrag,
|
||||
Pwd: w.localPwd,
|
||||
SessionID: w.sessionID,
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WorkerICE) InProgress() bool {
|
||||
func (w *ICE) InProgress() bool {
|
||||
w.muxAgent.Lock()
|
||||
defer w.muxAgent.Unlock()
|
||||
|
||||
return w.agentConnecting
|
||||
}
|
||||
|
||||
func (w *WorkerICE) Close() {
|
||||
func (w *ICE) Close() {
|
||||
w.muxAgent.Lock()
|
||||
defer w.muxAgent.Unlock()
|
||||
|
||||
@@ -212,10 +230,10 @@ func (w *WorkerICE) Close() {
|
||||
w.agent = nil
|
||||
}
|
||||
|
||||
func (w *WorkerICE) reCreateAgent(dialerCancel context.CancelFunc, candidates []ice.CandidateType) (*icemaker.ThreadSafeAgent, error) {
|
||||
func (w *ICE) reCreateAgent(ctx context.Context, dialerCancel context.CancelFunc, candidates []ice.CandidateType) (*icemaker.ThreadSafeAgent, error) {
|
||||
w.portForwardAttempted = false
|
||||
|
||||
agent, err := icemaker.NewAgent(w.ctx, w.iFaceDiscover, w.config.ICEConfig, candidates, w.localUfrag, w.localPwd)
|
||||
agent, err := icemaker.NewAgent(ctx, w.iFaceDiscover, w.iceConfig, candidates, w.localUfrag, w.localPwd)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create agent: %w", err)
|
||||
}
|
||||
@@ -237,7 +255,7 @@ func (w *WorkerICE) reCreateAgent(dialerCancel context.CancelFunc, candidates []
|
||||
return agent, nil
|
||||
}
|
||||
|
||||
func (w *WorkerICE) SessionID() ICESessionID {
|
||||
func (w *ICE) getSessionID() icemaker.SessionID {
|
||||
w.muxAgent.Lock()
|
||||
defer w.muxAgent.Unlock()
|
||||
|
||||
@@ -247,11 +265,11 @@ func (w *WorkerICE) SessionID() ICESessionID {
|
||||
// will block until connection succeeded
|
||||
// but it won't release if ICE Agent went into Disconnected or Failed state,
|
||||
// so we have to cancel it with the provided context once agent detected a broken connection
|
||||
func (w *WorkerICE) connect(ctx context.Context, agent *icemaker.ThreadSafeAgent, remoteOfferAnswer *OfferAnswer) {
|
||||
func (w *ICE) connect(ctx context.Context, dialerCancel context.CancelFunc, agent *icemaker.ThreadSafeAgent, remoteOfferAnswer *signaling.OfferAnswer) {
|
||||
w.log.Debugf("gather candidates")
|
||||
if err := agent.GatherCandidates(); err != nil {
|
||||
w.log.Warnf("failed to gather candidates: %s", err)
|
||||
w.closeAgent(agent, w.agentDialerCancel)
|
||||
w.closeAgent(agent, dialerCancel)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -259,19 +277,19 @@ func (w *WorkerICE) connect(ctx context.Context, agent *icemaker.ThreadSafeAgent
|
||||
remoteConn, err := w.turnAgentDial(ctx, agent, remoteOfferAnswer)
|
||||
if err != nil {
|
||||
w.log.Debugf("failed to dial the remote peer: %s", err)
|
||||
w.closeAgent(agent, w.agentDialerCancel)
|
||||
w.closeAgent(agent, dialerCancel)
|
||||
return
|
||||
}
|
||||
w.log.Debugf("agent dial succeeded")
|
||||
|
||||
pair, err := agent.GetSelectedCandidatePair()
|
||||
if err != nil {
|
||||
w.closeAgent(agent, w.agentDialerCancel)
|
||||
w.closeAgent(agent, dialerCancel)
|
||||
return
|
||||
}
|
||||
if pair == nil {
|
||||
w.log.Warnf("selected candidate pair is nil, cannot proceed")
|
||||
w.closeAgent(agent, w.agentDialerCancel)
|
||||
w.closeAgent(agent, dialerCancel)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -299,17 +317,22 @@ func (w *WorkerICE) connect(ctx context.Context, agent *icemaker.ThreadSafeAgent
|
||||
}
|
||||
w.log.Debugf("on ICE conn is ready to use")
|
||||
|
||||
w.log.Infof("connection succeeded with offer session: %s", remoteOfferAnswer.SessionIDString())
|
||||
w.muxAgent.Lock()
|
||||
if w.agent != agent {
|
||||
w.muxAgent.Unlock()
|
||||
w.log.Debugf("agent has been replaced during connect, dropping obsolete connection")
|
||||
return
|
||||
}
|
||||
w.agentConnecting = false
|
||||
w.lastSuccess = time.Now()
|
||||
w.connectedAgent = agent
|
||||
w.muxAgent.Unlock()
|
||||
|
||||
// todo: the potential problem is a race between the onConnectionStateChange
|
||||
w.conn.onICEConnectionIsReady(selectedPriority(pair), ci)
|
||||
w.log.Infof("connection succeeded with offer session: %s", remoteOfferAnswer.SessionIDString())
|
||||
w.onConnReady(selectedPriority(pair), ci)
|
||||
}
|
||||
|
||||
func (w *WorkerICE) closeAgent(agent *icemaker.ThreadSafeAgent, cancel context.CancelFunc) bool {
|
||||
func (w *ICE) closeAgent(agent *icemaker.ThreadSafeAgent, cancel context.CancelFunc) bool {
|
||||
cancel()
|
||||
if err := agent.Close(); err != nil {
|
||||
w.log.Warnf("failed to close ICE agent: %s", err)
|
||||
@@ -323,7 +346,7 @@ func (w *WorkerICE) closeAgent(agent *icemaker.ThreadSafeAgent, cancel context.C
|
||||
|
||||
if w.agent == agent {
|
||||
// consider to remove from here and move to the OnNewOffer
|
||||
sessionID, err := NewICESessionID()
|
||||
sessionID, err := icemaker.NewSessionID()
|
||||
if err != nil {
|
||||
w.log.Errorf("failed to create new session ID: %s", err)
|
||||
}
|
||||
@@ -335,7 +358,7 @@ func (w *WorkerICE) closeAgent(agent *icemaker.ThreadSafeAgent, cancel context.C
|
||||
return sessionChanged
|
||||
}
|
||||
|
||||
func (w *WorkerICE) punchRemoteWGPort(pair *ice.CandidatePair, remoteWgPort int) {
|
||||
func (w *ICE) punchRemoteWGPort(pair *ice.CandidatePair, remoteWgPort int) {
|
||||
// wait local endpoint configuration
|
||||
time.Sleep(time.Second)
|
||||
addr, err := net.ResolveUDPAddr("udp", net.JoinHostPort(pair.Remote.Address(), strconv.Itoa(remoteWgPort)))
|
||||
@@ -344,7 +367,7 @@ func (w *WorkerICE) punchRemoteWGPort(pair *ice.CandidatePair, remoteWgPort int)
|
||||
return
|
||||
}
|
||||
|
||||
mux, ok := w.config.ICEConfig.UDPMuxSrflx.(*udpmux.UniversalUDPMuxDefault)
|
||||
mux, ok := w.iceConfig.UDPMuxSrflx.(*udpmux.UniversalUDPMuxDefault)
|
||||
if !ok {
|
||||
w.log.Warn("invalid udp mux conversion")
|
||||
return
|
||||
@@ -357,7 +380,7 @@ func (w *WorkerICE) punchRemoteWGPort(pair *ice.CandidatePair, remoteWgPort int)
|
||||
|
||||
// onICECandidate is a callback attached to an ICE Agent to receive new local connection candidates
|
||||
// and then signals them to the remote peer
|
||||
func (w *WorkerICE) onICECandidate(candidate ice.Candidate) {
|
||||
func (w *ICE) onICECandidate(candidate ice.Candidate) {
|
||||
// nil means candidate gathering has been ended
|
||||
if candidate == nil {
|
||||
return
|
||||
@@ -366,9 +389,9 @@ func (w *WorkerICE) onICECandidate(candidate ice.Candidate) {
|
||||
// TODO: reported port is incorrect for CandidateTypeHost, makes understanding ICE use via logs confusing as port is ignored
|
||||
w.log.Debugf("discovered local candidate %s", candidate.String())
|
||||
go func() {
|
||||
err := w.signaler.SignalICECandidate(candidate, w.config.Key)
|
||||
err := w.signaler.SignalICECandidate(candidate, w.key)
|
||||
if err != nil {
|
||||
w.log.Errorf("failed signaling candidate to the remote peer %s %s", w.config.Key, err)
|
||||
w.log.Errorf("failed signaling candidate to the remote peer %s %s", w.key, err)
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -378,8 +401,8 @@ func (w *WorkerICE) onICECandidate(candidate ice.Candidate) {
|
||||
}
|
||||
|
||||
// injectPortForwardedCandidate signals an additional candidate using the pre-created port mapping.
|
||||
func (w *WorkerICE) injectPortForwardedCandidate(srflxCandidate ice.Candidate) {
|
||||
pfManager := w.conn.portForwardManager
|
||||
func (w *ICE) injectPortForwardedCandidate(srflxCandidate ice.Candidate) {
|
||||
pfManager := w.portForwardManager
|
||||
if pfManager == nil {
|
||||
return
|
||||
}
|
||||
@@ -407,7 +430,7 @@ func (w *WorkerICE) injectPortForwardedCandidate(srflxCandidate ice.Candidate) {
|
||||
forwardedCandidate.String(), mapping.InternalPort, mapping.ExternalPort, mapping.NATType, forwardedCandidate.Priority())
|
||||
|
||||
go func() {
|
||||
if err := w.signaler.SignalICECandidate(forwardedCandidate, w.config.Key); err != nil {
|
||||
if err := w.signaler.SignalICECandidate(forwardedCandidate, w.key); err != nil {
|
||||
w.log.Errorf("signal port-forwarded candidate: %v", err)
|
||||
}
|
||||
}()
|
||||
@@ -415,7 +438,7 @@ func (w *WorkerICE) injectPortForwardedCandidate(srflxCandidate ice.Candidate) {
|
||||
|
||||
// createForwardedCandidate creates a new server reflexive candidate with the forwarded port.
|
||||
// It uses the NAT gateway's external IP with the forwarded port.
|
||||
func (w *WorkerICE) createForwardedCandidate(srflxCandidate ice.Candidate, mapping *portforward.Mapping) (ice.Candidate, error) {
|
||||
func (w *ICE) createForwardedCandidate(srflxCandidate ice.Candidate, mapping *portforward.Mapping) (ice.Candidate, error) {
|
||||
var externalIP string
|
||||
if mapping.ExternalIP != nil && !mapping.ExternalIP.IsUnspecified() {
|
||||
externalIP = mapping.ExternalIP.String()
|
||||
@@ -460,9 +483,9 @@ func (w *WorkerICE) createForwardedCandidate(srflxCandidate ice.Candidate, mappi
|
||||
return candidate, nil
|
||||
}
|
||||
|
||||
func (w *WorkerICE) onICESelectedCandidatePair(agent *icemaker.ThreadSafeAgent, c1, c2 ice.Candidate) {
|
||||
func (w *ICE) onICESelectedCandidatePair(agent *icemaker.ThreadSafeAgent, c1, c2 ice.Candidate) {
|
||||
w.log.Debugf("selected candidate pair [local <-> remote] -> [%s <-> %s], peer %s", c1.String(), c2.String(),
|
||||
w.config.Key)
|
||||
w.key)
|
||||
|
||||
pairStat, ok := agent.GetSelectedCandidatePairStats()
|
||||
if !ok {
|
||||
@@ -471,14 +494,14 @@ func (w *WorkerICE) onICESelectedCandidatePair(agent *icemaker.ThreadSafeAgent,
|
||||
}
|
||||
|
||||
duration := time.Duration(pairStat.CurrentRoundTripTime * float64(time.Second))
|
||||
if err := w.statusRecorder.UpdateLatency(w.config.Key, duration); err != nil {
|
||||
if err := w.statusRecorder.UpdateLatency(w.key, duration); err != nil {
|
||||
w.log.Debugf("failed to update latency for peer: %s", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WorkerICE) logSuccessfulPaths(agent *icemaker.ThreadSafeAgent) {
|
||||
sessionID := w.SessionID()
|
||||
func (w *ICE) logSuccessfulPaths(agent *icemaker.ThreadSafeAgent) {
|
||||
sessionID := w.getSessionID()
|
||||
stats := agent.GetCandidatePairsStats()
|
||||
localCandidates, _ := agent.GetLocalCandidates()
|
||||
remoteCandidates, _ := agent.GetRemoteCandidates()
|
||||
@@ -508,32 +531,44 @@ func (w *WorkerICE) logSuccessfulPaths(agent *icemaker.ThreadSafeAgent) {
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WorkerICE) onConnectionStateChange(agent *icemaker.ThreadSafeAgent, dialerCancel context.CancelFunc) func(ice.ConnectionState) {
|
||||
func (w *ICE) onConnectionStateChange(agent *icemaker.ThreadSafeAgent, dialerCancel context.CancelFunc) func(ice.ConnectionState) {
|
||||
// per-agent state; pion delivers callbacks of one agent sequentially
|
||||
var connected bool
|
||||
return func(state ice.ConnectionState) {
|
||||
w.log.Debugf("ICE ConnectionState has changed to %s", state.String())
|
||||
switch state {
|
||||
case ice.ConnectionStateConnected:
|
||||
w.lastKnownState = ice.ConnectionStateConnected
|
||||
connected = true
|
||||
w.logSuccessfulPaths(agent)
|
||||
return
|
||||
case ice.ConnectionStateFailed, ice.ConnectionStateDisconnected, ice.ConnectionStateClosed:
|
||||
// ice.ConnectionStateClosed happens when we recreate the agent. For the P2P to TURN switch important to
|
||||
// notify the conn.onICEStateDisconnected changes to update the current used priority
|
||||
|
||||
sessionChanged := w.closeAgent(agent, dialerCancel)
|
||||
|
||||
if w.lastKnownState == ice.ConnectionStateConnected {
|
||||
w.lastKnownState = ice.ConnectionStateDisconnected
|
||||
w.conn.onICEStateDisconnected(sessionChanged)
|
||||
if !connected {
|
||||
return
|
||||
}
|
||||
default:
|
||||
return
|
||||
connected = false
|
||||
|
||||
w.muxAgent.Lock()
|
||||
stale := w.connectedAgent != agent
|
||||
if !stale {
|
||||
w.connectedAgent = nil
|
||||
}
|
||||
w.muxAgent.Unlock()
|
||||
|
||||
if stale {
|
||||
w.log.Debugf("suppress disconnected event of replaced ICE agent")
|
||||
return
|
||||
}
|
||||
w.onStatusDisconnect(sessionChanged)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WorkerICE) turnAgentDial(ctx context.Context, agent *icemaker.ThreadSafeAgent, remoteOfferAnswer *OfferAnswer) (*ice.Conn, error) {
|
||||
if isController(w.config) {
|
||||
func (w *ICE) turnAgentDial(ctx context.Context, agent *icemaker.ThreadSafeAgent, remoteOfferAnswer *signaling.OfferAnswer) (*ice.Conn, error) {
|
||||
if w.isController {
|
||||
return agent.Dial(ctx, remoteOfferAnswer.IceCredentials.UFrag, remoteOfferAnswer.IceCredentials.Pwd)
|
||||
} else {
|
||||
return agent.Accept(ctx, remoteOfferAnswer.IceCredentials.UFrag, remoteOfferAnswer.IceCredentials.Pwd)
|
||||
@@ -595,10 +630,10 @@ func isRelayed(pair *ice.CandidatePair) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func selectedPriority(pair *ice.CandidatePair) conntype.ConnPriority {
|
||||
func selectedPriority(pair *ice.CandidatePair) ConnPriority {
|
||||
if isRelayed(pair) {
|
||||
return conntype.ICETurn
|
||||
return ICETurn
|
||||
} else {
|
||||
return conntype.ICEP2P
|
||||
return ICEP2P
|
||||
}
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package peer
|
||||
package worker
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -10,22 +10,23 @@ import (
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal/peer/signaling"
|
||||
relayClient "github.com/netbirdio/netbird/shared/relay/client"
|
||||
)
|
||||
|
||||
type RelayConnInfo struct {
|
||||
relayedConn net.Conn
|
||||
rosenpassPubKey []byte
|
||||
rosenpassAddr string
|
||||
RelayedConn net.Conn
|
||||
RosenpassPubKey []byte
|
||||
RosenpassAddr string
|
||||
}
|
||||
|
||||
type WorkerRelay struct {
|
||||
peerCtx context.Context
|
||||
log *log.Entry
|
||||
isController bool
|
||||
config ConnConfig
|
||||
conn *Conn
|
||||
relayManager *relayClient.Manager
|
||||
log *log.Entry
|
||||
key string
|
||||
isController bool
|
||||
onConnReady func(RelayConnInfo)
|
||||
onDisconnected func()
|
||||
relayManager *relayClient.Manager
|
||||
|
||||
relayedConn net.Conn
|
||||
relayLock sync.Mutex
|
||||
@@ -33,19 +34,19 @@ type WorkerRelay struct {
|
||||
relaySupportedOnRemotePeer atomic.Bool
|
||||
}
|
||||
|
||||
func NewWorkerRelay(ctx context.Context, log *log.Entry, ctrl bool, config ConnConfig, conn *Conn, relayManager *relayClient.Manager) *WorkerRelay {
|
||||
func NewWorkerRelay(log *log.Entry, key string, isController bool, onConnReady func(RelayConnInfo), onDisconnected func(), relayManager *relayClient.Manager) *WorkerRelay {
|
||||
r := &WorkerRelay{
|
||||
peerCtx: ctx,
|
||||
log: log,
|
||||
isController: ctrl,
|
||||
config: config,
|
||||
conn: conn,
|
||||
relayManager: relayManager,
|
||||
log: log,
|
||||
key: key,
|
||||
isController: isController,
|
||||
onConnReady: onConnReady,
|
||||
onDisconnected: onDisconnected,
|
||||
relayManager: relayManager,
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func (w *WorkerRelay) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
|
||||
func (w *WorkerRelay) OnNewOffer(ctx context.Context, remoteOfferAnswer *signaling.OfferAnswer) {
|
||||
if !w.isRelaySupported(remoteOfferAnswer) {
|
||||
w.log.Infof("Relay is not supported by remote peer")
|
||||
w.relaySupportedOnRemotePeer.Store(false)
|
||||
@@ -66,7 +67,7 @@ func (w *WorkerRelay) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
|
||||
serverIP = remoteOfferAnswer.RelaySrvIP
|
||||
}
|
||||
|
||||
relayedConn, err := w.relayManager.OpenConn(w.peerCtx, srv, w.config.Key, serverIP)
|
||||
relayedConn, err := w.relayManager.OpenConn(ctx, srv, w.key, serverIP)
|
||||
if err != nil {
|
||||
if errors.Is(err, relayClient.ErrConnAlreadyExists) {
|
||||
w.log.Debugf("handled offer by reusing existing relay connection")
|
||||
@@ -88,10 +89,10 @@ func (w *WorkerRelay) OnNewOffer(remoteOfferAnswer *OfferAnswer) {
|
||||
}
|
||||
|
||||
w.log.Debugf("peer conn opened via Relay: %s", srv)
|
||||
go w.conn.onRelayConnectionIsReady(RelayConnInfo{
|
||||
relayedConn: relayedConn,
|
||||
rosenpassPubKey: remoteOfferAnswer.RosenpassPubKey,
|
||||
rosenpassAddr: remoteOfferAnswer.RosenpassAddr,
|
||||
w.onConnReady(RelayConnInfo{
|
||||
RelayedConn: relayedConn,
|
||||
RosenpassPubKey: remoteOfferAnswer.RosenpassPubKey,
|
||||
RosenpassAddr: remoteOfferAnswer.RosenpassAddr,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -119,7 +120,7 @@ func (w *WorkerRelay) CloseConn() {
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WorkerRelay) isRelaySupported(answer *OfferAnswer) bool {
|
||||
func (w *WorkerRelay) isRelaySupported(answer *signaling.OfferAnswer) bool {
|
||||
if !w.relayManager.HasRelayAddress() {
|
||||
return false
|
||||
}
|
||||
@@ -134,5 +135,5 @@ func (w *WorkerRelay) preferredRelayServer(myRelayAddress, remoteRelayAddress st
|
||||
}
|
||||
|
||||
func (w *WorkerRelay) onRelayClientDisconnected() {
|
||||
go w.conn.onRelayDisconnected()
|
||||
w.onDisconnected()
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package worker
|
||||
package peer
|
||||
|
||||
import (
|
||||
"sync/atomic"
|
||||
@@ -7,17 +7,17 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
StatusDisconnected Status = iota
|
||||
StatusConnected
|
||||
WorkerStatusDisconnected WorkerStatus = iota
|
||||
WorkerStatusConnected
|
||||
)
|
||||
|
||||
type Status int32
|
||||
type WorkerStatus int32
|
||||
|
||||
func (s Status) String() string {
|
||||
func (s WorkerStatus) String() string {
|
||||
switch s {
|
||||
case StatusDisconnected:
|
||||
case WorkerStatusDisconnected:
|
||||
return "Disconnected"
|
||||
case StatusConnected:
|
||||
case WorkerStatusConnected:
|
||||
return "Connected"
|
||||
default:
|
||||
log.Errorf("unknown status: %d", s)
|
||||
@@ -37,16 +37,16 @@ func NewAtomicStatus() *AtomicWorkerStatus {
|
||||
}
|
||||
|
||||
// Get returns the current connection status
|
||||
func (acs *AtomicWorkerStatus) Get() Status {
|
||||
return Status(acs.status.Load())
|
||||
func (acs *AtomicWorkerStatus) Get() WorkerStatus {
|
||||
return WorkerStatus(acs.status.Load())
|
||||
}
|
||||
|
||||
func (acs *AtomicWorkerStatus) SetConnected() {
|
||||
acs.status.Store(int32(StatusConnected))
|
||||
acs.status.Store(int32(WorkerStatusConnected))
|
||||
}
|
||||
|
||||
func (acs *AtomicWorkerStatus) SetDisconnected() {
|
||||
acs.status.Store(int32(StatusDisconnected))
|
||||
acs.status.Store(int32(WorkerStatusDisconnected))
|
||||
}
|
||||
|
||||
// String returns the string representation of the current status
|
||||
@@ -102,11 +102,6 @@ type ConfigInput struct {
|
||||
DNSLabels domain.List
|
||||
|
||||
MTU *uint16
|
||||
|
||||
// Owners replaces the profile's owner principal list when non-nil.
|
||||
// Shared replaces the profile's shared flag when non-nil.
|
||||
Owners []string
|
||||
Shared *bool
|
||||
}
|
||||
|
||||
// Config Configuration type
|
||||
@@ -189,16 +184,6 @@ type Config struct {
|
||||
|
||||
MTU uint16
|
||||
|
||||
// Owners lists the principals allowed to control this profile over the local
|
||||
// IPC, as typed strings: "uid:1000", "gid:1000", "group:netbird-admins"
|
||||
// (Unix, NSS-resolved) or "sid:S-1-5-..." (Windows user or group SID). Empty
|
||||
// with Shared=false means the profile is owned by nobody yet, until claimed
|
||||
Owners []string `json:"Owners,omitempty"`
|
||||
|
||||
// Shared, when true, lets any authenticated local caller control this profile
|
||||
// (opt-in). Takes precedence over Owners.
|
||||
Shared bool `json:"Shared,omitempty"`
|
||||
|
||||
// policy is the MDM policy that produced the currently-set values for
|
||||
// any MDM-enforced fields. Set by applyMDMPolicy at the tail of apply()
|
||||
// and reset on every apply() invocation. Never persisted to disk.
|
||||
@@ -657,18 +642,6 @@ func (config *Config) apply(input ConfigInput) (updated bool, err error) {
|
||||
updated = true
|
||||
}
|
||||
|
||||
if input.Owners != nil && !slices.Equal(config.Owners, input.Owners) {
|
||||
log.Infof("updating profile owners to %v", input.Owners)
|
||||
config.Owners = input.Owners
|
||||
updated = true
|
||||
}
|
||||
|
||||
if input.Shared != nil && *input.Shared != config.Shared {
|
||||
log.Infof("updating profile shared flag to %t", *input.Shared)
|
||||
config.Shared = *input.Shared
|
||||
updated = true
|
||||
}
|
||||
|
||||
// MDM is the last override layer: any key present in the policy
|
||||
// supersedes defaults, on-disk config, env vars and CLI input.
|
||||
config.applyMDMPolicy(loadMDMPolicy())
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user