Compare commits

...

1 Commits

Author SHA1 Message Date
Theodor S. Midtlien
f44040feb0 WIP 2026-07-21 17:15:15 +02:00
35 changed files with 1188 additions and 172 deletions

View File

@@ -6,6 +6,7 @@ import (
"fmt"
"io"
"io/fs"
"net"
"os"
"os/signal"
"path"
@@ -79,6 +80,8 @@ var (
updateSettingsDisabled bool
captureEnabled bool
networksDisabled bool
socketOwner string
strictSocketDisabled bool
rootCmd = &cobra.Command{
Use: "netbird",
@@ -143,10 +146,12 @@ func init() {
defaultDaemonAddr := "unix:///var/run/netbird.sock"
if runtime.GOOS == "windows" {
defaultDaemonAddr = "tcp://127.0.0.1:41731"
// Named pipe (not loopback TCP): the pipe SDDL gates who may connect and
// the pipe client token carries the caller's SID for per-RPC authorization.
defaultDaemonAddr = "npipe://netbird"
}
rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp]://[path|host:port]")
rootCmd.PersistentFlags().StringVar(&daemonAddr, "daemon-addr", defaultDaemonAddr, "Daemon service address to serve CLI requests [unix|tcp|npipe]://[path|host:port|name]")
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")
@@ -265,16 +270,31 @@ func FlagNameToEnvVar(cmdFlag string, prefix string) string {
}
// DialClientGRPCServer returns client connection to the daemon server.
//
// The daemon reads the caller's kernel identity from the transport (SO_PEERCRED
// on a Unix socket, the client token on a Windows named pipe), so the client
// side uses insecure (plaintext) credentials — it needs no cooperation to be
// identified. For npipe addresses we install a context dialer since gRPC's
// resolver does not understand Windows named pipes.
func DialClientGRPCServer(ctx context.Context, addr string) (*grpc.ClientConn, error) {
ctx, cancel := context.WithTimeout(ctx, time.Second*10)
defer cancel()
return grpc.DialContext(
ctx,
strings.TrimPrefix(addr, "tcp://"),
opts := []grpc.DialOption{
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithBlock(),
)
}
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 grpc.DialContext(ctx, target, opts...)
}
// WithBackOff execute function in backoff cycle.

View File

@@ -56,6 +56,9 @@ func init() {
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]://[path|host:port]. Requires --enable-json-socket to serve. To persist, use: netbird service install --enable-json-socket --json-socket")
serviceCmd.PersistentFlags().StringVar(&socketOwner, "socket-owner", "", "user to own the daemon control socket; restricts it to that user plus the netbird group (0660). If unset, the first client to connect claims ownership (trust-on-first-use). Persisted via: netbird service install --socket-owner")
serviceCmd.PersistentFlags().BoolVar(&strictSocketDisabled, "disable-strict-socket", false, "leave the daemon control socket world-writable (0666) instead of restricting it (root-only, discouraged). Persisted via: netbird service install --disable-strict-socket")
rootCmd.PersistentFlags().StringVarP(&serviceName, "service", "s", defaultServiceName, "Netbird system service name")
serviceEnvDesc := `Sets extra environment variables for the service. ` +
`You can specify a comma-separated list of KEY=VALUE pairs. ` +

View File

@@ -5,6 +5,7 @@ package cmd
import (
"context"
"fmt"
"runtime"
"time"
"github.com/kardianos/service"
@@ -13,12 +14,31 @@ 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 returns the gRPC server options that install peer-identity
// transport credentials on the daemon control channel. Identity extraction is
// only possible over a Unix socket (SO_PEERCRED) or Windows named pipe (client
// token); over TCP, or on platforms without a peer-credential primitive, the
// daemon runs without per-caller authorization and logs a warning.
func daemonServerOptions(network string) []grpc.ServerOption {
creds := ipcauth.NewTransportCredentials()
if creds == nil {
log.Warnf("daemon control channel 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)}
}
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")
@@ -37,8 +57,13 @@ func (p *program) Start(svc service.Service) error {
// Collect static system and platform information
system.UpdateStaticInfoAsync()
// in any case, even if configuration does not exists we run daemon to serve CLI gRPC API.
p.serv = grpc.NewServer()
network, _, err := parseListenAddress(daemonAddr)
if err != nil {
return fmt.Errorf("parse daemon address: %w", err)
}
// in any case, even if configuration does not exist we run daemon to serve the CLI gRPC API.
p.serv = grpc.NewServer(daemonServerOptions(network)...)
daemonListener, err := listenOnAddress(daemonAddr)
if err != nil {
@@ -62,7 +87,8 @@ func (p *program) Start(svc service.Service) error {
defer jsonListener.Close()
}
if err := daemonListener.chmodUnixSocket("daemon"); err != nil {
serveListener, err := secureDaemonListener(daemonListener)
if err != nil {
log.Error(err)
return
}
@@ -84,6 +110,7 @@ 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 as the daemon's own identity and BYPASSES per-caller authorization; restrict access to %s separately", jsonSocket)
if err := p.startJSONGateway(jsonListener, daemonAddr); err != nil {
log.Fatalf("failed to start daemon JSON server: %v", err)
}
@@ -92,7 +119,7 @@ func (p *program) Start(svc service.Service) error {
}
log.Printf("started daemon server: %v", daemonListener.address)
if err := p.serv.Serve(daemonListener.Listener); err != nil {
if err := p.serv.Serve(serveListener); err != nil {
log.Errorf("failed to serve daemon requests: %v", err)
}
}()

View File

@@ -71,6 +71,14 @@ func buildServiceArguments() []string {
args = append(args, "--enable-json-socket", "--json-socket", jsonSocket)
}
if socketOwner != "" {
args = append(args, "--socket-owner", socketOwner)
}
if strictSocketDisabled {
args = append(args, "--disable-strict-socket")
}
return args
}

View File

@@ -32,6 +32,8 @@ type serviceParams struct {
EnableCapture bool `json:"enable_capture,omitempty"`
DisableNetworks bool `json:"disable_networks,omitempty"`
EnableJSONSocket bool `json:"enable_json_socket,omitempty"`
SocketOwner string `json:"socket_owner,omitempty"`
DisableStrictSocket bool `json:"disable_strict_socket,omitempty"`
ServiceEnvVars map[string]string `json:"service_env_vars,omitempty"`
}
@@ -86,6 +88,8 @@ func currentServiceParams() *serviceParams {
EnableCapture: captureEnabled,
DisableNetworks: networksDisabled,
EnableJSONSocket: enableJSONSocket,
SocketOwner: socketOwner,
DisableStrictSocket: strictSocketDisabled,
}
if len(serviceEnvVars) > 0 {
@@ -165,6 +169,14 @@ func applyServiceParams(cmd *cobra.Command, params *serviceParams) {
networksDisabled = params.DisableNetworks
}
if !serviceCmd.PersistentFlags().Changed("socket-owner") {
socketOwner = params.SocketOwner
}
if !serviceCmd.PersistentFlags().Changed("disable-strict-socket") {
strictSocketDisabled = params.DisableStrictSocket
}
applyServiceEnvParams(cmd, params)
}

View File

@@ -0,0 +1,20 @@
//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)
}

View File

@@ -0,0 +1,32 @@
//go:build windows
package cmd
import (
"context"
"net"
"time"
"github.com/Microsoft/go-winio"
"github.com/netbirdio/netbird/client/internal/ipcauth"
)
// listenNamedPipe creates the daemon control named pipe with a tight SDDL
// (SYSTEM + Administrators + interactive users). ListenPipe fails if the pipe
// already exists (first-instance semantics), which prevents a squatting process
// from pre-creating it — we surface that error loudly rather than falling back.
func listenNamedPipe(path string) (net.Listener, error) {
return winio.ListenPipe(path, &winio.PipeConfig{
SecurityDescriptor: ipcauth.DefaultPipeSDDL(),
})
}
// dialNamedPipe connects to the daemon control named pipe.
func dialNamedPipe(ctx context.Context, path string) (net.Conn, error) {
if deadline, ok := ctx.Deadline(); ok {
timeout := time.Until(deadline)
return winio.DialPipe(path, &timeout)
}
return winio.DialPipeContext(ctx, path)
}

View File

@@ -26,6 +26,15 @@ 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)
}
@@ -45,13 +54,23 @@ func parseListenAddress(addr string) (string, string, error) {
}
switch network {
case "unix", "tcp":
case "unix", "tcp", "npipe":
return network, address, nil
default:
return "", "", fmt.Errorf("unsupported daemon address protocol: %v", network)
}
}
// pipePath maps a daemon-addr npipe name (e.g. "netbird" from "npipe://netbird")
// to a Windows named-pipe path (\\.\pipe\netbird). A caller may also pass a full
// \\.\pipe\ path, which is returned unchanged.
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 {

View File

@@ -0,0 +1,11 @@
//go:build windows
package cmd
import "net"
// secureDaemonListener is a no-op on Windows: the named-pipe SDDL gates who may
// connect (Layer 1), and the pipe client token supplies per-RPC identity.
func secureDaemonListener(l *socketListener) (net.Listener, error) {
return l.Listener, nil
}

View File

@@ -0,0 +1,225 @@
//go:build !windows && !ios && !android
package cmd
import (
"errors"
"fmt"
"net"
"os"
"os/exec"
"os/user"
"strconv"
"sync"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/client/internal/ipcauth"
"github.com/netbirdio/netbird/client/internal/shell"
)
// secureDaemonListener applies the Layer-1 access control to the daemon control
// socket and returns the listener to serve on. For a Unix socket this restricts
// the socket to an owner (plus the netbird group); for anything else it is a
// no-op (TCP is legacy/unauthenticated; named pipes are gated by their SDDL).
func secureDaemonListener(l *socketListener) (net.Listener, error) {
if l.network != "unix" {
return l.Listener, nil
}
owner := effectiveSocketOwner()
switch {
case strictSocketDisabled:
// Root-only opt-out (via service.json): leave it world-writable.
if err := os.Chmod(l.address, 0666); err != nil {
return nil, fmt.Errorf("set daemon socket permissions: %w", err)
}
log.Warnf("daemon control socket left world-writable (0666) by --disable-strict-socket")
return l.Listener, nil
case owner != "":
// Seeded owner (flag, MDM, or persisted TOFU result): restrict before
// serving so there is no open window.
uid, err := lookupUser(owner)
if err != nil {
return nil, fmt.Errorf("lookup socket owner %q: %w", owner, err)
}
if err := restrictSocket(l.address, uid); err != nil {
return nil, fmt.Errorf("restrict socket to %q: %w", owner, err)
}
return l.Listener, nil
default:
// Trust-on-first-use: open the socket now; tofuListener locks it to the
// first caller's uid on the first connection.
if err := os.Chmod(l.address, 0666); err != nil {
return nil, fmt.Errorf("set daemon socket permissions: %w", err)
}
return &tofuListener{Listener: l.Listener, path: l.address, owner: -1}, nil
}
}
func lookupUser(username string) (int, error) {
u, err := shell.LookupWithGetent(username)
if err != nil {
return -1, fmt.Errorf("lookup user %s: %w", username, err)
}
uid, err := strconv.Atoi(u.Uid)
if err != nil {
return -1, fmt.Errorf("parse uid %s: %w", u.Uid, err)
}
return uid, nil
}
// addGroup creates a system group if it doesn't already exist and returns the gid.
// Must run as root.
func addGroup(name string) (int, error) {
group, err := shell.LookupGroupWithGetent(name)
if err == nil {
gid, err := strconv.ParseInt(group.Gid, 10, 64)
return int(gid), err
}
groupadd, err := exec.LookPath("groupadd")
if err != nil {
// Fallback for Alpine/BusyBox systems.
if groupadd, err = exec.LookPath("addgroup"); err != nil {
return -1, errors.New("neither groupadd nor addgroup found")
}
}
// Use --system for a service/daemon group (no login, low GID).
out, err := exec.Command(groupadd, "--system", name).CombinedOutput()
if err != nil {
return -1, fmt.Errorf("create group %q: %w: %s", name, err, out)
}
if group, err := shell.LookupGroupWithGetent(name); err == nil {
gid, err := strconv.ParseInt(group.Gid, 10, 64)
return int(gid), err
}
return -1, fmt.Errorf("lookup group %q: %w", name, err)
}
// restrictSocket locks the unix socket down to the owner uid plus the netbird
// group (0660). If the group cannot be created or applied, it fails closed to
// owner-only 0600 — it never leaves the socket world-writable.
func restrictSocket(path string, uid int) error {
gid, err := addGroup("netbird")
if err != nil {
log.Errorf("create netbird group, failing closed to owner-only 0600: %v", err)
return chownChmod(path, uid, -1, 0600)
}
if err := chownChmod(path, uid, gid, 0660); err != nil {
log.Errorf("apply netbird group to socket, failing closed to owner-only 0600: %v", err)
return chownChmod(path, uid, -1, 0600)
}
return nil
}
// chownChmod sets ownership and mode on the socket. A gid of -1 leaves the
// group unchanged.
func chownChmod(path string, uid, gid int, mode os.FileMode) error {
if err := os.Chown(path, uid, gid); err != nil {
return fmt.Errorf("chown socket %s: %w", path, err)
}
if err := os.Chmod(path, mode); err != nil {
return fmt.Errorf("chmod socket %s: %w", path, err)
}
return nil
}
// tofuListener implements trust-on-first-use for the daemon control socket.
// The socket starts world-writable; the first caller's uid (read via SO_PEERCRED)
// becomes the owner. On that first connection the socket is restricted and the
// owner persisted so the open window never reopens on later starts. Connections
// that raced in during the open window and are neither the owner nor root are
// dropped. Changing the socket mode does not disturb the already-open
// connection, so the first caller's request is served normally.
type tofuListener struct {
net.Listener
path string
mu sync.Mutex
owner int // -1 until claimed
}
func (l *tofuListener) Accept() (net.Conn, error) {
for {
c, err := l.Listener.Accept()
if err != nil {
return nil, err
}
id, err := ipcauth.PeerIdentity(c)
if err != nil {
log.Errorf("read peer credentials, dropping connection: %v", err)
_ = c.Close()
continue
}
uid := int(id.UID)
l.mu.Lock()
if l.owner == -1 {
if err := restrictSocket(l.path, uid); err != nil {
l.mu.Unlock()
_ = c.Close()
// Refuse to serve on a socket we could not lock down.
return nil, fmt.Errorf("restrict socket on first connection: %w", err)
}
l.owner = uid
persistSocketOwner(uid)
log.Infof("control socket restricted to first caller (uid %d)", uid)
l.mu.Unlock()
return c, nil
}
owner := l.owner
l.mu.Unlock()
// New connects are already gated by the 0660 perms set above; this only
// drops anything that slipped in during the brief open window.
if uid != owner && uid != 0 {
log.Warnf("dropping non-owner connection (uid %d) during socket bootstrap", uid)
_ = c.Close()
continue
}
return c, nil
}
}
// effectiveSocketOwner returns the configured socket owner: the --socket-owner
// flag when set, otherwise the owner persisted by a previous TOFU migration.
func effectiveSocketOwner() string {
if socketOwner != "" {
return socketOwner
}
params, err := loadServiceParams()
if err != nil {
log.Errorf("load service params for socket owner: %v", err)
return ""
}
if params != nil {
return params.SocketOwner
}
return ""
}
// persistSocketOwner records the TOFU-selected owner (by username) so the next
// daemon start restricts the socket immediately, with no open window.
func persistSocketOwner(uid int) {
u, err := user.LookupId(strconv.Itoa(uid))
if err != nil {
log.Errorf("resolve uid %d to username for persistence: %v", uid, err)
return
}
params, err := loadServiceParams()
if err != nil {
log.Errorf("load service params to persist socket owner: %v", err)
return
}
if params == nil {
params = currentServiceParams()
}
params.SocketOwner = u.Username
if err := saveServiceParams(params); err != nil {
log.Errorf("persist socket owner: %v", err)
}
}

View File

@@ -0,0 +1,13 @@
//go:build !linux && !darwin && !freebsd && !windows
package ipcauth
import "google.golang.org/grpc/credentials"
// NewTransportCredentials returns nil on platforms without a peer-identity
// primitive. The daemon falls back to insecure credentials and skips per-RPC
// authorization (logging a warning), preserving pre-hardening behavior until
// the transport gains an identity primitive.
func NewTransportCredentials() credentials.TransportCredentials {
return nil
}

View File

@@ -0,0 +1,48 @@
//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. It is non-nil on platforms with a
// peer-credential primitive.
func NewTransportCredentials() credentials.TransportCredentials {
return unixCreds{}
}
// unixCreds implements credentials.TransportCredentials over a Unix socket.
// The server side reads SO_PEERCRED/LOCAL_PEERCRED during the handshake; the
// client side is a no-op (the kernel supplies the peer identity to the server
// without any client cooperation).
type unixCreds struct{}
func (unixCreds) ClientHandshake(_ context.Context, _ string, conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
return conn, AuthInfo{}, nil
}
// 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 := PeerIdentity(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: "netbird-ipc-peercred"}
}
func (unixCreds) Clone() credentials.TransportCredentials { return unixCreds{} }
func (unixCreds) OverrideServerName(string) error { return nil }

View File

@@ -0,0 +1,142 @@
//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")
procRevertToSelf = modadvapi32.NewProc("RevertToSelf")
)
// Windows group-SID attribute flags (winnt.h): a group only counts toward
// membership when it is enabled and not marked use-for-deny-only.
const (
seGroupEnabled = 0x00000004
seGroupUseForDenyOnly = 0x00000010
)
// DefaultPipeSDDL restricts the daemon control pipe to LocalSystem (SY), the
// Administrators group (BA), and interactive logon users (IU). It deliberately
// excludes Authenticated Users / Everyone so remote or arbitrary service
// principals cannot connect. This is the Layer-1 channel gate; the interceptor
// (Layer 2) further restricts by per-profile ownership.
func DefaultPipeSDDL() string {
return "D:P(A;;GA;;;SY)(A;;GA;;;BA)(A;;GA;;;IU)"
}
// NewTransportCredentials returns gRPC transport credentials that derive the
// caller's identity from the named-pipe client token, following Microsoft's
// "Verifying Client Access with ACLs" pattern: ImpersonateNamedPipeClient ->
// OpenThreadToken -> RevertToSelf. Per threat-model M-NOIMP, impersonation is
// used only to read the client token for identity, never to perform privileged work.
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
}
// ServerHandshake extracts the connecting client's identity from the pipe token.
// Fails closed if the handle or token cannot be read.
func (winpipeCreds) ServerHandshake(conn net.Conn) (net.Conn, credentials.AuthInfo, error) {
// go-winio's pipe connection embeds *win32File, which exposes Fd().
fdConn, ok := conn.(interface{ Fd() uintptr })
if !ok {
return nil, nil, fmt.Errorf("connection %T does not expose a pipe handle", conn)
}
handle := windows.Handle(fdConn.Fd())
id, err := pipeClientIdentity(handle)
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: "netbird-ipc-peercred"}
}
func (winpipeCreds) Clone() credentials.TransportCredentials { return winpipeCreds{} }
func (winpipeCreds) OverrideServerName(string) error { return nil }
// pipeClientIdentity reads the connecting client's user SID and enabled group
// SIDs from the named-pipe handle. The impersonation window is kept as small as
// possible and pinned to the OS thread (impersonation is thread-local).
func pipeClientIdentity(handle windows.Handle) (Identity, error) {
var pid uint32
hasPID := windows.GetNamedPipeClientProcessId(handle, &pid) == nil
runtime.LockOSThread()
defer runtime.UnlockOSThread()
if err := impersonateNamedPipeClient(handle); err != nil {
return Identity{}, fmt.Errorf("impersonate named pipe client: %w", err)
}
defer func() { _ = revertToSelf() }()
// 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&seGroupEnabled == 0 || g.Attributes&seGroupUseForDenyOnly != 0 {
continue
}
groups = append(groups, g.Sid.String())
}
return Identity{
SID: tu.User.Sid.String(),
Groups: groups,
PID: int32(pid),
HasPID: hasPID,
}, nil
}
func impersonateNamedPipeClient(h windows.Handle) error {
r, _, e := procImpersonateNamedPipeClient.Call(uintptr(h))
if r == 0 {
return e
}
return nil
}
func revertToSelf() error {
r, _, e := procRevertToSelf.Call()
if r == 0 {
return e
}
return nil
}

View File

@@ -0,0 +1,93 @@
// Package ipcauth provides kernel-authenticated caller identity for the daemon's
// local IPC (gRPC) channel and the transport credentials that populate it.
//
// It is the identity foundation shared by two layers of the local-IPC hardening:
// - the socket-permission layer (Layer 1, client/cmd), which reads the peer
// identity to gate who may connect and to run trust-on-first-use; and
// - the per-RPC authorization interceptor (Layer 2), which reads the same
// identity from the gRPC context to enforce per-profile ownership.
//
// 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 (the daemon logs a warning and stays open,
// preserving today's behavior until the transport is hardened).
package ipcauth
import (
"context"
"fmt"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/peer"
)
// Identity is the kernel-authenticated identity of a local IPC caller.
//
// The zero value is not a valid identity; callers obtain one via
// IdentityFromContext (which reports presence) or PeerIdentity.
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
// PID is the caller's process ID, for audit only. HasPID is false when the
// platform cannot supply it (e.g. Darwin/FreeBSD xucred carries no PID).
PID int32
HasPID bool
// 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 time (empty on Unix, where supplementary group
// membership is resolved on demand via NSS/getent by the authorizer).
Groups []string
}
// 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 != ""
}
// String renders the identity for audit logs.
func (i Identity) String() string {
if i.IsWindows() {
if i.HasPID {
return fmt.Sprintf("sid=%s pid=%d", i.SID, i.PID)
}
return fmt.Sprintf("sid=%s", i.SID)
}
if i.HasPID {
return fmt.Sprintf("uid=%d gid=%d pid=%d", i.UID, i.GID, i.PID)
}
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 (e.g. an unsupported platform, or a caller that
// did not come through the daemon socket) — 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
}

View File

@@ -0,0 +1,47 @@
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, PID: 4242, HasPID: true}
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_String(t *testing.T) {
assert.Equal(t, "uid=1000 gid=1000 pid=42", Identity{UID: 1000, GID: 1000, PID: 42, HasPID: true}.String())
assert.Equal(t, "uid=1000 gid=1000", Identity{UID: 1000, GID: 1000}.String())
assert.Equal(t, "sid=S-1-5-21-1 pid=42", Identity{SID: "S-1-5-21-1", PID: 42, HasPID: true}.String())
assert.Equal(t, "sid=S-1-5-21-1", Identity{SID: "S-1-5-21-1"}.String())
}
func TestIdentity_IsWindows(t *testing.T) {
assert.True(t, Identity{SID: "S-1-5-18"}.IsWindows())
assert.False(t, Identity{UID: 0}.IsWindows())
}

View File

@@ -0,0 +1,43 @@
//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 group list but no pid, so audit on these platforms is
// uid/gid-based (HasPID is false); PID via LOCAL_PEERPID is a possible follow-up.
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
}

View File

@@ -0,0 +1,43 @@
//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,
PID: cred.Pid,
HasPID: true,
}, nil
}

View File

@@ -0,0 +1,16 @@
//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)
}

View File

@@ -0,0 +1,113 @@
//go:build linux || darwin || freebsd
package ipcauth
import (
"net"
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// 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")
}
// TestUnixCreds_ServerHandshake exercises the transport-credentials path end to end.
func TestUnixCreds_ServerHandshake(t *testing.T) {
creds := NewTransportCredentials()
require.NotNil(t, creds)
sock := filepath.Join(t.TempDir(), "hs.sock")
ln, err := net.Listen("unix", sock)
require.NoError(t, err)
t.Cleanup(func() { _ = ln.Close() })
type result struct {
info interface{ AuthType() string }
err error
}
done := make(chan result, 1)
go func() {
c, aerr := ln.Accept()
if aerr != nil {
done <- result{err: aerr}
return
}
_, ai, herr := creds.ServerHandshake(c)
if herr != nil {
done <- result{err: herr}
return
}
done <- result{info: ai}
}()
client, err := net.Dial("unix", sock)
require.NoError(t, err)
t.Cleanup(func() { _ = client.Close() })
res := <-done
require.NoError(t, res.err)
ai, ok := res.info.(AuthInfo)
require.True(t, ok, "expected ipcauth.AuthInfo, got %T", res.info)
assert.Equal(t, uint32(os.Getuid()), ai.Identity.UID)
assert.Equal(t, "netbird-ipc-peercred", ai.AuthType())
}

View File

@@ -0,0 +1,29 @@
//go:build cgo && !osusergo && !windows
package shell
import "os/user"
// LookupWithGetent with CGO delegates directly to os/user.Lookup.
// When CGO is enabled, os/user uses libc (getpwnam_r) which goes through
// the NSS stack natively. If it fails, the user truly doesn't exist and
// getent would also fail.
func LookupWithGetent(username string) (*user.User, error) {
return user.Lookup(username)
}
// CurrentUserWithGetent with CGO delegates directly to os/user.Current.
func CurrentUserWithGetent() (*user.User, error) {
return user.Current()
}
// LookupGroupWithGetent returns the resolved group from either a gid or groupname.
func LookupGroupWithGetent(name string) (*user.Group, error) {
return user.LookupGroup(name)
}
// GroupIdsWithFallback with CGO delegates directly to user.GroupIds.
// libc's getgrouplist handles NSS groups natively.
func GroupIdsWithFallback(u *user.User) ([]string, error) {
return u.GroupIds()
}

View File

@@ -1,6 +1,6 @@
//go:build (!cgo || osusergo) && !windows
package server
package shell
import (
"os"
@@ -10,10 +10,10 @@ import (
log "github.com/sirupsen/logrus"
)
// lookupWithGetent looks up a user by name, falling back to getent if os/user fails.
// LookupWithGetent looks up a user by name, falling back to getent if os/user fails.
// Without CGO, os/user only reads /etc/passwd and misses NSS-provided users.
// getent goes through the host's NSS stack.
func lookupWithGetent(username string) (*user.User, error) {
func LookupWithGetent(username string) (*user.User, error) {
u, err := user.Lookup(username)
if err == nil {
return u, nil
@@ -22,7 +22,7 @@ func lookupWithGetent(username string) (*user.User, error) {
stdErr := err
log.Debugf("os/user.Lookup(%q) failed, trying getent: %v", username, err)
u, _, getentErr := runGetent(username)
u, _, getentErr := runGetentPasswd(username)
if getentErr != nil {
log.Debugf("getent fallback for %q also failed: %v", username, getentErr)
return nil, stdErr
@@ -31,8 +31,26 @@ func lookupWithGetent(username string) (*user.User, error) {
return u, nil
}
// currentUserWithGetent gets the current user, falling back to getent if os/user fails.
func currentUserWithGetent() (*user.User, error) {
// LookupGroupWithGetent returns the resolved group from either a gid or groupname,
// falling back to getent if os/user fails (NSS groups under nocgo).
func LookupGroupWithGetent(name string) (*user.Group, error) {
g, err := user.LookupGroup(name)
if err == nil {
return g, nil
}
stdErr := err
log.Debugf("os/user.LookupGroup(%q) failed, trying getent: %v", name, err)
g, getentErr := runGetentGroup(name)
if getentErr != nil {
log.Debugf("getent fallback for %q also failed: %v", name, getentErr)
return nil, stdErr
}
return g, nil
}
// CurrentUserWithGetent gets the current user, falling back to getent if os/user fails.
func CurrentUserWithGetent() (*user.User, error) {
u, err := user.Current()
if err == nil {
return u, nil
@@ -42,7 +60,7 @@ func currentUserWithGetent() (*user.User, error) {
uid := strconv.Itoa(os.Getuid())
log.Debugf("os/user.Current() failed, trying getent with UID %s: %v", uid, err)
u, _, getentErr := runGetent(uid)
u, _, getentErr := runGetentPasswd(uid)
if getentErr != nil {
return nil, stdErr
}
@@ -50,14 +68,14 @@ func currentUserWithGetent() (*user.User, error) {
return u, nil
}
// groupIdsWithFallback gets group IDs for a user via the id command first,
// GroupIdsWithFallback gets group IDs for a user via the id command first,
// falling back to user.GroupIds().
// NOTE: unlike lookupWithGetent/currentUserWithGetent which try stdlib first,
// NOTE: unlike LookupWithGetent/CurrentUserWithGetent which try stdlib first,
// this intentionally tries `id -G` first because without CGO, user.GroupIds()
// only reads /etc/group and silently returns incomplete results for NSS users
// (no error, just missing groups). The id command goes through NSS and returns
// the full set.
func groupIdsWithFallback(u *user.User) ([]string, error) {
func GroupIdsWithFallback(u *user.User) ([]string, error) {
ids, err := runIdGroups(u.Username)
if err == nil {
return ids, nil

View File

@@ -1,4 +1,4 @@
package server
package shell
import (
"os/user"
@@ -15,7 +15,7 @@ func TestLookupWithGetent_CurrentUser(t *testing.T) {
current, err := user.Current()
require.NoError(t, err)
u, err := lookupWithGetent(current.Username)
u, err := LookupWithGetent(current.Username)
require.NoError(t, err)
assert.Equal(t, current.Username, u.Username)
assert.Equal(t, current.Uid, u.Uid)
@@ -23,7 +23,7 @@ func TestLookupWithGetent_CurrentUser(t *testing.T) {
}
func TestLookupWithGetent_NonexistentUser(t *testing.T) {
_, err := lookupWithGetent("nonexistent_user_xyzzy_12345")
_, err := LookupWithGetent("nonexistent_user_xyzzy_12345")
require.Error(t, err, "should fail for nonexistent user")
}
@@ -31,7 +31,7 @@ func TestCurrentUserWithGetent(t *testing.T) {
stdUser, err := user.Current()
require.NoError(t, err)
u, err := currentUserWithGetent()
u, err := CurrentUserWithGetent()
require.NoError(t, err)
assert.Equal(t, stdUser.Uid, u.Uid)
assert.Equal(t, stdUser.Username, u.Username)
@@ -41,7 +41,7 @@ func TestGroupIdsWithFallback_CurrentUser(t *testing.T) {
current, err := user.Current()
require.NoError(t, err)
groups, err := groupIdsWithFallback(current)
groups, err := GroupIdsWithFallback(current)
require.NoError(t, err)
require.NotEmpty(t, groups, "current user should have at least one group")
@@ -56,7 +56,7 @@ func TestGroupIdsWithFallback_CurrentUser(t *testing.T) {
func TestGetShellFromGetent_CurrentUser(t *testing.T) {
if runtime.GOOS == "windows" {
// Windows stub always returns empty, which is correct
shell := getShellFromGetent("1000")
shell := GetShellFromGetent("1000")
assert.Empty(t, shell, "Windows stub should return empty")
return
}
@@ -65,9 +65,9 @@ func TestGetShellFromGetent_CurrentUser(t *testing.T) {
require.NoError(t, err)
// getent may not be available on all systems (e.g., macOS without Homebrew getent)
shell := getShellFromGetent(current.Uid)
shell := GetShellFromGetent(current.Uid)
if shell == "" {
t.Log("getShellFromGetent returned empty, getent may not be available")
t.Log("GetShellFromGetent returned empty, getent may not be available")
return
}
assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell)
@@ -78,7 +78,7 @@ func TestLookupWithGetent_RootUser(t *testing.T) {
t.Skip("no root user on Windows")
}
u, err := lookupWithGetent("root")
u, err := LookupWithGetent("root")
if err != nil {
t.Skip("root user not available on this system")
}
@@ -86,25 +86,25 @@ func TestLookupWithGetent_RootUser(t *testing.T) {
}
// TestIntegration_FullLookupChain exercises the complete user lookup chain
// against the real system, testing that all wrappers (lookupWithGetent,
// currentUserWithGetent, groupIdsWithFallback, getShellFromGetent) produce
// against the real system, testing that all wrappers (LookupWithGetent,
// CurrentUserWithGetent, GroupIdsWithFallback, GetShellFromGetent) produce
// consistent and correct results when composed together.
func TestIntegration_FullLookupChain(t *testing.T) {
// Step 1: currentUserWithGetent must resolve the running user.
current, err := currentUserWithGetent()
require.NoError(t, err, "currentUserWithGetent must resolve the running user")
// Step 1: CurrentUserWithGetent must resolve the running user.
current, err := CurrentUserWithGetent()
require.NoError(t, err, "CurrentUserWithGetent must resolve the running user")
require.NotEmpty(t, current.Uid)
require.NotEmpty(t, current.Username)
// Step 2: lookupWithGetent by the same username must return matching identity.
byName, err := lookupWithGetent(current.Username)
// Step 2: LookupWithGetent by the same username must return matching identity.
byName, err := LookupWithGetent(current.Username)
require.NoError(t, err)
assert.Equal(t, current.Uid, byName.Uid, "lookup by name should return same UID")
assert.Equal(t, current.Gid, byName.Gid, "lookup by name should return same GID")
assert.Equal(t, current.HomeDir, byName.HomeDir, "lookup by name should return same home")
// Step 3: groupIdsWithFallback must return at least the primary GID.
groups, err := groupIdsWithFallback(current)
// Step 3: GroupIdsWithFallback must return at least the primary GID.
groups, err := GroupIdsWithFallback(current)
require.NoError(t, err)
require.NotEmpty(t, groups, "user must have at least one group")
@@ -120,10 +120,10 @@ func TestIntegration_FullLookupChain(t *testing.T) {
}
assert.True(t, foundPrimary, "primary GID %s should appear in supplementary groups", current.Gid)
// Step 4: getShellFromGetent should either return a valid shell path or empty
// Step 4: GetShellFromGetent should either return a valid shell path or empty
// (empty is OK when getent is not available, e.g. macOS without Homebrew getent).
if runtime.GOOS != "windows" {
shell := getShellFromGetent(current.Uid)
shell := GetShellFromGetent(current.Uid)
if shell != "" {
assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell)
}
@@ -131,17 +131,17 @@ func TestIntegration_FullLookupChain(t *testing.T) {
}
// TestIntegration_LookupAndGroupsConsistency verifies that a user resolved via
// lookupWithGetent can have their groups resolved via groupIdsWithFallback,
// LookupWithGetent can have their groups resolved via GroupIdsWithFallback,
// testing the handoff between the two functions as used by the SSH server.
func TestIntegration_LookupAndGroupsConsistency(t *testing.T) {
current, err := user.Current()
require.NoError(t, err)
// Simulate the SSH server flow: lookup user, then get their groups.
resolved, err := lookupWithGetent(current.Username)
resolved, err := LookupWithGetent(current.Username)
require.NoError(t, err)
groups, err := groupIdsWithFallback(resolved)
groups, err := GroupIdsWithFallback(resolved)
require.NoError(t, err)
require.NotEmpty(t, groups, "resolved user must have groups")
@@ -156,7 +156,7 @@ func TestIntegration_LookupAndGroupsConsistency(t *testing.T) {
}
// TestIntegration_ShellLookupChain tests the full shell resolution chain
// (getShellFromPasswd -> getShellFromGetent -> $SHELL -> default) on Unix.
// (getShellFromPasswd -> GetShellFromGetent -> $SHELL -> default) on Unix.
func TestIntegration_ShellLookupChain(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("Unix shell lookup not applicable on Windows")
@@ -165,8 +165,8 @@ func TestIntegration_ShellLookupChain(t *testing.T) {
current, err := user.Current()
require.NoError(t, err)
// getUserShell is the top-level function used by the SSH server.
shell := getUserShell(current.Uid)
require.NotEmpty(t, shell, "getUserShell must always return a shell")
// GetUserShell is the top-level function used by the SSH server.
shell := GetUserShell(current.Uid)
require.NotEmpty(t, shell, "GetUserShell must always return a shell")
assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell)
}

View File

@@ -1,6 +1,6 @@
//go:build !windows
package server
package shell
import (
"context"
@@ -14,19 +14,26 @@ import (
const getentTimeout = 5 * time.Second
// getShellFromGetent gets a user's login shell via getent by UID.
// GetShellFromGetent gets a user's login shell via getent by UID.
// This is needed even with CGO because getShellFromPasswd reads /etc/passwd
// directly and won't find NSS-provided users there.
func getShellFromGetent(userID string) string {
_, shell, err := runGetent(userID)
func GetShellFromGetent(userID string) string {
_, shell, err := runGetentPasswd(userID)
if err != nil {
return ""
}
return shell
}
// runGetent executes `getent passwd <query>` and returns the user and login shell.
func runGetent(query string) (*user.User, string, error) {
// GetUserFromGetent returns the resolved user from either a uid or username,
// going through the host's NSS stack.
func GetUserFromGetent(query string) (*user.User, error) {
u, _, err := runGetentPasswd(query)
return u, err
}
// runGetentPasswd executes `getent passwd <query>` and returns the user and login shell.
func runGetentPasswd(query string) (*user.User, string, error) {
if !validateGetentInput(query) {
return nil, "", fmt.Errorf("invalid getent input: %q", query)
}
@@ -42,7 +49,24 @@ func runGetent(query string) (*user.User, string, error) {
return parseGetentPasswd(string(out))
}
// parseGetentPasswd parses getent passwd output: "name:x:uid:gid:gecos:home:shell"
// runGetentGroup executes `getent group <query>` and returns the group.
func runGetentGroup(query string) (*user.Group, error) {
if !validateGetentInput(query) {
return nil, fmt.Errorf("invalid getent input: %q", query)
}
ctx, cancel := context.WithTimeout(context.Background(), getentTimeout)
defer cancel()
out, err := exec.CommandContext(ctx, "getent", "group", query).Output()
if err != nil {
return nil, fmt.Errorf("getent group %s: %w", query, err)
}
return parseGetentGroup(string(out))
}
// parseGetentPasswd parses getent passwd output: "name:x:uid:gid:gecos:home:shell".
func parseGetentPasswd(output string) (*user.User, string, error) {
fields := strings.SplitN(strings.TrimSpace(output), ":", 8)
if len(fields) < 6 {
@@ -67,6 +91,20 @@ func parseGetentPasswd(output string) (*user.User, string, error) {
}, shell, nil
}
// parseGetentGroup parses getent group output: "group:x:gid:members".
func parseGetentGroup(output string) (*user.Group, error) {
fields := strings.SplitN(strings.TrimSpace(output), ":", 8)
if len(fields) < 4 {
return nil, fmt.Errorf("unexpected getent output (need 4+ fields): %q", output)
}
if fields[0] == "" || fields[2] == "" {
return nil, fmt.Errorf("missing required fields in getent output: %q", output)
}
return &user.Group{Gid: fields[2], Name: fields[0]}, nil
}
// validateGetentInput checks that the input is safe to pass to getent or id.
// Allows POSIX usernames, numeric UIDs, and common NSS extensions
// (@ for Kerberos, $ for Samba, + for NIS compat). A leading hyphen is

View File

@@ -1,6 +1,6 @@
//go:build !windows
package server
package shell
import (
"os/exec"
@@ -198,7 +198,7 @@ func TestRunGetent_RootUser(t *testing.T) {
t.Skip("getent not available on this system")
}
u, shell, err := runGetent("root")
u, shell, err := runGetentPasswd("root")
require.NoError(t, err)
assert.Equal(t, "root", u.Username)
assert.Equal(t, "0", u.Uid)
@@ -211,7 +211,7 @@ func TestRunGetent_ByUID(t *testing.T) {
t.Skip("getent not available on this system")
}
u, _, err := runGetent("0")
u, _, err := runGetentPasswd("0")
require.NoError(t, err)
assert.Equal(t, "root", u.Username)
assert.Equal(t, "0", u.Uid)
@@ -222,15 +222,15 @@ func TestRunGetent_NonexistentUser(t *testing.T) {
t.Skip("getent not available on this system")
}
_, _, err := runGetent("nonexistent_user_xyzzy_12345")
_, _, err := runGetentPasswd("nonexistent_user_xyzzy_12345")
assert.Error(t, err)
}
func TestRunGetent_InvalidInput(t *testing.T) {
_, _, err := runGetent("")
_, _, err := runGetentPasswd("")
assert.Error(t, err)
_, _, err = runGetent("user\x00name")
_, _, err = runGetentPasswd("user\x00name")
assert.Error(t, err)
}
@@ -239,7 +239,7 @@ func TestRunGetent_NotAvailable(t *testing.T) {
t.Skip("getent is available, can't test missing case")
}
_, _, err := runGetent("root")
_, _, err := runGetentPasswd("root")
assert.Error(t, err, "should fail when getent is not installed")
}
@@ -286,7 +286,7 @@ func TestGetentResultsMatchStdlib(t *testing.T) {
current, err := user.Current()
require.NoError(t, err)
getentUser, _, err := runGetent(current.Username)
getentUser, _, err := runGetentPasswd(current.Username)
require.NoError(t, err)
assert.Equal(t, current.Username, getentUser.Username, "username should match")
@@ -303,7 +303,7 @@ func TestGetentResultsMatchStdlib_ByUID(t *testing.T) {
current, err := user.Current()
require.NoError(t, err)
getentUser, _, err := runGetent(current.Uid)
getentUser, _, err := runGetentPasswd(current.Uid)
require.NoError(t, err)
assert.Equal(t, current.Username, getentUser.Username, "username should match when looked up by UID")
@@ -359,7 +359,7 @@ func TestGetShellFromPasswd_CurrentUser(t *testing.T) {
assert.True(t, shell[0] == '/', "shell should be an absolute path, got %q", shell)
if _, err := exec.LookPath("getent"); err == nil {
_, getentShell, getentErr := runGetent(current.Uid)
_, getentShell, getentErr := runGetentPasswd(current.Uid)
if getentErr == nil && getentShell != "" {
assert.Equal(t, getentShell, shell, "shell from /etc/passwd should match getent")
}
@@ -403,7 +403,7 @@ func TestGetShellFromPasswd_MatchesGetentForKnownUsers(t *testing.T) {
continue
}
_, getentShell, err := runGetent(uid)
_, getentShell, err := runGetentPasswd(uid)
if err != nil {
continue
}

View File

@@ -0,0 +1,31 @@
//go:build windows
package shell
import "os/user"
// LookupWithGetent on Windows just delegates to os/user.Lookup.
// Windows does not use NSS/getent; its user lookup works without CGO.
func LookupWithGetent(username string) (*user.User, error) {
return user.Lookup(username)
}
// CurrentUserWithGetent on Windows just delegates to os/user.Current.
func CurrentUserWithGetent() (*user.User, error) {
return user.Current()
}
// LookupGroupWithGetent on Windows just delegates to os/user.LookupGroup.
func LookupGroupWithGetent(name string) (*user.Group, error) {
return user.LookupGroup(name)
}
// GetShellFromGetent is a no-op on Windows; shell resolution uses PowerShell detection.
func GetShellFromGetent(_ string) string {
return ""
}
// GroupIdsWithFallback on Windows just delegates to u.GroupIds().
func GroupIdsWithFallback(u *user.User) ([]string, error) {
return u.GroupIds()
}

View File

@@ -1,17 +1,14 @@
package server
package shell
import (
"bufio"
"fmt"
"net"
"os"
"os/exec"
"os/user"
"runtime"
"strconv"
"strings"
"github.com/gliderlabs/ssh"
log "github.com/sirupsen/logrus"
)
@@ -22,9 +19,9 @@ const (
powershellExe = "powershell.exe"
)
// getUserShell returns the appropriate shell for the given user ID
// Handles all platform-specific logic and fallbacks consistently
func getUserShell(userID string) string {
// GetUserShell returns the appropriate shell for the given user ID.
// Handles all platform-specific logic and fallbacks consistently.
func GetUserShell(userID string) string {
switch runtime.GOOS {
case "windows":
return getWindowsUserShell()
@@ -56,7 +53,7 @@ func getUnixUserShell(userID string) string {
return shell
}
if shell := getShellFromGetent(userID); shell != "" {
if shell := GetShellFromGetent(userID); shell != "" {
return shell
}
@@ -67,7 +64,7 @@ func getUnixUserShell(userID string) string {
return defaultUnixShell
}
// getShellFromPasswd reads the shell from /etc/passwd for the given user ID
// getShellFromPasswd reads the shell from /etc/passwd for the given user ID.
func getShellFromPasswd(userID string) string {
file, err := os.Open("/etc/passwd")
if err != nil {
@@ -101,8 +98,8 @@ func getShellFromPasswd(userID string) string {
return ""
}
// prepareUserEnv prepares environment variables for user execution
func prepareUserEnv(user *user.User, shell string) []string {
// PrepareUserEnv prepares environment variables for user execution.
func PrepareUserEnv(user *user.User, shell string) []string {
pathValue := "/usr/local/bin:/usr/bin:/bin:/usr/local/games:/usr/games"
if runtime.GOOS == "windows" {
pathValue = `C:\Windows\System32;C:\Windows;C:\Windows\System32\Wbem;C:\Windows\System32\WindowsPowerShell\v1.0`
@@ -117,9 +114,9 @@ func prepareUserEnv(user *user.User, shell string) []string {
}
}
// acceptEnv checks if environment variable from SSH client should be accepted
// This is a whitelist of variables that SSH clients can send to the server
func acceptEnv(envVar string) bool {
// AcceptEnv checks if an environment variable from an SSH client should be accepted.
// This is a whitelist of variables that SSH clients can send to the server.
func AcceptEnv(envVar string) bool {
varName := envVar
if idx := strings.Index(envVar, "="); idx != -1 {
varName = envVar[:idx]
@@ -156,29 +153,3 @@ func acceptEnv(envVar string) bool {
return false
}
// prepareSSHEnv prepares SSH protocol-specific environment variables
// These variables provide information about the SSH connection itself
func prepareSSHEnv(session ssh.Session) []string {
remoteAddr := session.RemoteAddr()
localAddr := session.LocalAddr()
remoteHost, remotePort, err := net.SplitHostPort(remoteAddr.String())
if err != nil {
remoteHost = remoteAddr.String()
remotePort = "0"
}
localHost, localPort, err := net.SplitHostPort(localAddr.String())
if err != nil {
localHost = localAddr.String()
localPort = strconv.Itoa(InternalSSHPort)
}
return []string{
// SSH_CLIENT format: "client_ip client_port server_port"
fmt.Sprintf("SSH_CLIENT=%s %s %s", remoteHost, remotePort, localPort),
// SSH_CONNECTION format: "client_ip client_port server_ip server_port"
fmt.Sprintf("SSH_CONNECTION=%s %s %s %s", remoteHost, remotePort, localHost, localPort),
}
}

View File

@@ -20,6 +20,8 @@ import (
"github.com/creack/pty"
"github.com/gliderlabs/ssh"
log "github.com/sirupsen/logrus"
shellutil "github.com/netbirdio/netbird/client/internal/shell"
)
// ptyManager manages Pty file operations with thread safety
@@ -146,10 +148,10 @@ func (s *Server) createShellCommand(ctx context.Context, shell string, args []st
// prepareCommandEnv prepares environment variables for command execution on Unix
func (s *Server) prepareCommandEnv(_ *log.Entry, localUser *user.User, session ssh.Session) []string {
env := prepareUserEnv(localUser, getUserShell(localUser.Uid))
env := shellutil.PrepareUserEnv(localUser, shellutil.GetUserShell(localUser.Uid))
env = append(env, prepareSSHEnv(session)...)
for _, v := range session.Environ() {
if acceptEnv(v) {
if shellutil.AcceptEnv(v) {
env = append(env, v)
}
}

View File

@@ -15,6 +15,7 @@ import (
"golang.org/x/sys/windows"
"golang.org/x/sys/windows/registry"
shellutil "github.com/netbirdio/netbird/client/internal/shell"
"github.com/netbirdio/netbird/client/ssh/server/winpty"
)
@@ -247,10 +248,10 @@ func (s *Server) prepareCommandEnv(logger *log.Entry, localUser *user.User, sess
userEnv, err := s.getUserEnvironment(logger, username, domain)
if err != nil {
log.Debugf("failed to get user environment for %s\\%s, using fallback: %v", domain, username, err)
env := prepareUserEnv(localUser, getUserShell(localUser.Uid))
env := shellutil.PrepareUserEnv(localUser, shellutil.GetUserShell(localUser.Uid))
env = append(env, prepareSSHEnv(session)...)
for _, v := range session.Environ() {
if acceptEnv(v) {
if shellutil.AcceptEnv(v) {
env = append(env, v)
}
}
@@ -260,7 +261,7 @@ func (s *Server) prepareCommandEnv(logger *log.Entry, localUser *user.User, sess
env := userEnv
env = append(env, prepareSSHEnv(session)...)
for _, v := range session.Environ() {
if acceptEnv(v) {
if shellutil.AcceptEnv(v) {
env = append(env, v)
}
}
@@ -273,7 +274,7 @@ func (s *Server) handlePtyLogin(logger *log.Entry, session ssh.Session, privileg
return false
}
shell := getUserShell(privilegeResult.User.Uid)
shell := shellutil.GetUserShell(privilegeResult.User.Uid)
logger.Infof("starting interactive shell: %s", shell)
s.executeCommandWithPty(logger, session, nil, privilegeResult, ptyReq, nil)
@@ -384,7 +385,7 @@ func (s *Server) executeCommandWithPty(logger *log.Entry, session ssh.Session, _
}
username, domain := s.parseUsername(localUser.Username)
shell := getUserShell(localUser.Uid)
shell := shellutil.GetUserShell(localUser.Uid)
req := PtyExecutionRequest{
Shell: shell,

View File

@@ -1,24 +0,0 @@
//go:build cgo && !osusergo && !windows
package server
import "os/user"
// lookupWithGetent with CGO delegates directly to os/user.Lookup.
// When CGO is enabled, os/user uses libc (getpwnam_r) which goes through
// the NSS stack natively. If it fails, the user truly doesn't exist and
// getent would also fail.
func lookupWithGetent(username string) (*user.User, error) {
return user.Lookup(username)
}
// currentUserWithGetent with CGO delegates directly to os/user.Current.
func currentUserWithGetent() (*user.User, error) {
return user.Current()
}
// groupIdsWithFallback with CGO delegates directly to user.GroupIds.
// libc's getgrouplist handles NSS groups natively.
func groupIdsWithFallback(u *user.User) ([]string, error) {
return u.GroupIds()
}

View File

@@ -1,26 +0,0 @@
//go:build windows
package server
import "os/user"
// lookupWithGetent on Windows just delegates to os/user.Lookup.
// Windows does not use NSS/getent; its user lookup works without CGO.
func lookupWithGetent(username string) (*user.User, error) {
return user.Lookup(username)
}
// currentUserWithGetent on Windows just delegates to os/user.Current.
func currentUserWithGetent() (*user.User, error) {
return user.Current()
}
// getShellFromGetent is a no-op on Windows; shell resolution uses PowerShell detection.
func getShellFromGetent(_ string) string {
return ""
}
// groupIdsWithFallback on Windows just delegates to u.GroupIds().
func groupIdsWithFallback(u *user.User) ([]string, error) {
return u.GroupIds()
}

View File

@@ -0,0 +1,35 @@
package server
import (
"fmt"
"net"
"strconv"
"github.com/gliderlabs/ssh"
)
// prepareSSHEnv prepares SSH protocol-specific environment variables
// These variables provide information about the SSH connection itself
func prepareSSHEnv(session ssh.Session) []string {
remoteAddr := session.RemoteAddr()
localAddr := session.LocalAddr()
remoteHost, remotePort, err := net.SplitHostPort(remoteAddr.String())
if err != nil {
remoteHost = remoteAddr.String()
remotePort = "0"
}
localHost, localPort, err := net.SplitHostPort(localAddr.String())
if err != nil {
localHost = localAddr.String()
localPort = strconv.Itoa(InternalSSHPort)
}
return []string{
// SSH_CLIENT format: "client_ip client_port server_port"
fmt.Sprintf("SSH_CLIENT=%s %s %s", remoteHost, remotePort, localPort),
// SSH_CONNECTION format: "client_ip client_port server_ip server_port"
fmt.Sprintf("SSH_CONNECTION=%s %s %s %s", remoteHost, remotePort, localHost, localPort),
}
}

View File

@@ -9,6 +9,8 @@ import (
"strings"
log "github.com/sirupsen/logrus"
shellutil "github.com/netbirdio/netbird/client/internal/shell"
)
var (
@@ -23,8 +25,8 @@ func isPlatformUnix() bool {
// Dependency injection variables for testing - allows mocking dynamic runtime checks
var (
getCurrentUser = currentUserWithGetent
lookupUser = lookupWithGetent
getCurrentUser = shellutil.CurrentUserWithGetent
lookupUser = shellutil.LookupWithGetent
getCurrentOS = func() string { return runtime.GOOS }
getIsProcessPrivileged = isCurrentProcessPrivileged

View File

@@ -16,6 +16,8 @@ import (
"github.com/gliderlabs/ssh"
log "github.com/sirupsen/logrus"
shellutil "github.com/netbirdio/netbird/client/internal/shell"
)
// POSIX portable filename character set regex: [a-zA-Z0-9._-]
@@ -160,7 +162,7 @@ func (s *Server) parseUserCredentials(localUser *user.User) (uint32, uint32, []u
// getSupplementaryGroups retrieves supplementary group IDs for a user.
// Uses id/getent fallback for NSS users in CGO_ENABLED=0 builds.
func (s *Server) getSupplementaryGroups(u *user.User) ([]uint32, error) {
groupIDStrings, err := groupIdsWithFallback(u)
groupIDStrings, err := shellutil.GroupIdsWithFallback(u)
if err != nil {
return nil, fmt.Errorf("get group IDs for user %s: %w", u.Username, err)
}
@@ -196,7 +198,7 @@ func (s *Server) createExecutorCommand(logger *log.Entry, session ssh.Session, l
GID: gid,
Groups: groups,
WorkingDir: localUser.HomeDir,
Shell: getUserShell(localUser.Uid),
Shell: shellutil.GetUserShell(localUser.Uid),
Command: session.RawCommand(),
PTY: hasPty,
}
@@ -228,7 +230,7 @@ func (s *Server) createPtyCommand(privilegeResult PrivilegeCheckResult, ptyReq s
func (s *Server) createDirectPtyCommand(session ssh.Session, localUser *user.User, ptyReq ssh.Pty) *exec.Cmd {
log.Debugf("creating direct Pty command for user %s (no user switching needed)", localUser.Username)
shell := getUserShell(localUser.Uid)
shell := shellutil.GetUserShell(localUser.Uid)
args := s.getShellCommandArgs(shell, session.RawCommand())
cmd := s.createShellCommand(session.Context(), shell, args)
@@ -245,12 +247,12 @@ func (s *Server) preparePtyEnv(localUser *user.User, ptyReq ssh.Pty, session ssh
termType = "xterm-256color"
}
env := prepareUserEnv(localUser, getUserShell(localUser.Uid))
env := shellutil.PrepareUserEnv(localUser, shellutil.GetUserShell(localUser.Uid))
env = append(env, prepareSSHEnv(session)...)
env = append(env, fmt.Sprintf("TERM=%s", termType))
for _, v := range session.Environ() {
if acceptEnv(v) {
if shellutil.AcceptEnv(v) {
env = append(env, v)
}
}

View File

@@ -13,6 +13,8 @@ import (
"github.com/gliderlabs/ssh"
log "github.com/sirupsen/logrus"
"golang.org/x/sys/windows"
shellutil "github.com/netbirdio/netbird/client/internal/shell"
)
// validateUsername validates Windows usernames according to SAM Account Name rules
@@ -104,7 +106,7 @@ func (s *Server) createExecutorCommand(logger *log.Entry, session ssh.Session, l
func (s *Server) createUserSwitchCommand(logger *log.Entry, session ssh.Session, localUser *user.User) (*exec.Cmd, func(), error) {
username, domain := s.parseUsername(localUser.Username)
shell := getUserShell(localUser.Uid)
shell := shellutil.GetUserShell(localUser.Uid)
rawCmd := session.RawCommand()
var command string

2
go.mod
View File

@@ -30,6 +30,7 @@ require (
require (
github.com/DeRuina/timberjack v1.4.2
github.com/Microsoft/go-winio v0.6.2
github.com/awnumar/memguard v0.23.0
github.com/aws/aws-sdk-go-v2 v1.38.3
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.1
@@ -156,7 +157,6 @@ require (
github.com/Masterminds/goutils v1.1.1 // indirect
github.com/Masterminds/semver/v3 v3.4.0 // indirect
github.com/Masterminds/sprig/v3 v3.3.0 // indirect
github.com/Microsoft/go-winio v0.6.2 // indirect
github.com/adrg/xdg v0.5.3 // indirect
github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be // indirect
github.com/apapsch/go-jsonmerge/v2 v2.0.0 // indirect