session-helper: move srv.RunAndExit to session/reexec (#65263)

* Move the reexec parts of lib/srv to session/reexec

* Update references to the moved bits of lib/srv

* Avoid testutils in session/reexec

* Shuffle some constants around to avoid imports in session/reexec

* Vendor in the relevant parts of logutils in session/reexec

* Remove unnecessary symlink checks when opening files

* Inline the last two things from lib/utils in session/reexec

* Inline moved constants

* Deprecate group name consts and fix missed inlines

* Add missing session/reexec.TestMain with reexec check

* Reuse existing constants from the log constants package in session

* Fix broken godoc link
This commit is contained in:
Edoardo Spadolini
2026-04-08 09:40:52 +00:00
committed by GitHub
parent c3d5543cba
commit 825b23049e
44 changed files with 2252 additions and 800 deletions
+14
View File
@@ -610,3 +610,17 @@ const (
// EICEDisabledMessage is the message that gets returned to the user when they try to use this functionality.
EICEDisabledMessage = "support for accessing EC2 instances using EC2 Instance Connect Endpoint was removed"
)
const (
// TeleportDropGroup is a default group that users of the teleport automated user
// provisioning system get added to when provisioned in INSECURE_DROP mode. This
// prevents already existing users from being tampered with or deleted.
TeleportDropGroup = "teleport-system"
// TeleportKeepGroup is a default group that users of the teleport automated user
// provisioning system get added to when provisioned in KEEP mode. This prevents
// already existing users from being tampered with or deleted.
TeleportKeepGroup = "teleport-keep"
// TeleportStaticGroup is a default group that static host users get added to. This
// prevents already existing users from being tampered with or deleted.
TeleportStaticGroup = "teleport-static"
)
+14 -3
View File
@@ -17,6 +17,7 @@ limitations under the License.
package types
import (
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/types/common"
)
@@ -1763,18 +1764,28 @@ var KubernetesCoreResourceKinds = map[string]struct{}{
"services": {},
}
// TODO(espadolini): delete in v20
const (
// TeleportDropGroup is a default group that users of the teleport automated user
// provisioning system get added to when provisioned in INSECURE_DROP mode. This
// prevents already existing users from being tampered with or deleted.
TeleportDropGroup = "teleport-system"
//
// Deprecated: use [constants.TeleportDropGroup].
//go:fix inline
TeleportDropGroup = constants.TeleportDropGroup
// TeleportKeepGroup is a default group that users of the teleport automated user
// provisioning system get added to when provisioned in KEEP mode. This prevents
// already existing users from being tampered with or deleted.
TeleportKeepGroup = "teleport-keep"
//
// Deprecated: use [constants.TeleportKeepGroup].
//go:fix inline
TeleportKeepGroup = constants.TeleportKeepGroup
// TeleportStaticGroup is a default group that static host users get added to. This
// prevents already existing users from being tampered with or deleted.
TeleportStaticGroup = "teleport-static"
//
// Deprecated: use [constants.TeleportStaticGroup].
//go:fix inline
TeleportStaticGroup = constants.TeleportStaticGroup
)
const (
-40
View File
@@ -577,24 +577,6 @@ const (
JumpCloud = "jumpcloud"
)
const (
// RemoteCommandSuccess is returned when a command has successfully executed.
RemoteCommandSuccess = 0
// RemoteCommandFailure is returned when a command has failed to execute and
// we don't have another status code for it.
RemoteCommandFailure = 255
// HomeDirNotFound is returned when the "teleport checkhomedir" command cannot
// find the user's home directory.
HomeDirNotFound = 254
// HomeDirNotAccessible is returned when the "teleport checkhomedir" command has
// found the user's home directory, but the user does NOT have permissions to
// access it.
HomeDirNotAccessible = 253
// UnexpectedCredentials is returned when a command is no longer running with the expected
// credentials.
UnexpectedCredentials = 252
)
// MaxResourceSize is the maximum size (in bytes) of a serialized resource. This limit is
// typically only enforced against resources that are likely to arbitrarily grow (e.g. PluginData).
const MaxResourceSize = 1000000
@@ -958,28 +940,6 @@ const (
)
const (
// ExecSubCommand is the sub-command Teleport uses to re-exec itself for
// command execution (exec and shells).
ExecSubCommand = "exec"
// NetworkingSubCommand is the sub-command Teleport uses to re-exec itself
// for networking operations. e.g. local/remote port forwarding, agent forwarding,
// or x11 forwarding.
NetworkingSubCommand = "networking"
// CheckHomeDirSubCommand is the sub-command Teleport uses to re-exec itself
// to check if the user's home directory exists.
CheckHomeDirSubCommand = "checkhomedir"
// ParkSubCommand is the sub-command Teleport uses to re-exec itself as a
// specific UID to prevent the matching user from being deleted before
// spawning the intended child process.
ParkSubCommand = "park"
// SFTPSubCommand is the sub-command Teleport uses to re-exec itself to
// handle SFTP connections.
SFTPSubCommand = "sftp"
// WaitSubCommand is the sub-command Teleport uses to wait
// until a domain name stops resolving. Its main use is to ensure no
// auth instances are still running the previous major version.
+4 -4
View File
@@ -16,17 +16,17 @@
package teleport
import "github.com/gravitational/teleport/session/logutils"
import "github.com/gravitational/teleport/session/logconstants"
// static assertions that [logutils.ComponentKey] and [logutils.ComponentFields]
// static assertions that [logconstants.ComponentKey] and [logconstants.ComponentFields]
// are equal to the respective consts defined in this package; we can't just
// define them to be equal because the true definition belongs here and we want
// to avoid circular module requirements
func _() {
const mustBeTrue = ComponentKey == logutils.ComponentKey
const mustBeTrue = ComponentKey == logconstants.ComponentKey
_ = map[bool]struct{}{false: struct{}{}, mustBeTrue: struct{}{}}
}
func _() {
const mustBeTrue = ComponentFields == logutils.ComponentFields
const mustBeTrue = ComponentFields == logconstants.ComponentFields
_ = map[bool]struct{}{false: struct{}{}, mustBeTrue: struct{}{}}
}
+2 -2
View File
@@ -26,8 +26,8 @@ import (
"github.com/gravitational/teleport/lib/cryptosuites/cryptosuitestest"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/utils/log/logtest"
"github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/tool/teleport/common"
)
@@ -42,7 +42,7 @@ func TestMainImplementation(m *testing.M) {
modules.SetInsecureTestMode(true)
// If the test is re-executing itself, execute the command that comes over
// the pipe.
if srv.IsReexec() {
if reexec.IsReexec() {
defer cancel()
common.Run(common.Options{Args: os.Args[1:]})
return
+9 -8
View File
@@ -40,6 +40,7 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
apiconstants "github.com/gravitational/teleport/api/constants"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
labelv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/label/v1"
userprovisioningpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/userprovisioning/v2"
@@ -253,7 +254,7 @@ func TestRootHostUsers(t *testing.T) {
closer, err := users.UpsertUser(testuser, &decisionpb.HostUsersInfo{Groups: testGroups, Mode: decisionpb.HostUserMode_HOST_USER_MODE_DROP})
require.NoError(t, err)
testGroups = append(testGroups, types.TeleportDropGroup)
testGroups = append(testGroups, apiconstants.TeleportDropGroup)
t.Cleanup(func() { cleanupUsersAndGroups([]string{testuser}, testGroups) })
u, err := user.Lookup(testuser)
@@ -282,7 +283,7 @@ func TestRootHostUsers(t *testing.T) {
})
require.NoError(t, err)
t.Cleanup(func() { cleanupUsersAndGroups([]string{testuser}, []string{types.TeleportDropGroup}) })
t.Cleanup(func() { cleanupUsersAndGroups([]string{testuser}, []string{apiconstants.TeleportDropGroup}) })
group, err := user.LookupGroupId(testGID)
require.NoError(t, err)
@@ -307,7 +308,7 @@ func TestRootHostUsers(t *testing.T) {
closer, err := users.UpsertUser(testuser, &decisionpb.HostUsersInfo{Mode: decisionpb.HostUserMode_HOST_USER_MODE_KEEP})
require.NoError(t, err)
require.Nil(t, closer)
t.Cleanup(func() { cleanupUsersAndGroups([]string{testuser}, []string{types.TeleportKeepGroup}) })
t.Cleanup(func() { cleanupUsersAndGroups([]string{testuser}, []string{apiconstants.TeleportKeepGroup}) })
u, err := user.Lookup(testuser)
require.NoError(t, err)
@@ -373,7 +374,7 @@ func TestRootHostUsers(t *testing.T) {
t.Cleanup(func() {
cleanupUsersAndGroups(
[]string{"teleport-user1", "teleport-user2", "teleport-user3", "teleport-user4"},
[]string{types.TeleportDropGroup, types.TeleportKeepGroup})
[]string{apiconstants.TeleportDropGroup, apiconstants.TeleportKeepGroup})
})
err = users.DeleteAllUsers()
@@ -598,13 +599,13 @@ func TestRootHostUsers(t *testing.T) {
})
t.Run("Test migrate unmanaged user", func(t *testing.T) {
t.Cleanup(func() { cleanupUsersAndGroups([]string{testuser}, []string{types.TeleportKeepGroup}) })
t.Cleanup(func() { cleanupUsersAndGroups([]string{testuser}, []string{apiconstants.TeleportKeepGroup}) })
users := srv.NewHostUsers(context.Background(), presence, "host_uuid")
_, err := host.UserAdd(testuser, nil, host.UserOpts{})
require.NoError(t, err)
closer, err := users.UpsertUser(testuser, &decisionpb.HostUsersInfo{Mode: decisionpb.HostUserMode_HOST_USER_MODE_KEEP, Groups: []string{types.TeleportKeepGroup}})
closer, err := users.UpsertUser(testuser, &decisionpb.HostUsersInfo{Mode: decisionpb.HostUserMode_HOST_USER_MODE_KEEP, Groups: []string{apiconstants.TeleportKeepGroup}})
require.NoError(t, err)
require.Nil(t, closer)
@@ -614,7 +615,7 @@ func TestRootHostUsers(t *testing.T) {
gids, err := u.GroupIds()
require.NoError(t, err)
keepGroup, err := user.LookupGroup(types.TeleportKeepGroup)
keepGroup, err := user.LookupGroup(apiconstants.TeleportKeepGroup)
require.NoError(t, err)
require.Contains(t, gids, keepGroup.Gid)
})
@@ -888,7 +889,7 @@ func testStaticHostUsers(t *testing.T, nodeUUID, goodLogin, goodLoginWithShell,
userGroups = append(userGroups, group.Name)
}
require.Subset(t, userGroups, groups)
require.Contains(t, userGroups, types.TeleportStaticGroup)
require.Contains(t, userGroups, apiconstants.TeleportStaticGroup)
// Check that the sudoers file was created.
require.FileExists(t, sudoersPath(goodLogin, nodeUUID))
userShells, err := getUserShells("/etc/passwd")
+10 -9
View File
@@ -61,6 +61,7 @@ import (
"github.com/gravitational/teleport/session/envutils"
"github.com/gravitational/teleport/session/networking/x11"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec"
)
var ctxID int32
@@ -207,7 +208,7 @@ type Server interface {
// ChildLogConfig is the log configuration for handling logs from child processes.
type ChildLogConfig struct {
ExecLogConfig
reexec.ExecLogConfig
// Writer is the output writer to use for the logger. May be nil.
Writer io.Writer
@@ -1066,7 +1067,7 @@ func (c *ServerContext) LogValue() slog.Value {
)
}
func getPAMConfig(c *ServerContext) (*PAMConfig, error) {
func getPAMConfig(c *ServerContext) (*reexec.PAMConfig, error) {
// PAM should be disabled.
if c.srv.Component() != teleport.ComponentNode {
return nil, nil
@@ -1117,7 +1118,7 @@ func getPAMConfig(c *ServerContext) (*PAMConfig, error) {
}
}
return &PAMConfig{
return &reexec.PAMConfig{
UsePAMAuth: localPAMConfig.UsePAMAuth,
ServiceName: localPAMConfig.ServiceName,
Environment: environment,
@@ -1126,7 +1127,7 @@ func getPAMConfig(c *ServerContext) (*PAMConfig, error) {
// ExecCommand takes a *ServerContext and extracts the parts needed to create
// an *execCommand which can be re-sent to Teleport.
func (c *ServerContext) ExecCommand() (*ExecCommand, error) {
func (c *ServerContext) ExecCommand() (*reexec.ExecCommand, error) {
// Extract the command to be executed. This only exists if command execution
// (exec or shell) is being requested, port forwarding has no command to
// execute.
@@ -1163,7 +1164,7 @@ func (c *ServerContext) ExecCommand() (*ExecCommand, error) {
}
// Create the execCommand that will be sent to the child process.
return &ExecCommand{
return &reexec.ExecCommand{
LogConfig: c.srv.ChildLogConfig().ExecLogConfig,
Command: command,
DestinationAddress: c.DstAddr,
@@ -1285,10 +1286,10 @@ func closeAll(closers ...io.Closer) error {
return trace.NewAggregate(errs...)
}
func newUaccMetadata(c *ServerContext) (*UaccMetadata, error) {
func newUaccMetadata(c *ServerContext) (*reexec.UaccMetadata, error) {
utmpPath, wtmpPath, btmpPath, wtmpdbPath := c.srv.GetUserAccountingPaths()
return &UaccMetadata{
RemoteAddr: utils.FromAddr(c.ConnectionContext.ServerConn.RemoteAddr()),
return &reexec.UaccMetadata{
RemoteAddr: reexec.NetAddrFromAddr(c.ConnectionContext.ServerConn.RemoteAddr()),
UtmpPath: utmpPath,
WtmpPath: wtmpPath,
BtmpPath: btmpPath,
@@ -1408,7 +1409,7 @@ func (c *ServerContext) WaitForChild(ctx context.Context) error {
// Session Recording events to the SSH session.
var waitErr error
if bpfService.Enabled() {
if waitErr = waitForSignal(ctx, c.readyr, childReadyWaitTimeout); waitErr != nil {
if waitErr = reexec.WaitForSignal(ctx, c.readyr, childReadyWaitTimeout); waitErr != nil {
c.Logger.ErrorContext(ctx, "Child process never became ready.", "error", waitErr)
}
}
+99 -74
View File
@@ -19,8 +19,8 @@
package srv
import (
"bufio"
"context"
"encoding/json"
"errors"
"fmt"
"io"
@@ -43,16 +43,9 @@ import (
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/sshutils/reexec"
"github.com/gravitational/teleport/lib/utils"
)
const (
defaultPath = "/bin:/usr/bin:/usr/local/bin:/sbin"
defaultEnvPath = "PATH=" + defaultPath
defaultRootPath = "/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"
defaultEnvRootPath = "PATH=" + defaultRootPath
defaultTerm = "xterm"
defaultLoginDefsPath = "/etc/login.defs"
"github.com/gravitational/teleport/session/envutils"
sessionreexec "github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
// ExecResult is used internally to send the result of a command execution from
@@ -591,66 +584,6 @@ func emitExecAuditEvent(ctx *ServerContext, cmd string, execErr error) {
}
}
// getDefaultEnvPath returns the default value of PATH environment variable for
// new logins (prior to shell) based on login.defs. Returns a string which
// looks like "PATH=/usr/bin:/bin"
func getDefaultEnvPath(uid string, loginDefsPath string) string {
envPath := defaultEnvPath
envRootPath := defaultEnvRootPath
// open file, if it doesn't exist return a default path and move on
f, err := utils.OpenFileAllowingUnsafeLinks(loginDefsPath)
if err != nil {
if uid == "0" {
slog.DebugContext(context.Background(), "Unable to open login.defs, returning default su path", "login_defs_path", loginDefsPath, "error", err, "default_path", defaultEnvRootPath)
return defaultEnvRootPath
}
slog.DebugContext(context.Background(), "Unable to open login.defs, returning default path", "login_defs_path", loginDefsPath, "error", err, "default_path", defaultEnvPath)
return defaultEnvPath
}
defer f.Close()
// read path from login.defs file (/etc/login.defs) line by line:
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
// skip comments and empty lines:
if line == "" || line[0] == '#' {
continue
}
// look for a line that starts with ENV_PATH or ENV_SUPATH
fields := strings.Fields(line)
if len(fields) > 1 {
if fields[0] == "ENV_PATH" {
envPath = fields[1]
}
if fields[0] == "ENV_SUPATH" {
envRootPath = fields[1]
}
}
}
// if any error occurs while reading the file, return the default value
err = scanner.Err()
if err != nil {
if uid == "0" {
slog.WarnContext(context.Background(), "Unable to read login.defs, returning default su path", "login_defs_path", loginDefsPath, "error", err, "default_path", defaultEnvRootPath)
return defaultEnvRootPath
}
slog.WarnContext(context.Background(), "Unable to read login.defs, returning default path", "login_defs_path", loginDefsPath, "error", err, "default_path", defaultEnvPath)
return defaultEnvPath
}
// if requesting path for uid 0 and no ENV_SUPATH is given, fallback to
// ENV_PATH first, then the default path.
if uid == "0" {
return envRootPath
}
return envPath
}
// parseSecureCopy will parse a command and return if it's secure copy or not.
func parseSecureCopy(path string) (string, string, bool, error) {
parts := strings.Fields(path)
@@ -687,7 +620,7 @@ func parseSecureCopy(path string) (string, string, bool, error) {
func exitCode(err error) int {
// If no error occurred, return 0 (success).
if err == nil {
return teleport.RemoteCommandSuccess
return reexecconstants.RemoteCommandSuccess
}
var execExitErr *exec.ExitError
@@ -697,7 +630,7 @@ func exitCode(err error) int {
case errors.As(err, &execExitErr):
waitStatus, ok := execExitErr.Sys().(syscall.WaitStatus)
if !ok {
return teleport.RemoteCommandFailure
return reexecconstants.RemoteCommandFailure
}
return waitStatus.ExitStatus()
// Remote execution.
@@ -706,6 +639,98 @@ func exitCode(err error) int {
// An error occurred, but the type is unknown, return a generic 255 code.
default:
slog.DebugContext(context.Background(), "Unknown error returned when executing command", "error", err)
return teleport.RemoteCommandFailure
return reexecconstants.RemoteCommandFailure
}
}
// ConfigureCommand creates a command fully configured to execute. This
// function is used by Teleport to re-execute itself and pass whatever data
// is need to the child to actually execute the shell.
func ConfigureCommand(ctx *ServerContext, extraFiles ...*os.File) (*exec.Cmd, error) {
// Create a os.Pipe and start copying over the payload to execute. While the
// pipe buffer is quite large (64k) some users have run into the pipe
// blocking writes on much smaller buffers (7k) leading to Teleport being
// unable to run some exec commands.
//
// To not depend on the OS implementation of a pipe, instead the copy should
// be non-blocking. The io.Copy will be closed when either when the child
// process has fully read in the payload or the process exits with an error
// (and closes all child file descriptors).
//
// See the below for details.
//
// https://man7.org/linux/man-pages/man7/pipe.7.html
cmdmsg, err := ctx.ExecCommand()
if err != nil {
return nil, trace.Wrap(err)
}
go copyCommand(ctx.CancelContext(), ctx.cmdw, cmdmsg)
// Find the Teleport executable and its directory on disk.
executable, err := os.Executable()
if err != nil {
return nil, trace.Wrap(err)
}
// The channel/request type determines the subcommand to execute.
var subCommand string
switch ctx.ExecType {
case reexecconstants.NetworkingSubCommand:
subCommand = reexecconstants.NetworkingSubCommand
default:
subCommand = reexecconstants.ExecSubCommand
}
// Build the list of arguments to have Teleport re-exec itself. The "-d" flag
// is appended if Teleport is running in debug mode.
args := []string{executable, subCommand}
// build env for `teleport exec`
env := &envutils.SafeEnv{}
env.AddExecEnvironment()
// Build the "teleport exec" command.
cmd := &exec.Cmd{
Path: executable,
Args: args,
Env: *env,
ExtraFiles: []*os.File{
ctx.cmdr,
ctx.logw,
ctx.contr,
ctx.readyw,
ctx.killShellr,
},
}
// Add extra files if applicable.
if len(extraFiles) > 0 {
cmd.ExtraFiles = append(cmd.ExtraFiles, extraFiles...)
}
// Perform OS-specific tweaks to the command.
sessionreexec.CommandOSTweaks(cmd)
return cmd, nil
}
// copyCommand will copy the provided command to the child process over the
// pipe attached to the context.
func copyCommand(ctx context.Context, cmdw *os.File, cmdmsg *sessionreexec.ExecCommand) {
defer func() {
err := cmdw.Close()
if err != nil {
slog.ErrorContext(ctx, "Failed to close command pipe", "error", err)
}
// Set to nil so the close in the context doesn't attempt to re-close.
cmdw = nil
}()
// Write command bytes to pipe. The child process will read the command
// to execute from this pipe.
if err := json.NewEncoder(cmdw).Encode(cmdmsg); err != nil {
slog.ErrorContext(ctx, "Failed to copy command over pipe", "error", err)
return
}
}
+26 -14
View File
@@ -31,11 +31,13 @@ import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
"github.com/gravitational/teleport/lib/utils/testutils"
"github.com/gravitational/teleport/session/host"
"github.com/gravitational/teleport/session/reexec"
)
func TestOSCommandPrep(t *testing.T) {
@@ -74,7 +76,7 @@ func TestOSCommandPrep(t *testing.T) {
require.NoError(t, os.Chown(tempHome, uid, -1))
expectedEnv := []string{
"LANG=en_US.UTF-8",
getDefaultEnvPath(usr.Uid, defaultLoginDefsPath),
reexec.GetDefaultEnvPath(usr.Uid),
fmt.Sprintf("HOME=%s", usr.HomeDir),
fmt.Sprintf("USER=%s", username),
"SHELL=/bin/sh",
@@ -92,11 +94,11 @@ func TestOSCommandPrep(t *testing.T) {
// Empty command (simple shell).
execCmd, err := scx.ExecCommand()
require.NoError(t, err)
execCmd.stdin = os.Stdin
execCmd.stdout = os.Stdout
execCmd.stderr = os.Stderr
execCmd.Stdin = os.Stdin
execCmd.Stdout = os.Stdout
execCmd.Stderr = os.Stderr
cmd, err := buildCommand(execCmd, usr, nil)
cmd, err := reexec.BuildCommand(execCmd, usr, nil)
require.NoError(t, err)
require.NotNil(t, cmd)
@@ -110,11 +112,11 @@ func TestOSCommandPrep(t *testing.T) {
scx.execRequest.SetCommand("ls -lh /etc")
execCmd, err = scx.ExecCommand()
require.NoError(t, err)
execCmd.stdin = os.Stdin
execCmd.stdout = os.Stdout
execCmd.stderr = os.Stderr
execCmd.Stdin = os.Stdin
execCmd.Stdout = os.Stdout
execCmd.Stderr = os.Stderr
cmd, err = buildCommand(execCmd, usr, nil)
cmd, err = reexec.BuildCommand(execCmd, usr, nil)
require.NoError(t, err)
require.NotNil(t, cmd)
@@ -128,11 +130,11 @@ func TestOSCommandPrep(t *testing.T) {
scx.execRequest.SetCommand("top")
execCmd, err = scx.ExecCommand()
require.NoError(t, err)
execCmd.stdin = os.Stdin
execCmd.stdout = os.Stdout
execCmd.stderr = os.Stderr
execCmd.Stdin = os.Stdin
execCmd.Stdout = os.Stdout
execCmd.Stderr = os.Stderr
cmd, err = buildCommand(execCmd, usr, nil)
cmd, err = reexec.BuildCommand(execCmd, usr, nil)
require.NoError(t, err)
require.Equal(t, "/bin/sh", cmd.Path)
@@ -145,7 +147,7 @@ func TestOSCommandPrep(t *testing.T) {
usr.HomeDir = "/wrong/place"
root := string(os.PathSeparator)
expectedEnv[2] = "HOME=/wrong/place"
cmd, err = buildCommand(execCmd, usr, nil)
cmd, err = reexec.BuildCommand(execCmd, usr, nil)
require.NoError(t, err)
require.Equal(t, root, cmd.Dir)
@@ -214,3 +216,13 @@ func TestContinue(t *testing.T) {
require.NoError(t, err)
}
}
func changeHomeDir(t *testing.T, username, home string) {
usermodBin, err := exec.LookPath("usermod")
assert.NoError(t, err, "usermod binary must be present")
cmd := exec.Command(usermodBin, "--home", home, username)
_, err = cmd.CombinedOutput()
assert.NoError(t, err, "changing home should not error")
assert.Equal(t, 0, cmd.ProcessState.ExitCode(), "changing home should exit 0")
}
+7 -17
View File
@@ -28,12 +28,13 @@ import (
"github.com/stretchr/testify/require"
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/sshutils"
"github.com/gravitational/teleport/lib/utils/log/logtest"
"github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
// TestMain will re-execute Teleport to run a command if "exec" is passed to
@@ -43,8 +44,8 @@ func TestMain(m *testing.M) {
modules.SetInsecureTestMode(true)
// If the test is re-executing itself, execute the command that comes over
// the pipe.
if IsReexec() {
RunAndExit(os.Args[1])
if reexec.IsReexec() {
reexec.RunAndExit(os.Args[1])
return
}
@@ -91,21 +92,21 @@ func TestEmitExecAuditEvent(t *testing.T) {
inCommand: "exit 0",
inError: nil,
outCommand: "exit 0",
outCode: strconv.Itoa(teleport.RemoteCommandSuccess),
outCode: strconv.Itoa(reexecconstants.RemoteCommandSuccess),
},
// Exited with error.
{
inCommand: "exit 255",
inError: fmt.Errorf("unknown error"),
outCommand: "exit 255",
outCode: strconv.Itoa(teleport.RemoteCommandFailure),
outCode: strconv.Itoa(reexecconstants.RemoteCommandFailure),
},
// Command injection.
{
inCommand: "/bin/teleport scp --remote-addr=127.0.0.1:50862 --local-addr=127.0.0.1:54895 -f ~/file.txt && touch /tmp/new.txt",
inError: fmt.Errorf("unknown error"),
outCommand: "/bin/teleport scp --remote-addr=127.0.0.1:50862 --local-addr=127.0.0.1:54895 -f ~/file.txt && touch /tmp/new.txt",
outCode: strconv.Itoa(teleport.RemoteCommandFailure),
outCode: strconv.Itoa(reexecconstants.RemoteCommandFailure),
},
}
for _, tt := range tests {
@@ -125,17 +126,6 @@ func TestEmitExecAuditEvent(t *testing.T) {
}
}
func TestLoginDefsParser(t *testing.T) {
t.Parallel()
expectedEnvSuPath := "PATH=/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin:/bar"
expectedSuPath := "PATH=/usr/local/bin:/usr/bin:/bin:/foo"
require.Equal(t, expectedEnvSuPath, getDefaultEnvPath("0", "../../fixtures/login.defs"))
require.Equal(t, expectedSuPath, getDefaultEnvPath("1000", "../../fixtures/login.defs"))
require.Equal(t, defaultEnvPath, getDefaultEnvPath("1000", "bad/file"))
}
func newExecServerContext(t *testing.T, srv Server) *ServerContext {
scx := newTestServerContext(t, srv, nil, nil)
+2 -1
View File
@@ -60,6 +60,7 @@ import (
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/session/networking/x11"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec"
)
// Server is a forwarding server. Server is used to create a single in-memory
@@ -568,7 +569,7 @@ func (s *Server) GetLockWatcher() *services.LockWatcher {
// does not spawn child processes.
func (s *Server) ChildLogConfig() srv.ChildLogConfig {
return srv.ChildLogConfig{
ExecLogConfig: srv.ExecLogConfig{},
ExecLogConfig: reexec.ExecLogConfig{},
Writer: io.Discard,
}
}
+2 -1
View File
@@ -48,6 +48,7 @@ import (
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec"
)
var (
@@ -765,7 +766,7 @@ func (s *ForwardServer) GetSELinuxEnabled() bool {
// does not spawn child processes.
func (s *ForwardServer) ChildLogConfig() srv.ChildLogConfig {
return srv.ChildLogConfig{
ExecLogConfig: srv.ExecLogConfig{},
ExecLogConfig: reexec.ExecLogConfig{},
Writer: io.Discard,
}
}
+2 -1
View File
@@ -55,6 +55,7 @@ import (
"github.com/gravitational/teleport/lib/utils/clocki"
"github.com/gravitational/teleport/lib/utils/log/logtest"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec"
)
func newTestServerContext(t *testing.T, srv Server, sessionJoiningRoleSet services.RoleSet, accessPermit *decisionpb.SSHAccessPermit) *ServerContext {
@@ -348,7 +349,7 @@ func (m *mockServer) GetSELinuxEnabled() bool {
// ChildLogConfig returns a noop log configuration.
func (m *mockServer) ChildLogConfig() ChildLogConfig {
return ChildLogConfig{
ExecLogConfig: ExecLogConfig{
ExecLogConfig: reexec.ExecLogConfig{
Level: slog.LevelDebug,
Format: "json",
},
+26
View File
@@ -0,0 +1,26 @@
// Teleport
// Copyright (C) 2026 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package srv
import (
"github.com/gravitational/teleport/session/reexec"
)
//go:fix inline
func IsReexec() bool {
return reexec.IsReexec()
}
+2 -296
View File
@@ -29,168 +29,23 @@ import (
"net/http"
"net/http/httptest"
"os"
"os/exec"
"os/user"
"path/filepath"
"strconv"
"syscall"
"testing"
"github.com/gravitational/trace"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/crypto/ssh/agent"
"github.com/gravitational/teleport"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/sshagent"
"github.com/gravitational/teleport/lib/utils/testutils"
"github.com/gravitational/teleport/session/host"
"github.com/gravitational/teleport/session/networking"
"github.com/gravitational/teleport/session/networking/x11"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
type stubUser struct {
gid string
uid string
groupIDS []string
}
func (s *stubUser) GID() string {
return s.gid
}
func (s *stubUser) UID() string {
return s.uid
}
func (s *stubUser) GroupIds() ([]string, error) {
return s.groupIDS, nil
}
func TestStartNewParker(t *testing.T) {
currentUser, err := user.Current()
require.NoError(t, err)
currentUID, err := strconv.ParseUint(currentUser.Uid, 10, 32)
require.NoError(t, err)
currentGID, err := strconv.ParseUint(currentUser.Gid, 10, 32)
require.NoError(t, err)
t.Parallel()
type args struct {
credential *syscall.Credential
loginAsUser string
localUser *stubUser
}
tests := []struct {
name string
args args
newOsPack func(t *testing.T) (*osWrapper, func())
wantErr require.ErrorAssertionFunc
}{
{
name: "empty credentials does nothing",
wantErr: require.NoError,
newOsPack: func(t *testing.T) (*osWrapper, func()) {
return &osWrapper{}, func() {}
},
},
{
name: "missing Teleport group returns no error",
wantErr: require.NoError,
newOsPack: func(t *testing.T) (*osWrapper, func()) {
return &osWrapper{
LookupGroup: func(name string) (*user.Group, error) {
require.Equal(t, types.TeleportDropGroup, name)
return nil, user.UnknownGroupError(types.TeleportDropGroup)
},
}, func() {}
},
},
{
name: "different group doesn't start parker",
wantErr: require.NoError,
newOsPack: func(t *testing.T) (*osWrapper, func()) {
return &osWrapper{
LookupGroup: func(name string) (*user.Group, error) {
require.Equal(t, types.TeleportDropGroup, name)
return &user.Group{Gid: "1234"}, nil
},
CommandContext: func(ctx context.Context, name string, arg ...string) *exec.Cmd {
require.FailNow(t, "CommandContext should not be called")
return nil
},
}, func() {}
},
args: args{
credential: &syscall.Credential{Gid: 1000},
localUser: &stubUser{
uid: "1001",
gid: "1003",
groupIDS: []string{"1003"},
},
},
},
{
name: "parker is started",
wantErr: require.NoError,
newOsPack: func(t *testing.T) (*osWrapper, func()) {
parkerStarted := false
return &osWrapper{
LookupGroup: func(name string) (*user.Group, error) {
require.Equal(t, types.TeleportDropGroup, name)
return &user.Group{Gid: currentUser.Gid}, nil
},
CommandContext: func(ctx context.Context, name string, arg ...string) *exec.Cmd {
require.NotNil(t, ctx)
require.Len(t, arg, 1)
require.Equal(t, teleport.ParkSubCommand, arg[0])
parkerStarted = true
return exec.CommandContext(ctx, name, arg...)
},
LookupUser: func(username string) (*user.User, error) {
return &user.User{Uid: currentUser.Uid, Gid: currentUser.Gid}, nil
},
}, func() {
require.True(t, parkerStarted, "parker process didn't start")
}
},
args: args{
credential: &syscall.Credential{
Uid: uint32(currentUID),
Gid: uint32(currentGID),
// Changing to false causes "fork/exec /proc/self/exe: operation not permitted"
// to be returned when creating the park process.
NoSetGroups: true,
},
localUser: &stubUser{
uid: currentUser.Uid,
gid: currentUser.Gid,
groupIDS: []string{currentUser.Gid},
},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
osPack, assertExpected := tt.newOsPack(t)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel) // cancel to stop the park process.
err := osPack.startNewParker(ctx, tt.args.credential, tt.args.loginAsUser, tt.args.localUser)
tt.wantErr(t, err, fmt.Sprintf("startNewParker(%v, %+v, %v, %+v)", ctx, tt.args.credential, tt.args.loginAsUser, tt.args.localUser))
assertExpected()
})
}
}
func newHTTPTestServer(t *testing.T, listener net.Listener) *httptest.Server {
var err error
if listener == nil {
@@ -235,7 +90,7 @@ func testNetworkingCommand(t *testing.T, login string) {
srv := newMockServer(t)
scx := newTestServerContext(t, srv, nil, &decisionpb.SSHAccessPermit{})
scx.ExecType = teleport.NetworkingSubCommand
scx.ExecType = reexecconstants.NetworkingSubCommand
if login != "" {
scx.Identity.Login = login
}
@@ -412,152 +267,3 @@ func testX11Forward(ctx context.Context, t *testing.T, proc *networking.Process,
require.NoError(t, err)
require.Equal(t, fakeXauthEntry, readXauthEntry)
}
func TestRootCheckHomeDir(t *testing.T) {
testutils.RequireRoot(t)
// this test manipulates global state, ensure we're not going to run it in
// parallel with something else
t.Setenv("foo", "bar")
tmp := t.TempDir()
require.NoError(t, os.Chmod(filepath.Dir(tmp), 0777))
require.NoError(t, os.Chmod(tmp, 0777))
home := filepath.Join(tmp, "home")
noAccess := filepath.Join(tmp, "no_access")
file := filepath.Join(tmp, "file")
notFound := filepath.Join(tmp, "not_found")
require.NoError(t, os.Mkdir(home, 0700))
require.NoError(t, os.Mkdir(noAccess, 0700))
_, err := os.Create(file)
require.NoError(t, err)
login := testutils.GenerateLocalUsername(t)
_, err = host.UserAdd(login, nil, host.UserOpts{Home: home})
require.NoError(t, err)
t.Cleanup(func() {
// change back to accessible home so deletion works
changeHomeDir(t, login, home)
_, err := host.UserDel(login)
require.NoError(t, err)
})
testUser, err := user.Lookup(login)
require.NoError(t, err)
uid, err := strconv.Atoi(testUser.Uid)
require.NoError(t, err)
gid, err := strconv.Atoi(testUser.Gid)
require.NoError(t, err)
require.NoError(t, os.Chown(home, uid, gid))
require.NoError(t, os.Chown(file, uid, gid))
hasAccess, err := CheckHomeDir(testUser)
require.NoError(t, err)
require.True(t, hasAccess)
changeHomeDir(t, login, file)
hasAccess, err = CheckHomeDir(testUser)
require.NoError(t, err)
require.False(t, hasAccess)
changeHomeDir(t, login, notFound)
hasAccess, err = CheckHomeDir(testUser)
require.NoError(t, err)
require.False(t, hasAccess)
changeHomeDir(t, login, noAccess)
hasAccess, err = CheckHomeDir(testUser)
require.NoError(t, err)
require.False(t, hasAccess)
}
func changeHomeDir(t *testing.T, username, home string) {
usermodBin, err := exec.LookPath("usermod")
assert.NoError(t, err, "usermod binary must be present")
cmd := exec.Command(usermodBin, "--home", home, username)
_, err = cmd.CombinedOutput()
assert.NoError(t, err, "changing home should not error")
assert.Equal(t, 0, cmd.ProcessState.ExitCode(), "changing home should exit 0")
}
func TestRootOpenFileAsUser(t *testing.T) {
testutils.RequireRoot(t)
euid := os.Geteuid()
egid := os.Getegid()
username := "processing-user"
arg := os.Args[1]
os.Args[1] = teleport.ExecSubCommand
defer func() {
os.Args[1] = arg
}()
_, err := host.UserAdd(username, nil, host.UserOpts{})
require.NoError(t, err)
t.Cleanup(func() {
_, err := host.UserDel(username)
require.NoError(t, err)
})
tmp := t.TempDir()
testFile := filepath.Join(tmp, "testfile")
fileContent := "one does not simply open without permission"
err = os.WriteFile(testFile, []byte(fileContent), 0777)
require.NoError(t, err)
testUser, err := user.Lookup(username)
require.NoError(t, err)
// no access
file, err := openFileAsUser(testUser, testFile)
require.True(t, trace.IsAccessDenied(err))
require.Nil(t, file)
// ensure we fallback to root after
file, err = os.Open(testFile)
require.NoError(t, err)
require.NotNil(t, file)
file.Close()
// has access
uid, err := strconv.Atoi(testUser.Uid)
require.NoError(t, err)
gid, err := strconv.Atoi(testUser.Gid)
require.NoError(t, err)
err = os.Chown(filepath.Dir(tmp), uid, gid)
require.NoError(t, err)
err = os.Chown(tmp, uid, gid)
require.NoError(t, err)
err = os.Chown(testFile, uid, gid)
require.NoError(t, err)
file, err = openFileAsUser(testUser, testFile)
require.NoError(t, err)
require.NotNil(t, file)
data, err := io.ReadAll(file)
file.Close()
require.NoError(t, err)
require.Equal(t, fileContent, string(data))
// not exist
file, err = openFileAsUser(testUser, filepath.Join(tmp, "no_exist"))
require.ErrorIs(t, err, os.ErrNotExist)
require.Nil(t, file)
require.Equal(t, euid, os.Geteuid())
require.Equal(t, egid, os.Getegid())
}
+2 -1
View File
@@ -29,6 +29,7 @@ import (
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/session/reexec"
)
type homeDirSubsys struct {
@@ -50,7 +51,7 @@ func (h *homeDirSubsys) Start(_ context.Context, serverConn *ssh.ServerConn, ch
return trace.Wrap(err)
}
exists, err := srv.CheckHomeDir(localUser)
exists, err := reexec.CheckHomeDir(localUser)
if err != nil {
return trace.Wrap(err)
}
+2 -1
View File
@@ -41,6 +41,7 @@ import (
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/sshutils/reexec"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
type sftpSubsys struct {
@@ -108,7 +109,7 @@ func (s *sftpSubsys) Start(ctx context.Context,
defer auditPipeIn.Close()
// Create child process to handle SFTP connection
execRequest, err := srv.NewExecRequest(serverCtx, teleport.SFTPSubCommand)
execRequest, err := srv.NewExecRequest(serverCtx, reexecconstants.SFTPSubCommand)
if err != nil {
return trace.Wrap(err)
}
+5 -3
View File
@@ -76,6 +76,8 @@ import (
"github.com/gravitational/teleport/session/networking"
"github.com/gravitational/teleport/session/networking/x11"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
// Server implements SSH server that uses configuration backend and
@@ -370,7 +372,7 @@ func (s *Server) ChildLogConfig() srv.ChildLogConfig {
// return a noop log configuration
return srv.ChildLogConfig{
ExecLogConfig: srv.ExecLogConfig{},
ExecLogConfig: reexec.ExecLogConfig{},
Writer: io.Discard,
}
}
@@ -824,7 +826,7 @@ func SetImmutableLabels(labels map[string]string) ServerOption {
func SetChildLogConfig(cfg *servicecfg.Config) ServerOption {
return func(s *Server) error {
s.childLogConfig = &srv.ChildLogConfig{
ExecLogConfig: srv.ExecLogConfig{
ExecLogConfig: reexec.ExecLogConfig{
Level: cfg.LoggerLevel.Level(),
Format: strings.ToLower(cfg.LogConfig.Format),
ExtraFields: cfg.LogConfig.ExtraFields,
@@ -1303,7 +1305,7 @@ func (s *Server) startNetworkingProcess(scx *srv.ServerContext) (*networking.Pro
return nil, trace.Wrap(err)
}
nsctx.SessionRecordingConfig.SetMode(types.RecordOff)
nsctx.ExecType = teleport.NetworkingSubCommand
nsctx.ExecType = reexecconstants.NetworkingSubCommand
scx.Parent().AddCloser(nsctx)
// Create command to re-exec Teleport which will handle networking requests. The
+4 -3
View File
@@ -86,6 +86,7 @@ import (
"github.com/gravitational/teleport/lib/utils/log/logtest"
"github.com/gravitational/teleport/session/networking/x11"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec"
)
// teleportTestUser is additional user used for tests
@@ -101,8 +102,8 @@ var wildcardAllow = types.Labels{
func TestMain(m *testing.M) {
logtest.InitLogger(testing.Verbose)
modules.SetInsecureTestMode(true)
if srv.IsReexec() {
srv.RunAndExit(os.Args[1])
if reexec.IsReexec() {
reexec.RunAndExit(os.Args[1])
return
}
@@ -198,7 +199,7 @@ func (f *sshTestFixture) newSSHClient(ctx context.Context, t testing.TB, user *u
func setChildLogConfigForTest() ServerOption {
return func(s *Server) error {
s.childLogConfig = &srv.ChildLogConfig{
ExecLogConfig: srv.ExecLogConfig{
ExecLogConfig: reexec.ExecLogConfig{
Level: slog.LevelDebug,
},
Writer: os.Stderr,
+2 -1
View File
@@ -36,6 +36,7 @@ import (
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/utils/testutils"
"github.com/gravitational/teleport/session/host"
"github.com/gravitational/teleport/session/reexec"
)
// BenchmarkRootExecCommand measures performance of running multiple exec requests
@@ -70,7 +71,7 @@ func BenchmarkRootExecCommand(b *testing.B) {
// Re-enable once we have a better way to benchmark with logging enabled.
func(s *Server) error {
s.childLogConfig = &srv.ChildLogConfig{
ExecLogConfig: srv.ExecLogConfig{
ExecLogConfig: reexec.ExecLogConfig{
Level: slog.LevelError,
},
Writer: io.Discard,
+10 -4
View File
@@ -41,6 +41,11 @@ import (
tracessh "github.com/gravitational/teleport/api/observability/tracing/ssh"
rsession "github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/sshutils/reexec"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
const (
defaultTerm = "xterm"
)
// LookupUser is used to mock the value returned by user.Lookup(string).
@@ -197,8 +202,9 @@ func (t *terminal) AddParty(delta int) {
// Replace \n with \r\n so the message is correctly aligned.
var crlfReplacer = strings.NewReplacer("\r\n", "\r\n", "\n", "\r\n")
// Run will run the terminal. If the shell fails to start due to a [teleport.RemoteCommandFailure],
// the error will be written to the given error writer.
// Run will run the terminal. If the shell fails to start due to a
// [reexecconstants.RemoteCommandFailure], the error will be written to the
// given error writer.
func (t *terminal) Run(ctx context.Context, errorWriter io.Writer) error {
select {
case <-ctx.Done():
@@ -651,13 +657,13 @@ func (t *remoteTerminal) Wait() (*ExecResult, error) {
}
return &ExecResult{
Code: teleport.RemoteCommandFailure,
Code: reexecconstants.RemoteCommandFailure,
Command: execRequest.GetCommand(),
}, err
}
return &ExecResult{
Code: teleport.RemoteCommandSuccess,
Code: reexecconstants.RemoteCommandSuccess,
Command: execRequest.GetCommand(),
}, nil
}
+17 -16
View File
@@ -36,6 +36,7 @@ import (
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
apiconstants "github.com/gravitational/teleport/api/constants"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
@@ -157,7 +158,7 @@ type userCloser struct {
}
func (u *userCloser) Close() error {
teleportGroup, err := u.backend.LookupGroup(types.TeleportDropGroup)
teleportGroup, err := u.backend.LookupGroup(apiconstants.TeleportDropGroup)
if err != nil {
return trace.Wrap(err)
}
@@ -301,7 +302,7 @@ func (u *HostUserManagement) updateUser(hostUser HostUser, ui *decisionpb.HostUs
)
if ui.Mode == decisionpb.HostUserMode_HOST_USER_MODE_KEEP {
_, hasKeepGroup := hostUser.Groups[types.TeleportKeepGroup]
_, hasKeepGroup := hostUser.Groups[apiconstants.TeleportKeepGroup]
if !hasKeepGroup {
home, err := u.backend.GetDefaultHomeDirectory(hostUser.Name)
if err != nil {
@@ -580,17 +581,17 @@ func isUnknownGroupError(err error, groupName string) bool {
strings.HasSuffix(err.Error(), syscall.ESRCH.Error())
}
// DeleteAllUsers removes all temporary users in the [types.TeleportDropGroup]
// DeleteAllUsers removes all temporary users in the [apiconstants.TeleportDropGroup]
// without any active sessions.
func (u *HostUserManagement) DeleteAllUsers() error {
users, err := u.backend.GetAllUsers()
if err != nil {
return trace.Wrap(err)
}
teleportGroup, err := u.backend.LookupGroup(types.TeleportDropGroup)
teleportGroup, err := u.backend.LookupGroup(apiconstants.TeleportDropGroup)
if err != nil {
if isUnknownGroupError(err, types.TeleportDropGroup) {
u.log.DebugContext(u.ctx, "Target group not found, not deleting users", "group", types.TeleportDropGroup)
if isUnknownGroupError(err, apiconstants.TeleportDropGroup) {
u.log.DebugContext(u.ctx, "Target group not found, not deleting users", "group", apiconstants.TeleportDropGroup)
return nil
}
return trace.Wrap(err)
@@ -747,23 +748,23 @@ func ResolveGroups(logger *slog.Logger, hostUser *HostUser, ui *decisionpb.HostU
}
// because teleport-keep migration requires adding the group to host_groups, we need to note that before wiping the teleport system groups
_, hasExplicitKeepGroup := groups[types.TeleportKeepGroup]
_, hasExplicitKeepGroup := groups[apiconstants.TeleportKeepGroup]
// only one teleport system group should be resolved for a given user, so we remove any of them that might occur within the configured host
// groups since we'll compute the correct group below
delete(groups, types.TeleportKeepGroup)
delete(groups, types.TeleportDropGroup)
delete(groups, types.TeleportStaticGroup)
delete(groups, apiconstants.TeleportKeepGroup)
delete(groups, apiconstants.TeleportDropGroup)
delete(groups, apiconstants.TeleportStaticGroup)
// if we assign a teleport group, it will always coincide with the mode we're currently in, so we can compute it right away
teleportGroup := ""
switch ui.Mode {
case decisionpb.HostUserMode_HOST_USER_MODE_DROP:
teleportGroup = types.TeleportDropGroup
teleportGroup = apiconstants.TeleportDropGroup
case decisionpb.HostUserMode_HOST_USER_MODE_KEEP:
teleportGroup = types.TeleportKeepGroup
teleportGroup = apiconstants.TeleportKeepGroup
case decisionpb.HostUserMode_HOST_USER_MODE_STATIC:
teleportGroup = types.TeleportStaticGroup
teleportGroup = apiconstants.TeleportStaticGroup
}
log := logger.With("teleport_group", teleportGroup)
@@ -774,14 +775,14 @@ func ResolveGroups(logger *slog.Logger, hostUser *HostUser, ui *decisionpb.HostU
// 2. We reconcile an existing managed user
// 3. We migrate an existing unmanaged user
// functionally, there's no difference between 2 and 3 so if we check against all failure modes we can handle all other cases at once
_, hasDropGroup := hostUser.Groups[types.TeleportDropGroup]
_, hasKeepGroup := hostUser.Groups[types.TeleportKeepGroup]
_, hasDropGroup := hostUser.Groups[apiconstants.TeleportDropGroup]
_, hasKeepGroup := hostUser.Groups[apiconstants.TeleportKeepGroup]
migrateStaticUser := takeOwnership && ui.Mode == decisionpb.HostUserMode_HOST_USER_MODE_STATIC
migrateKeepUser := hasExplicitKeepGroup && ui.Mode == decisionpb.HostUserMode_HOST_USER_MODE_KEEP
managedUser := hasKeepGroup || hasDropGroup
_, staticUser := hostUser.Groups[types.TeleportStaticGroup]
_, staticUser := hostUser.Groups[apiconstants.TeleportStaticGroup]
inStaticMode := ui.Mode == decisionpb.HostUserMode_HOST_USER_MODE_STATIC
if (inStaticMode && managedUser) || (!inStaticMode && staticUser) {
+86 -86
View File
@@ -30,8 +30,8 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
apiconstants "github.com/gravitational/teleport/api/constants"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/backend/memory"
"github.com/gravitational/teleport/lib/services/local"
"github.com/gravitational/teleport/lib/utils/log/logtest"
@@ -248,7 +248,7 @@ func TestUserMgmt_CreateTemporaryUser(t *testing.T) {
// temporary users must always include the teleport-service group
require.ElementsMatch(t, []string{
"hello", "sudo", types.TeleportDropGroup,
"hello", "sudo", apiconstants.TeleportDropGroup,
}, backend.users["bob"])
// try create the same user again
@@ -315,7 +315,7 @@ func TestUserMgmtSudoers_CreateTemporaryUser(t *testing.T) {
Mode: decisionpb.HostUserMode_HOST_USER_MODE_DROP,
})
require.ErrorIs(t, err, errUnmanagedUser)
backend.CreateGroup(types.TeleportDropGroup, "")
backend.CreateGroup(apiconstants.TeleportDropGroup, "")
_, err = users.UpsertUser("testuser", &decisionpb.HostUsersInfo{
Mode: decisionpb.HostUserMode_HOST_USER_MODE_DROP,
})
@@ -332,8 +332,8 @@ func TestUserMgmt_DeleteAllTeleportSystemUsers(t *testing.T) {
}
usersDB := []userAndGroups{
{"fgh", []string{types.TeleportDropGroup}},
{"xyz", []string{types.TeleportDropGroup}},
{"fgh", []string{apiconstants.TeleportDropGroup}},
{"xyz", []string{apiconstants.TeleportDropGroup}},
{"pqr", []string{"not-deleted"}},
{"abc", []string{"not-deleted"}},
}
@@ -355,7 +355,7 @@ func TestUserMgmt_DeleteAllTeleportSystemUsers(t *testing.T) {
for _, group := range user.groups {
mgmt.CreateGroup(group, "")
}
if slices.Contains(user.groups, types.TeleportDropGroup) {
if slices.Contains(user.groups, apiconstants.TeleportDropGroup) {
users.UpsertUser(user.user, &decisionpb.HostUsersInfo{
Groups: user.groups,
Mode: decisionpb.HostUserMode_HOST_USER_MODE_DROP,
@@ -442,8 +442,8 @@ func Test_UpdateUserGroups_Keep(t *testing.T) {
assert.NoError(t, err)
assert.Equal(t, nil, closer)
assert.Zero(t, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportDropGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportDropGroup)
// Update user with new groups.
userinfo.Groups = slices.Clone(allGroups[2:])
@@ -452,16 +452,16 @@ func Test_UpdateUserGroups_Keep(t *testing.T) {
assert.NoError(t, err)
assert.Equal(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportDropGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportDropGroup)
// Upsert again with same groups should not call UpdateUser.
closer, err = users.UpsertUser("alice", &userinfo)
assert.NoError(t, err)
assert.Equal(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportDropGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportDropGroup)
// Do not convert the managed user to static.
userinfo.Mode = decisionpb.HostUserMode_HOST_USER_MODE_STATIC
@@ -469,7 +469,7 @@ func Test_UpdateUserGroups_Keep(t *testing.T) {
assert.ErrorIs(t, err, errStaticConversion)
assert.Equal(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportKeepGroup), backend.users["alice"])
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportKeepGroup), backend.users["alice"])
// Updates with INSECURE_DROP mode should convert the managed user
userinfo.Mode = decisionpb.HostUserMode_HOST_USER_MODE_DROP
@@ -478,8 +478,8 @@ func Test_UpdateUserGroups_Keep(t *testing.T) {
assert.NoError(t, err)
assert.NotEqual(t, nil, closer)
assert.Equal(t, 2, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportDropGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportKeepGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportDropGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportKeepGroup)
}
func Test_UpdateUserGroups_Drop(t *testing.T) {
@@ -498,8 +498,8 @@ func Test_UpdateUserGroups_Drop(t *testing.T) {
assert.NoError(t, err)
assert.NotEqual(t, nil, closer)
assert.Zero(t, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportDropGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportKeepGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportDropGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportKeepGroup)
// Update user with new groups.
userinfo.Groups = slices.Clone(allGroups[2:])
@@ -508,16 +508,16 @@ func Test_UpdateUserGroups_Drop(t *testing.T) {
assert.NoError(t, err)
assert.NotEqual(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportDropGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportKeepGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportDropGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportKeepGroup)
// Upsert again with same groups should not call SetUserGroups.
closer, err = users.UpsertUser("alice", &userinfo)
assert.NoError(t, err)
assert.NotEqual(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportDropGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportKeepGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportDropGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportKeepGroup)
// Do not convert the managed user to static.
userinfo.Mode = decisionpb.HostUserMode_HOST_USER_MODE_STATIC
@@ -525,7 +525,7 @@ func Test_UpdateUserGroups_Drop(t *testing.T) {
assert.ErrorIs(t, err, errStaticConversion)
assert.Equal(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportDropGroup), backend.users["alice"])
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportDropGroup), backend.users["alice"])
// Updates with KEEP mode should convert the ephemeral user
userinfo.Mode = decisionpb.HostUserMode_HOST_USER_MODE_KEEP
@@ -535,8 +535,8 @@ func Test_UpdateUserGroups_Drop(t *testing.T) {
assert.Equal(t, nil, closer)
assert.Equal(t, 2, backend.updateUserCalls)
assert.Equal(t, 1, backend.createHomeDirectoryCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportDropGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportDropGroup)
}
func Test_UpdateUserGroups_Static(t *testing.T) {
@@ -554,7 +554,7 @@ func Test_UpdateUserGroups_Static(t *testing.T) {
assert.NoError(t, err)
assert.Equal(t, nil, closer)
assert.Zero(t, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportStaticGroup), backend.users["alice"])
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportStaticGroup), backend.users["alice"])
// Update user with new groups.
userinfo.Groups = slices.Clone(allGroups[2:])
@@ -562,14 +562,14 @@ func Test_UpdateUserGroups_Static(t *testing.T) {
assert.NoError(t, err)
assert.Equal(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportStaticGroup), backend.users["alice"])
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportStaticGroup), backend.users["alice"])
// Upsert again with same groups should not call SetUserGroups.
closer, err = users.UpsertUser("alice", &userinfo)
assert.NoError(t, err)
assert.Equal(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportStaticGroup), backend.users["alice"])
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportStaticGroup), backend.users["alice"])
// Do not convert to KEEP.
userinfo.Mode = decisionpb.HostUserMode_HOST_USER_MODE_KEEP
@@ -577,7 +577,7 @@ func Test_UpdateUserGroups_Static(t *testing.T) {
assert.ErrorIs(t, err, errStaticConversion)
assert.Equal(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(slices.Clone(allGroups[2:]), types.TeleportStaticGroup), backend.users["alice"])
assert.ElementsMatch(t, append(slices.Clone(allGroups[2:]), apiconstants.TeleportStaticGroup), backend.users["alice"])
// Do not convert to INSECURE_DROP.
userinfo.Mode = decisionpb.HostUserMode_HOST_USER_MODE_DROP
@@ -585,7 +585,7 @@ func Test_UpdateUserGroups_Static(t *testing.T) {
assert.ErrorIs(t, err, errStaticConversion)
assert.Equal(t, nil, closer)
assert.Equal(t, 1, backend.updateUserCalls)
assert.ElementsMatch(t, append(slices.Clone(allGroups[2:]), types.TeleportStaticGroup), backend.users["alice"])
assert.ElementsMatch(t, append(slices.Clone(allGroups[2:]), apiconstants.TeleportStaticGroup), backend.users["alice"])
}
func Test_DontManageExistingUser(t *testing.T) {
@@ -674,7 +674,7 @@ func Test_DontUpdateUnmanagedUsers(t *testing.T) {
func Test_AllowExplicitlyManageExistingUsers(t *testing.T) {
t.Parallel()
allGroups := []string{"foo", types.TeleportKeepGroup, types.TeleportDropGroup}
allGroups := []string{"foo", apiconstants.TeleportKeepGroup, apiconstants.TeleportDropGroup}
users, backend := initBackend(t, allGroups)
assert.NoError(t, backend.CreateUser("alice-keep", []string{}, host.UserOpts{}))
@@ -692,7 +692,7 @@ func Test_AllowExplicitlyManageExistingUsers(t *testing.T) {
assert.Equal(t, 1, backend.updateUserCalls)
// slice off the end because teleport-system should be explicitly excluded
assert.ElementsMatch(t, allGroups[:2], backend.users["alice-keep"])
assert.NotContains(t, backend.users["alice-keep"], types.TeleportDropGroup)
assert.NotContains(t, backend.users["alice-keep"], apiconstants.TeleportDropGroup)
// Take ownership of existing user when in STATIC mode
userinfo.Mode = decisionpb.HostUserMode_HOST_USER_MODE_STATIC
@@ -701,9 +701,9 @@ func Test_AllowExplicitlyManageExistingUsers(t *testing.T) {
assert.Equal(t, nil, closer)
assert.Equal(t, 2, backend.updateUserCalls)
assert.Contains(t, backend.users["alice-static"], "foo")
assert.Contains(t, backend.users["alice-static"], types.TeleportStaticGroup)
assert.NotContains(t, backend.users["alice-static"], types.TeleportKeepGroup)
assert.NotContains(t, backend.users["alice-static"], types.TeleportDropGroup)
assert.Contains(t, backend.users["alice-static"], apiconstants.TeleportStaticGroup)
assert.NotContains(t, backend.users["alice-static"], apiconstants.TeleportKeepGroup)
assert.NotContains(t, backend.users["alice-static"], apiconstants.TeleportDropGroup)
// Don't take ownership of existing user when in DROP mode
userinfo.Mode = decisionpb.HostUserMode_HOST_USER_MODE_DROP
@@ -719,8 +719,8 @@ func Test_AllowExplicitlyManageExistingUsers(t *testing.T) {
assert.NoError(t, err)
assert.NotEqual(t, nil, closer)
assert.Equal(t, 2, backend.updateUserCalls)
assert.ElementsMatch(t, []string{"foo", types.TeleportDropGroup}, backend.users["bob"])
assert.NotContains(t, backend.users["bob"], types.TeleportKeepGroup)
assert.ElementsMatch(t, []string{"foo", apiconstants.TeleportDropGroup}, backend.users["bob"])
assert.NotContains(t, backend.users["bob"], apiconstants.TeleportKeepGroup)
}
func initBackend(t *testing.T, groups []string) (HostUserManagement, *testHostUserBackend) {
@@ -811,7 +811,7 @@ func TestHostUsersResolveGroups(t *testing.T) {
Groups: []string{"foo", "bar"},
},
expectGroups: []string{"foo", "bar", types.TeleportDropGroup},
expectGroups: []string{"foo", "bar", apiconstants.TeleportDropGroup},
},
{
name: "create keep user",
@@ -822,7 +822,7 @@ func TestHostUsersResolveGroups(t *testing.T) {
Groups: []string{"foo", "bar"},
},
expectGroups: []string{"foo", "bar", types.TeleportKeepGroup},
expectGroups: []string{"foo", "bar", apiconstants.TeleportKeepGroup},
},
{
name: "create static user",
@@ -833,15 +833,15 @@ func TestHostUsersResolveGroups(t *testing.T) {
Groups: []string{"foo", "bar"},
},
expectGroups: []string{"foo", "bar", types.TeleportStaticGroup},
expectGroups: []string{"foo", "bar", apiconstants.TeleportStaticGroup},
},
{
name: "update drop user",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportDropGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportDropGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -849,16 +849,16 @@ func TestHostUsersResolveGroups(t *testing.T) {
Groups: []string{"baz", "qux"},
},
expectGroups: []string{"baz", "qux", types.TeleportDropGroup},
expectGroups: []string{"baz", "qux", apiconstants.TeleportDropGroup},
},
{
name: "update keep user",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportKeepGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportKeepGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -866,16 +866,16 @@ func TestHostUsersResolveGroups(t *testing.T) {
Groups: []string{"baz", "qux"},
},
expectGroups: []string{"baz", "qux", types.TeleportKeepGroup},
expectGroups: []string{"baz", "qux", apiconstants.TeleportKeepGroup},
},
{
name: "update static user",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportStaticGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportStaticGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -883,16 +883,16 @@ func TestHostUsersResolveGroups(t *testing.T) {
Groups: []string{"baz", "qux"},
},
expectGroups: []string{"baz", "qux", types.TeleportStaticGroup},
expectGroups: []string{"baz", "qux", apiconstants.TeleportStaticGroup},
},
{
name: "convert drop to keep",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportDropGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportDropGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -900,16 +900,16 @@ func TestHostUsersResolveGroups(t *testing.T) {
Groups: []string{"baz", "qux"},
},
expectGroups: []string{"baz", "qux", types.TeleportKeepGroup},
expectGroups: []string{"baz", "qux", apiconstants.TeleportKeepGroup},
},
{
name: "convert keep to drop",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportKeepGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportKeepGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -917,16 +917,16 @@ func TestHostUsersResolveGroups(t *testing.T) {
Groups: []string{"baz", "qux"},
},
expectGroups: []string{"baz", "qux", types.TeleportDropGroup},
expectGroups: []string{"baz", "qux", apiconstants.TeleportDropGroup},
},
{
name: "don't convert drop to static",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportDropGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportDropGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -942,9 +942,9 @@ func TestHostUsersResolveGroups(t *testing.T) {
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportKeepGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportKeepGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -960,9 +960,9 @@ func TestHostUsersResolveGroups(t *testing.T) {
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportStaticGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportStaticGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -977,9 +977,9 @@ func TestHostUsersResolveGroups(t *testing.T) {
name: "don't convert static to drop",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportStaticGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportStaticGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
@@ -1001,7 +1001,7 @@ func TestHostUsersResolveGroups(t *testing.T) {
},
ui: &decisionpb.HostUsersInfo{
Mode: decisionpb.HostUserMode_HOST_USER_MODE_DROP,
Groups: []string{"baz", "qux", types.TeleportDropGroup}, // similarly including TeleportDropGroup to ensure no-op
Groups: []string{"baz", "qux", apiconstants.TeleportDropGroup}, // similarly including TeleportDropGroup to ensure no-op
},
takeOwnership: true, // this flag should be a no-op for DROP, so we include it to ensure that behavior
@@ -1050,10 +1050,10 @@ func TestHostUsersResolveGroups(t *testing.T) {
},
ui: &decisionpb.HostUsersInfo{
Mode: decisionpb.HostUserMode_HOST_USER_MODE_KEEP,
Groups: []string{"baz", "qux", types.TeleportKeepGroup},
Groups: []string{"baz", "qux", apiconstants.TeleportKeepGroup},
},
expectGroups: []string{"baz", "qux", types.TeleportKeepGroup},
expectGroups: []string{"baz", "qux", apiconstants.TeleportKeepGroup},
},
{
name: "take over unmanaged user in static mode when migrating",
@@ -1070,38 +1070,38 @@ func TestHostUsersResolveGroups(t *testing.T) {
takeOwnership: true,
expectGroups: []string{"baz", "qux", types.TeleportStaticGroup},
expectGroups: []string{"baz", "qux", apiconstants.TeleportStaticGroup},
},
{
name: "ignore explicitly configured teleport system groups",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportDropGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportDropGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
Mode: decisionpb.HostUserMode_HOST_USER_MODE_DROP,
Groups: []string{"baz", types.TeleportStaticGroup, types.TeleportKeepGroup, types.TeleportDropGroup},
Groups: []string{"baz", apiconstants.TeleportStaticGroup, apiconstants.TeleportKeepGroup, apiconstants.TeleportDropGroup},
},
expectGroups: []string{"baz", types.TeleportDropGroup},
expectGroups: []string{"baz", apiconstants.TeleportDropGroup},
},
{
name: "return no groups if no change is necessary",
hostUser: &HostUser{
Groups: map[string]struct{}{
"foo": {},
"bar": {},
types.TeleportDropGroup: {},
"foo": {},
"bar": {},
apiconstants.TeleportDropGroup: {},
},
},
ui: &decisionpb.HostUsersInfo{
Mode: decisionpb.HostUserMode_HOST_USER_MODE_DROP,
Groups: []string{"foo", "bar", types.TeleportDropGroup},
Groups: []string{"foo", "bar", apiconstants.TeleportDropGroup},
},
expectGroups: nil,
@@ -1144,8 +1144,8 @@ func TestRegressionGroupErrorDoesNotPanic(t *testing.T) {
assert.NoError(t, err)
assert.Equal(t, nil, closer)
assert.Zero(t, backend.updateUserCalls)
assert.ElementsMatch(t, append(userinfo.Groups, types.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], types.TeleportDropGroup)
assert.ElementsMatch(t, append(userinfo.Groups, apiconstants.TeleportKeepGroup), backend.users["alice"])
assert.NotContains(t, backend.users["alice"], apiconstants.TeleportDropGroup)
backend.groupDatabaseErr = errors.New("could not find group")
_, err = users.UpsertUser("alice", &userinfo)
+4 -4
View File
@@ -24,7 +24,7 @@ import (
"log/slog"
"syscall"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
// errorWithExitStatus defines an interface that provides an ExitStatus
@@ -48,7 +48,7 @@ type execExitError interface {
func ExitCodeFromExecError(err error) int {
// If no error occurred, return 0 (success).
if err == nil {
return teleport.RemoteCommandSuccess
return reexecconstants.RemoteCommandSuccess
}
var execExitErr execExitError
@@ -57,7 +57,7 @@ func ExitCodeFromExecError(err error) int {
case errors.As(err, &execExitErr):
waitStatus, ok := execExitErr.Sys().(syscall.WaitStatus)
if !ok {
return teleport.RemoteCommandFailure
return reexecconstants.RemoteCommandFailure
}
return waitStatus.ExitStatus()
case errors.As(err, &exitErr):
@@ -65,6 +65,6 @@ func ExitCodeFromExecError(err error) int {
// An error occurred, but the type is unknown, return a generic 255 code.
default:
slog.DebugContext(context.Background(), "Unknown error returned when executing command", "error", err)
return teleport.RemoteCommandFailure
return reexecconstants.RemoteCommandFailure
}
}
+4 -4
View File
@@ -27,7 +27,7 @@ import (
"github.com/stretchr/testify/require"
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
type mockErrorWithExitStatus struct {
@@ -66,7 +66,7 @@ func TestExitCodeFromExecError(t *testing.T) {
{
name: "success",
input: nil,
want: teleport.RemoteCommandSuccess,
want: reexecconstants.RemoteCommandSuccess,
},
{
name: "exec exit error",
@@ -76,7 +76,7 @@ func TestExitCodeFromExecError(t *testing.T) {
{
name: "exec exit error with unknown sys",
input: mockExecExitError{sys: "unknown"},
want: teleport.RemoteCommandFailure,
want: reexecconstants.RemoteCommandFailure,
},
{
name: "ssh exit error",
@@ -86,7 +86,7 @@ func TestExitCodeFromExecError(t *testing.T) {
{
name: "unknown error",
input: errors.New("unknown error"),
want: teleport.RemoteCommandFailure,
want: reexecconstants.RemoteCommandFailure,
},
}
+3 -2
View File
@@ -159,6 +159,7 @@ import (
"github.com/gravitational/teleport/lib/web/terminal"
webui "github.com/gravitational/teleport/lib/web/ui"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec"
)
const hostID = "00000000-0000-0000-0000-000000000000"
@@ -189,8 +190,8 @@ func TestMain(m *testing.M) {
modules.SetInsecureTestMode(true)
// If the test is re-executing itself, execute the command that comes over
// the pipe.
if srv.IsReexec() {
srv.RunAndExit(os.Args[1])
if reexec.IsReexec() {
reexec.RunAndExit(os.Args[1])
return
}
@@ -14,7 +14,7 @@
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package logutils
package logconstants
// these constants are asserted to be the same as the ones in
// github.com/gravitational/teleport, they are inlined here to avoid circular
+6 -6
View File
@@ -59,7 +59,7 @@ import (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/session/logutils"
"github.com/gravitational/teleport/session/logconstants"
"github.com/gravitational/teleport/session/pam/pamcfg"
)
@@ -102,7 +102,7 @@ func init() {
}
// logComponent is the value for the component attribute in logs for this
// package; it should be paired with [logutils.ComponentKey]. Using a
// package; it should be paired with [logconstants.ComponentKey]. Using a
// package-level logger with this attribute created with [slog.With] would not
// synchronize correctly with [slog.SetDefault] and copying the packagelogger
// construction from lib/utils/log doesn't seem to be worth it when we have
@@ -154,7 +154,7 @@ func writeCallback(index C.int, stream C.int, s *C.char) {
if err != nil {
slog.ErrorContext(context.Background(),
"Unable to write to output stream",
logutils.ComponentKey, logComponent,
logconstants.ComponentKey, logComponent,
"error", err,
)
return
@@ -174,7 +174,7 @@ func readCallback(index C.int, e C.int) *C.char {
if err != nil {
slog.ErrorContext(context.Background(),
"Unable to read from input stream",
logutils.ComponentKey, logComponent,
logconstants.ComponentKey, logComponent,
"error", err,
)
return nil
@@ -190,7 +190,7 @@ func readCallback(index C.int, e C.int) *C.char {
if err != nil {
slog.ErrorContext(context.Background(),
"Unable to read from input stream",
logutils.ComponentKey, logComponent,
logconstants.ComponentKey, logComponent,
"error", err,
)
return nil
@@ -456,7 +456,7 @@ func (p *PAM) free() {
if retval != C.PAM_SUCCESS {
slog.WarnContext(context.Background(),
"Failed to end PAM transaction",
logutils.ComponentKey, logComponent,
logconstants.ComponentKey, logComponent,
"error", p.codeToError(retval),
)
}
+125
View File
@@ -0,0 +1,125 @@
// Teleport
// Copyright (C) 2026 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package reexec
import (
"bufio"
"context"
"errors"
"log/slog"
"os"
"os/exec"
"strings"
"syscall"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
const (
defaultPath = "/bin:/usr/bin:/usr/local/bin:/sbin"
defaultEnvPath = "PATH=" + defaultPath
defaultRootPath = "/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"
defaultEnvRootPath = "PATH=" + defaultRootPath
defaultLoginDefsPath = "/etc/login.defs"
)
// GetDefaultEnvPath returns the default value of PATH environment variable for
// new logins (prior to shell) based on login.defs. Returns a string which
// looks like "PATH=/usr/bin:/bin"
func GetDefaultEnvPath(uid string) string {
return getDefaultEnvPathWithLoginDefs(uid, defaultLoginDefsPath)
}
func getDefaultEnvPathWithLoginDefs(uid string, loginDefsPath string) string {
envPath := defaultEnvPath
envRootPath := defaultEnvRootPath
// open file, if it doesn't exist return a default path and move on
f, err := os.Open(loginDefsPath)
if err != nil {
if uid == "0" {
slog.DebugContext(context.Background(), "Unable to open login.defs, returning default su path", "login_defs_path", loginDefsPath, "error", err, "default_path", defaultEnvRootPath)
return defaultEnvRootPath
}
slog.DebugContext(context.Background(), "Unable to open login.defs, returning default path", "login_defs_path", loginDefsPath, "error", err, "default_path", defaultEnvPath)
return defaultEnvPath
}
defer f.Close()
// read path from login.defs file (/etc/login.defs) line by line:
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
// skip comments and empty lines:
if line == "" || line[0] == '#' {
continue
}
// look for a line that starts with ENV_PATH or ENV_SUPATH
fields := strings.Fields(line)
if len(fields) > 1 {
if fields[0] == "ENV_PATH" {
envPath = fields[1]
}
if fields[0] == "ENV_SUPATH" {
envRootPath = fields[1]
}
}
}
// if any error occurs while reading the file, return the default value
err = scanner.Err()
if err != nil {
if uid == "0" {
slog.WarnContext(context.Background(), "Unable to read login.defs, returning default su path", "login_defs_path", loginDefsPath, "error", err, "default_path", defaultEnvRootPath)
return defaultEnvRootPath
}
slog.WarnContext(context.Background(), "Unable to read login.defs, returning default path", "login_defs_path", loginDefsPath, "error", err, "default_path", defaultEnvPath)
return defaultEnvPath
}
// if requesting path for uid 0 and no ENV_SUPATH is given, fallback to
// ENV_PATH first, then the default path.
if uid == "0" {
return envRootPath
}
return envPath
}
// exitCode extracts and returns the exit code from the error.
func exitCode(err error) int {
// If no error occurred, return 0 (success).
if err == nil {
return reexecconstants.RemoteCommandSuccess
}
var execExitErr *exec.ExitError
switch {
// Local execution.
case errors.As(err, &execExitErr):
waitStatus, ok := execExitErr.Sys().(syscall.WaitStatus)
if !ok {
return reexecconstants.RemoteCommandFailure
}
return waitStatus.ExitStatus()
// An error occurred, but the type is unknown, return a generic 255 code.
default:
slog.DebugContext(context.Background(), "Unknown error returned when executing command", "error", err)
return reexecconstants.RemoteCommandFailure
}
}
+60
View File
@@ -0,0 +1,60 @@
// Teleport
// Copyright (C) 2026 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package reexec
import (
"flag"
"log/slog"
"os"
"testing"
"github.com/stretchr/testify/require"
)
func TestMain(m *testing.M) {
if IsReexec() {
RunAndExit(os.Args[1])
return
}
if !flag.Parsed() {
flag.Parse()
}
if testing.Verbose() {
slog.SetDefault(slog.New(slog.NewJSONHandler(
os.Stderr,
&slog.HandlerOptions{
Level: slog.LevelDebug,
},
)))
} else {
slog.SetDefault(slog.New(slog.DiscardHandler))
}
os.Exit(m.Run())
}
func TestLoginDefsParser(t *testing.T) {
t.Parallel()
expectedEnvSuPath := "PATH=/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin:/bar"
expectedSuPath := "PATH=/usr/local/bin:/usr/bin:/bin:/foo"
require.Equal(t, expectedEnvSuPath, getDefaultEnvPathWithLoginDefs("0", "../../fixtures/login.defs"))
require.Equal(t, expectedSuPath, getDefaultEnvPathWithLoginDefs("1000", "../../fixtures/login.defs"))
require.Equal(t, defaultEnvPath, getDefaultEnvPathWithLoginDefs("1000", "bad/file"))
}
@@ -0,0 +1,62 @@
// Copyright 2022 The Go Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
package logutils
import "sync"
// buffer adapted from go/src/fmt/print.go
type buffer []byte
// Having an initial size gives a dramatic speedup.
var bufPool = sync.Pool{
New: func() any {
b := make([]byte, 0, 1024)
return (*buffer)(&b)
},
}
func newBuffer() *buffer {
return bufPool.Get().(*buffer)
}
func (b *buffer) Len() int {
return len(*b)
}
func (b *buffer) SetLen(n int) {
*b = (*b)[:n]
}
func (b *buffer) Free() {
// To reduce peak allocation, return only smaller buffers to the pool.
const maxBufferSize = 16 << 10
if cap(*b) <= maxBufferSize {
*b = (*b)[:0]
bufPool.Put(b)
}
}
func (b *buffer) Reset() {
*b = (*b)[:0]
}
func (b *buffer) Write(p []byte) (int, error) {
*b = append(*b, p...)
return len(p), nil
}
func (b *buffer) WriteString(s string) (int, error) {
*b = append(*b, s...)
return len(s), nil
}
func (b *buffer) WriteByte(c byte) error {
*b = append(*b, c)
return nil
}
func (b *buffer) String() string {
return string(*b)
}
@@ -0,0 +1,339 @@
// Copyright 2022 The Go Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
package logutils
import (
"encoding"
"fmt"
"log/slog"
"reflect"
"strconv"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/session/logconstants"
)
// handleState adapted from go/src/log/slog/handler.go
type handleState struct {
h *SlogTextHandler
buf *buffer
freeBuf bool // should buf be freed?
prefix *buffer // for text: key prefix
groups *[]string // pool-allocated slice of active groups, for ReplaceAttr
}
var groupPool = sync.Pool{New: func() any {
s := make([]string, 0, 10)
return &s
}}
func (s *handleState) free() {
if s.freeBuf {
s.buf.Free()
}
if gs := s.groups; gs != nil {
*gs = (*gs)[:0]
groupPool.Put(gs)
}
s.prefix.Free()
}
func (s *handleState) openGroups() {
for _, n := range s.h.groups[s.h.nOpenGroups:] {
s.openGroup(n)
}
}
// openGroup starts a new group of attributes
// with the given name.
func (s *handleState) openGroup(name string) {
s.prefix.WriteString(name)
s.prefix.WriteByte('.')
// Collect group names for ReplaceAttr.
if s.groups != nil {
*s.groups = append(*s.groups, name)
}
}
// closeGroup ends the group with the given name.
func (s *handleState) closeGroup(name string) {
*s.prefix = (*s.prefix)[:len(*s.prefix)-len(name)-1 /* for keyComponentSep */]
if s.groups != nil {
*s.groups = (*s.groups)[:len(*s.groups)-1]
}
}
// appendAttrs appends the slice of Attrs.
// It reports whether something was appended.
func (s *handleState) appendAttrs(as []slog.Attr) bool {
nonEmpty := false
for _, a := range as {
if s.appendAttr(a) {
nonEmpty = true
}
}
return nonEmpty
}
// appendAttr appends the Attr's key and value.
// It handles replacement and checking for an empty key.
// It reports whether something was appended.
func (s *handleState) appendAttr(a slog.Attr) bool {
a.Value = a.Value.Resolve()
if rep := s.h.cfg.ReplaceAttr; rep != nil && a.Value.Kind() != slog.KindGroup {
var gs []string
if s.groups != nil {
gs = *s.groups
}
// a.Value is resolved before calling ReplaceAttr, so the user doesn't have to.
a = rep(gs, a)
// The ReplaceAttr function may return an unresolved Attr.
a.Value = a.Value.Resolve()
}
// Elide empty Attrs.
if a.Equal(slog.Attr{}) {
return false
}
// Handle nested attributes from within component fields.
if a.Key == logconstants.ComponentFields {
nonEmpty := false
switch fields := a.Value.Any().(type) {
case map[string]any:
for k, v := range fields {
if s.appendAttr(slog.Any(k, v)) {
nonEmpty = true
}
}
return nonEmpty
}
}
// Handle special cases before formatting.
var traceError bool
if a.Value.Kind() == slog.KindAny {
switch v := a.Value.Any().(type) {
case *slog.Source:
a.Value = slog.StringValue(fmt.Sprintf(" %s:%d", v.File, v.Line))
case trace.Error:
traceError = true
a.Value = slog.StringValue("[" + v.DebugReport() + "]")
case error:
a.Value = slog.StringValue(fmt.Sprintf("[%v]", v))
}
}
if a.Value.Kind() == slog.KindGroup {
attrs := a.Value.Group()
// Output only non-empty groups.
if len(attrs) > 0 {
// The group may turn out to be empty even though it has attrs (for
// example, ReplaceAttr may delete all the attrs).
// So remember where we are in the buffer, to restore the position
// later if necessary.
pos := s.buf.Len()
// Inline a group with an empty key.
if a.Key != "" {
s.openGroup(a.Key)
}
if !s.appendAttrs(attrs) {
s.buf.SetLen(pos)
return false
}
if a.Key != "" {
s.closeGroup(a.Key)
}
}
return true
}
s.appendKey(a.Key)
// Write the level key to avoid quoting color formatting that exists or
// [trace.Error]s so that the debug report is output in it's entirety.
if traceError || a.Key == slog.LevelKey {
s.buf.WriteString(a.Value.String())
} else {
s.appendValue(a.Value)
}
return true
}
func (s *handleState) appendError(err error) {
s.appendString(fmt.Sprintf("!ERROR:%v", err))
}
func (s *handleState) appendKey(key string) {
if s.buf.Len() > 0 {
s.buf.WriteString(" ")
}
// These keys should not be included in the output to match
// the behavior of the lorgus formatter.
if key == slog.TimeKey ||
key == logconstants.ComponentKey ||
key == slog.LevelKey ||
key == CallerField ||
key == slog.MessageKey ||
key == slog.SourceKey {
return
}
if s.prefix != nil && len(*s.prefix) > 0 {
// TODO: optimize by avoiding allocation.
s.appendString(string(*s.prefix) + key)
} else {
s.appendString(key)
}
s.buf.WriteByte(':')
}
func (s *handleState) appendString(str string) {
if str == "" {
return
}
if needsQuoting(str) {
*s.buf = strconv.AppendQuote(*s.buf, str)
} else {
s.buf.WriteString(str)
}
}
func (s *handleState) appendValue(v slog.Value) {
defer func() {
if r := recover(); r != nil {
// If it panics with a nil pointer, the most likely cases are
// an encoding.TextMarshaler or error fails to guard against nil,
// in which case "<nil>" seems to be the feasible choice.
//
// Adapted from the code in fmt/print.go.
if v := reflect.ValueOf(v.Any()); v.Kind() == reflect.Pointer && v.IsNil() {
s.appendString("<nil>")
return
}
// Otherwise just print the original panic message.
s.appendString(fmt.Sprintf("!PANIC: %v", r))
}
}()
if err := appendTextValue(s, v); err != nil {
s.appendError(err)
}
}
func (s *handleState) appendTime(t time.Time) {
*s.buf = appendRFC3339Millis(*s.buf, t)
}
func (s *handleState) appendNonBuiltIns(r slog.Record) {
// preformatted Attrs
if pfa := s.h.preformatted; len(pfa) > 0 {
s.buf.WriteString(" ")
s.buf.Write(pfa)
}
// Attrs in Record -- unlike the built-in ones, they are in groups started
// from WithGroup.
// If the record has no Attrs, don't output any groups.
if r.NumAttrs() > 0 {
s.prefix.WriteString(s.h.groupPrefix)
// The group may turn out to be empty even though it has attrs (for
// example, ReplaceAttr may delete all the attrs).
// So remember where we are in the buffer, to restore the position
// later if necessary.
pos := s.buf.Len()
s.openGroups()
empty := true
r.Attrs(func(a slog.Attr) bool {
// The component is handled by the top level handler.
if a.Key == logconstants.ComponentKey {
return true
}
if s.appendAttr(a) {
empty = false
}
return true
})
if empty {
s.buf.SetLen(pos)
}
}
}
func byteSlice(a any) ([]byte, bool) {
if bs, ok := a.([]byte); ok {
return bs, true
}
// Like Printf's %s, we allow both the slice type and the byte element type to be named.
t := reflect.TypeOf(a)
if t != nil && t.Kind() == reflect.Slice && t.Elem().Kind() == reflect.Uint8 {
return reflect.ValueOf(a).Bytes(), true
}
return nil, false
}
func appendTextValue(s *handleState, v slog.Value) error {
switch v.Kind() {
case slog.KindString:
s.appendString(v.String())
case slog.KindTime:
s.appendTime(v.Time())
case slog.KindAny:
if tm, ok := v.Any().(encoding.TextMarshaler); ok {
data, err := tm.MarshalText()
if err != nil {
return err
}
// TODO: avoid the conversion to string.
s.appendString(string(data))
return nil
}
if bs, ok := byteSlice(v.Any()); ok {
// As of Go 1.19, this only allocates for strings longer than 32 bytes.
s.buf.WriteString(strconv.Quote(string(bs)))
return nil
}
s.appendString(fmt.Sprintf("%+v", v.Any()))
case slog.KindInt64:
*s.buf = strconv.AppendInt(*s.buf, v.Int64(), 10)
case slog.KindUint64:
*s.buf = strconv.AppendUint(*s.buf, v.Uint64(), 10)
case slog.KindFloat64:
*s.buf = strconv.AppendFloat(*s.buf, v.Float64(), 'g', -1, 64)
case slog.KindBool:
*s.buf = strconv.AppendBool(*s.buf, v.Bool())
case slog.KindDuration:
*s.buf = append(*s.buf, v.Duration().String()...)
case slog.KindGroup:
*s.buf = fmt.Append(*s.buf, v.Group())
case slog.KindLogValuer:
*s.buf = fmt.Append(*s.buf, v.Any())
default:
panic(fmt.Sprintf("bad kind: %s", v.Kind()))
}
return nil
}
func appendRFC3339Millis(b []byte, t time.Time) []byte {
// Format according to time.RFC3339Nano since it is highly optimized,
// but truncate it to use millisecond resolution.
// Unfortunately, that format trims trailing 0s, so add 1/10 millisecond
// to guarantee that there are exactly 4 digits after the period.
const prefixLen = len("2006-01-02T15:04:05.000")
n := len(b)
t = t.Truncate(time.Millisecond).Add(time.Millisecond / 10)
b = t.AppendFormat(b, time.RFC3339Nano)
b = append(b[:n+prefixLen], b[n+prefixLen+1:]...) // drop the 4th digit
return b
}
+129
View File
@@ -0,0 +1,129 @@
/*
* Teleport
* Copyright (C) 2023 Gravitational, Inc.
*
* This program is free software: you can redistribute it and/or modify
* it under the terms of the GNU Affero General Public License as published by
* the Free Software Foundation, either version 3 of the License, or
* (at your option) any later version.
*
* This program is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU Affero General Public License for more details.
*
* You should have received a copy of the GNU Affero General Public License
* along with this program. If not, see <http://www.gnu.org/licenses/>.
*/
package logutils
import (
"context"
"log/slog"
"strings"
"unicode"
"github.com/gravitational/trace"
)
const (
// TraceLevel is the logging level when set to Trace verbosity.
TraceLevel = slog.LevelDebug - 1
// TraceLevelText is the text representation of Trace verbosity.
TraceLevelText = "TRACE"
noColor = -1
red = 31
yellow = 33
blue = 36
gray = 37
// LevelField is the log field that stores the verbosity.
LevelField = "level"
// ComponentField is the log field that stores the calling component.
ComponentField = "component"
// CallerField is the log field that stores the calling file and line number.
CallerField = "caller"
// TimestampField is the field that stores the timestamp the log was emitted.
TimestampField = "timestamp"
messageField = "message"
// defaultComponentPadding is a default padding for component field
defaultComponentPadding = 11
// defaultLevelPadding is a default padding for level field
defaultLevelPadding = 4
)
// SupportedLevelsText lists the supported log levels in their text
// representation. All strings are in uppercase.
var SupportedLevelsText = []string{
TraceLevelText,
slog.LevelDebug.String(),
slog.LevelInfo.String(),
slog.LevelWarn.String(),
slog.LevelError.String(),
}
func addTracingContextToRecord(ctx context.Context, r *slog.Record) {
// there can't be a span from context from OTEL because we don't include
// OTEL in this module
}
var defaultFormatFields = []string{LevelField, ComponentField, CallerField, TimestampField}
var knownFormatFields = map[string]struct{}{
LevelField: {},
ComponentField: {},
CallerField: {},
TimestampField: {},
}
// ValidateFields ensures the provided fields map to the allowed fields. An error
// is returned if any of the fields are invalid.
func ValidateFields(formatInput []string) (result []string, err error) {
for _, component := range formatInput {
component = strings.TrimSpace(component)
if _, ok := knownFormatFields[component]; !ok {
return nil, trace.BadParameter("invalid log format key: %q", component)
}
result = append(result, component)
}
return result, nil
}
// needsQuoting returns true if any non-printable characters are found.
func needsQuoting(text string) bool {
for _, r := range text {
if !unicode.IsPrint(r) {
return true
}
}
return false
}
func padMax(in string, chars int) string {
switch {
case len(in) < chars:
return in + strings.Repeat(" ", chars-len(in))
default:
return in[:chars]
}
}
// getCaller retrieves source information from the attribute
// and returns the file and line of the caller. The file is
// truncated from the absolute path to package/filename.
func getCaller(s *slog.Source) (file string, line int) {
count := 0
idx := strings.LastIndexFunc(s.File, func(r rune) bool {
if r == '/' {
count++
}
return count == 2
})
file = s.File[idx+1:]
line = s.Line
return file, line
}
@@ -0,0 +1,165 @@
// Teleport
// Copyright (C) 2024 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package logutils
import (
"context"
"fmt"
"io"
"log/slog"
"slices"
"strings"
"time"
"github.com/gravitational/teleport/session/logconstants"
)
// SlogJSONHandlerConfig allows the SlogJSONHandler functionality
// to be tweaked.
type SlogJSONHandlerConfig struct {
// Level is the minimum record level that will be logged.
Level slog.Leveler
// ConfiguredFields are fields explicitly set by users to be included in
// the output message. If there are any entries configured, they will be honored.
// If empty, the default fields will be populated and included in the output.
ConfiguredFields []string
// ReplaceAttr is called to rewrite each non-group attribute before
// it is logged.
ReplaceAttr func(groups []string, a slog.Attr) slog.Attr
}
// SlogJSONHandler is a [slog.Handler] that outputs messages in a json
// format per the config file.
type SlogJSONHandler struct {
*slog.JSONHandler
}
// NewSlogJSONHandler creates a SlogJSONHandler that outputs to w.
func NewSlogJSONHandler(w io.Writer, cfg SlogJSONHandlerConfig) *SlogJSONHandler {
withCaller := len(cfg.ConfiguredFields) == 0 || slices.Contains(cfg.ConfiguredFields, CallerField)
withComponent := len(cfg.ConfiguredFields) == 0 || slices.Contains(cfg.ConfiguredFields, ComponentField)
withTimestamp := len(cfg.ConfiguredFields) == 0 || slices.Contains(cfg.ConfiguredFields, TimestampField)
return &SlogJSONHandler{
JSONHandler: slog.NewJSONHandler(w, &slog.HandlerOptions{
AddSource: true,
Level: cfg.Level,
ReplaceAttr: func(groups []string, a slog.Attr) slog.Attr {
switch a.Key {
case logconstants.ComponentKey:
if !withComponent {
return slog.Attr{}
}
if a.Value.Kind() != slog.KindString {
return a
}
a.Key = ComponentField
case slog.LevelKey:
// The slog.JSONHandler will inject "level" Attr.
// However, this lib's consumer might add an Attr using the same key ("level") and we end up with two records named "level".
// We must check its type before assuming this was injected by the slog.JSONHandler.
lvl, ok := a.Value.Any().(slog.Level)
if !ok {
return a
}
var level string
switch lvl {
case TraceLevel:
level = "trace"
case slog.LevelDebug:
level = "debug"
case slog.LevelInfo:
level = "info"
case slog.LevelWarn:
level = "warning"
case slog.LevelError:
level = "error"
default:
level = strings.ToLower(lvl.String())
}
a.Value = slog.StringValue(level)
case slog.TimeKey:
if !withTimestamp {
return slog.Attr{}
}
// The slog.JSONHandler will inject "time" Attr.
// However, this lib's consumer might add an Attr using the same key ("time") and we end up with two records named "time".
// We must check its type before assuming this was injected by the slog.JSONHandler.
if a.Value.Kind() != slog.KindTime {
return a
}
t := a.Value.Time()
if t.IsZero() {
return a
}
a.Key = TimestampField
a.Value = slog.StringValue(t.Format(time.RFC3339))
case slog.MessageKey:
// The slog.JSONHandler will inject "msg" Attr.
// However, this lib's consumer might add an Attr using the same key ("msg") and we end up with two records named "msg".
// We must check its type before assuming this was injected by the slog.JSONHandler.
if a.Value.Kind() != slog.KindString {
return a
}
a.Key = messageField
case slog.SourceKey:
if !withCaller {
return slog.Attr{}
}
// The slog.JSONHandler will inject "source" Attr when AddSource is true.
// However, this lib's consumer might add an Attr using the same key ("source") and we end up with two records named "source".
// We must check its type before assuming this was injected by the slog.JSONHandler.
s, ok := a.Value.Any().(*slog.Source)
if !ok {
return a
}
file, line := getCaller(s)
a = slog.String(CallerField, fmt.Sprintf("%s:%d", file, line))
}
// Convert [slog.KindAny] values that are backed by an [error] or [fmt.Stringer]
// to strings so that only the message is output instead of a json object. The kind is
// first checked to avoid allocating an interface for the values stored inline
// in [slog.Attr].
if a.Value.Kind() == slog.KindAny {
if err, ok := a.Value.Any().(error); ok {
a.Value = slog.StringValue(err.Error())
}
if stringer, ok := a.Value.Any().(fmt.Stringer); ok {
a.Value = slog.StringValue(stringer.String())
}
}
return a
},
}),
}
}
func (j *SlogJSONHandler) Handle(ctx context.Context, r slog.Record) error {
addTracingContextToRecord(ctx, &r)
return j.JSONHandler.Handle(ctx, r)
}
@@ -0,0 +1,387 @@
// Teleport
// Copyright (C) 2024 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package logutils
import (
"context"
"fmt"
"io"
"log/slog"
"runtime"
"slices"
"strings"
"sync"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/session/logconstants"
)
// SlogTextHandler is a [slog.Handler] that outputs messages in a textual
// manner as configured by the Teleport configuration.
type SlogTextHandler struct {
cfg SlogTextHandlerConfig
out slogTextHandlerWriter
// withCaller indicates whether the location the log was emitted from
// should be included in the output message.
withCaller bool
// withTimestamp indicates whether the times that the log was emitted at
// should be included in the output message.
withTimestamp bool
// rawComponent is the Teleport subcomponent that emitted the log, e.g., "tsh".
rawComponent string
// component is rawComponent wrapped in square brackets and truncated if necessary to not exceed
// cfg.Padding, e.g., "[TSH]".
component string
// preformatted data from previous calls to WithGroup and WithAttrs.
preformatted []byte
// groupPrefix is for the text handler only.
// It holds the prefix for groups that were already pre-formatted.
// A group will appear here when a call to WithGroup is followed by
// a call to WithAttrs.
groupPrefix string
// groups passed in via WithGroup and WithAttrs.
groups []string
// nOpenGroups the number of groups opened in preformatted.
nOpenGroups int
}
type slogTextHandlerWriter interface {
Write(bytes []byte, component string, level slog.Level) error
}
// SlogTextHandlerConfig allow the SlogTextHandler functionality
// to be tweaked.
type SlogTextHandlerConfig struct {
// Level is the minimum record level that will be logged.
Level slog.Leveler
// EnableColors allows the level to be printed in color.
EnableColors bool
// Padding to use for [ComponentField] to ensure that the initial columns in the output line up.
// The component is wrapped in square brackets. If the length of the component exceeds Padding+2,
// the component is truncated. If the length is less than Padding+2, the component in square
// brackets is followed by spaces to pad it to the given Padding.
//
// If set to zero, no padding is done and components are not truncated.
Padding int
// ConfiguredFields are fields explicitly set by users to be included in
// the output message. If there are any entries configured, they will be honored.
// If empty, the default fields will be populated and included in the output.
ConfiguredFields []string
// ReplaceAttr is called to rewrite each non-group attribute before
// it is logged.
ReplaceAttr func(groups []string, a slog.Attr) slog.Attr
}
// NewSlogTextHandler creates a SlogTextHandler that writes messages to w.
func NewSlogTextHandler(w io.Writer, cfg SlogTextHandlerConfig) *SlogTextHandler {
if cfg.Padding == 0 {
cfg.Padding = defaultComponentPadding
}
handler := SlogTextHandler{
cfg: cfg,
out: newIOWriter(w),
withCaller: len(cfg.ConfiguredFields) == 0 || slices.Contains(cfg.ConfiguredFields, CallerField),
withTimestamp: len(cfg.ConfiguredFields) == 0 || slices.Contains(cfg.ConfiguredFields, TimestampField),
}
if handler.cfg.ConfiguredFields == nil {
handler.cfg.ConfiguredFields = defaultFormatFields
}
return &handler
}
// Enabled returns whether the provided level will be included in output.
func (s *SlogTextHandler) Enabled(ctx context.Context, level slog.Level) bool {
minLevel := slog.LevelInfo
if s.cfg.Level != nil {
minLevel = s.cfg.Level.Level()
}
return level >= minLevel
}
func (s *SlogTextHandler) newHandleState(buf *buffer, freeBuf bool) handleState {
state := handleState{
h: s,
buf: buf,
freeBuf: freeBuf,
prefix: newBuffer(),
}
if s.cfg.ReplaceAttr != nil {
state.groups = groupPool.Get().(*[]string)
*state.groups = append(*state.groups, s.groups[:s.nOpenGroups]...)
}
return state
}
// Handle formats the provided record and writes the output to the
// destination.
func (s *SlogTextHandler) Handle(ctx context.Context, r slog.Record) error {
state := s.newHandleState(newBuffer(), true)
defer state.free()
addTracingContextToRecord(ctx, &r)
// Built-in attributes. They are not in a group.
stateGroups := state.groups
state.groups = nil // So ReplaceAttrs sees no groups instead of the pre groups.
rep := s.cfg.ReplaceAttr
if s.withTimestamp && !r.Time.IsZero() {
if rep == nil {
state.appendKey(slog.TimeKey)
state.appendTime(r.Time.Round(0))
} else {
state.appendAttr(slog.Time(slog.TimeKey, r.Time.Round(0)))
}
}
rawComponent := s.rawComponent
// Processing fields in this manner allows users to
// configure the level and component position in the output.
// This matches the behavior of the original logrus formatter. All other
// fields location in the output message are static.
for _, field := range s.cfg.ConfiguredFields {
switch field {
case LevelField:
level := formatLevel(r.Level, s.cfg.EnableColors)
if rep == nil {
state.appendKey(slog.LevelKey)
// Write the level directly to stat to avoid quoting
// color formatting that exists.
state.buf.WriteString(level)
} else {
state.appendAttr(slog.String(slog.LevelKey, level))
}
case ComponentField:
// If a component is provided with the attributes, it should be used instead of
// the component set on the handler. Note that if there are multiple components
// specified in the arguments, the one with the lowest index is used and the others are ignored.
// In the example below, the resulting component in the message output would be "alpaca".
//
// logger := logger.With(teleport.ComponentKey, "fish")
// logger.InfoContext(ctx, "llama llama llama", teleport.ComponentKey, "alpaca", "foo", 123, teleport.ComponentKey, "shark")
component := s.component
r.Attrs(func(attr slog.Attr) bool {
if attr.Key != logconstants.ComponentKey {
return true
}
rawComponent = attr.Value.String()
component = formatComponent(attr.Value, s.cfg.Padding)
return false
})
if rep == nil {
state.appendKey(logconstants.ComponentKey)
state.appendString(component)
} else {
state.appendAttr(slog.String(logconstants.ComponentKey, component))
}
default:
if _, ok := knownFormatFields[field]; !ok {
return trace.BadParameter("invalid log format key: %v", field)
}
}
}
if rep == nil {
state.appendKey(slog.MessageKey)
state.appendString(r.Message)
} else {
state.appendAttr(slog.String(slog.MessageKey, r.Message))
}
state.groups = stateGroups // Restore groups passed to ReplaceAttrs.
state.appendNonBuiltIns(r)
if r.PC != 0 && s.withCaller {
fs := runtime.CallersFrames([]uintptr{r.PC})
f, _ := fs.Next()
src := slog.Source{
Function: f.Function,
File: f.File,
Line: f.Line,
}
src.File, src.Line = getCaller(&src)
if rep == nil {
state.appendKey(slog.SourceKey)
state.appendString(fmt.Sprintf("%s:%d", src.File, src.Line))
} else {
state.appendAttr(slog.Any(slog.SourceKey, &src))
}
}
state.buf.WriteByte('\n')
return s.out.Write(*state.buf, rawComponent, r.Level)
}
func formatLevel(value slog.Level, enableColors bool) string {
var color int
var level string
switch value {
case TraceLevel:
level = "TRACE"
color = gray
case slog.LevelDebug:
level = "DEBUG"
color = gray
case slog.LevelInfo:
level = "INFO"
color = blue
case slog.LevelWarn:
level = "WARN"
color = yellow
case slog.LevelError:
level = "ERROR"
color = red
default:
color = blue
level = value.String()
}
if !enableColors {
color = noColor
}
level = padMax(level, defaultLevelPadding)
if color != noColor {
level = fmt.Sprintf("\u001B[%dm%s\u001B[0m", color, level)
}
return level
}
func formatComponent(value slog.Value, padding int) string {
component := strings.ToUpper(fmt.Sprintf("[%v]", value))
if padding <= 0 {
return component
}
component = padMax(component, padding)
if component[len(component)-1] != ' ' {
component = component[:len(component)-1] + "]"
}
return component
}
func (s *SlogTextHandler) clone() *SlogTextHandler {
return &SlogTextHandler{
cfg: s.cfg,
withCaller: s.withCaller,
withTimestamp: s.withTimestamp,
component: s.component,
rawComponent: s.rawComponent,
preformatted: slices.Clip(s.preformatted),
groupPrefix: s.groupPrefix,
groups: slices.Clip(s.groups),
nOpenGroups: s.nOpenGroups,
out: s.out,
}
}
// WithAttrs clones the current handler with the provided attributes
// added to any existing attributes. The values are preformatted here
// so that they do not need to be formatted per call to Handle.
func (s *SlogTextHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
if len(attrs) == 0 {
return s
}
s2 := s.clone()
// Pre-format the attributes as an optimization.
state := s2.newHandleState((*buffer)(&s2.preformatted), false)
defer state.free()
state.prefix.WriteString(s.groupPrefix)
// Remember the position in the buffer, in case all attrs are empty.
pos := state.buf.Len()
state.openGroups()
nonEmpty := false
for _, a := range attrs {
switch a.Key {
case logconstants.ComponentKey:
component := strings.ToUpper(fmt.Sprintf("[%v]", a.Value.String()))
if s.cfg.Padding > 0 {
component = padMax(component, s.cfg.Padding)
if component[len(component)-1] != ' ' {
component = component[:len(component)-1] + "]"
}
}
s2.component = component
s2.rawComponent = a.Value.String()
case logconstants.ComponentFields:
switch fields := a.Value.Any().(type) {
case map[string]any:
for k, v := range fields {
if state.appendAttr(slog.Any(k, v)) {
nonEmpty = true
}
}
}
default:
if state.appendAttr(a) {
nonEmpty = true
}
}
}
if !nonEmpty {
state.buf.SetLen(pos)
} else {
// Remember the new prefix for later keys.
s2.groupPrefix = state.prefix.String()
// Remember how many opened groups are in preformattedAttrs,
// so we don't open them again when we handle a Record.
s2.nOpenGroups = len(s2.groups)
}
return s2
}
// WithGroup opens a new group.
func (s *SlogTextHandler) WithGroup(name string) slog.Handler {
s2 := s.clone()
s2.groups = append(s2.groups, name)
return s2
}
type ioWriter struct {
mu sync.Mutex
out io.Writer
}
func newIOWriter(w io.Writer) *ioWriter {
return &ioWriter{out: w}
}
func (o *ioWriter) Write(bytes []byte, rawComponent string, level slog.Level) error {
o.mu.Lock()
defer o.mu.Unlock()
_, err := o.out.Write(bytes)
return err
}
+141 -176
View File
@@ -16,7 +16,7 @@
* along with this program. If not, see <http://www.gnu.org/licenses/>.
*/
package srv
package reexec
import (
"bytes"
@@ -44,19 +44,18 @@ import (
ocselinux "github.com/opencontainers/selinux/go-selinux"
"golang.org/x/sys/unix"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/sshutils"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
apiconstants "github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/session/auditd"
"github.com/gravitational/teleport/session/envutils"
"github.com/gravitational/teleport/session/host"
"github.com/gravitational/teleport/session/logconstants"
"github.com/gravitational/teleport/session/loginuid"
"github.com/gravitational/teleport/session/networking"
"github.com/gravitational/teleport/session/networking/x11"
"github.com/gravitational/teleport/session/pam"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec/internal/logutils"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
"github.com/gravitational/teleport/session/selinux"
"github.com/gravitational/teleport/session/shell"
"github.com/gravitational/teleport/session/uacc"
@@ -141,9 +140,9 @@ func fdName(f FileFD) string {
// ExecCommand contains the payload to "teleport exec" which will be used to
// construct and execute a shell.
type ExecCommand struct {
stdin io.Reader
stdout io.Writer
stderr io.Writer
Stdin io.Reader `json:"-"`
Stdout io.Writer `json:"-"`
Stderr io.Writer `json:"-"`
// LogConfig is the log configuration for the child process.
LogConfig ExecLogConfig `json:"log_config"`
@@ -245,7 +244,7 @@ type PAMConfig struct {
// UaccMetadata contains information the child needs from the parent for user accounting.
type UaccMetadata struct {
// RemoteAddr is the address of the remote host.
RemoteAddr utils.NetAddr `json:"remote_addr"`
RemoteAddr NetAddr `json:"remote_addr"`
// UtmpPath is the path of the system utmp database.
UtmpPath string `json:"utmp_path,omitempty"`
@@ -260,6 +259,28 @@ type UaccMetadata struct {
WtmpdbPath string `json:"wtmpdb_path,omitempty"`
}
// NetAddrFromAddr returns NetAddr from golang standard net.Addr
func NetAddrFromAddr(a net.Addr) NetAddr {
return NetAddr{AddrNetwork: a.Network(), Addr: a.String()}
}
type NetAddr struct {
AddrNetwork string `json:"network"`
Addr string `json:"addr"`
}
var _ net.Addr = (*NetAddr)(nil)
// Network implements [net.Addr].
func (n *NetAddr) Network() string {
return n.AddrNetwork
}
// String implements [net.Addr].
func (n *NetAddr) String() string {
return n.Addr
}
// RunCommand reads in the command to run from the parent process (over a
// pipe) then constructs and runs the command. This function may change
// system state related to the process and/or thread for PAM and SELinux.
@@ -361,28 +382,28 @@ func RunCommand() (exitErr error, err error) {
if tty == nil {
return nil, trace.BadParameter("tty not found")
}
c.stdin = tty
c.stdout = tty
c.stderr = tty
} else if c.RequestType == sshutils.SubsystemRequest && c.Command == teleport.SFTPSubsystem {
c.Stdin = tty
c.Stdout = tty
c.Stderr = tty
} else if c.RequestType == "subsystem" && c.Command == "sftp" {
// std{in/out} is not used by the SFTP sub process, just collect stderr.
c.stdin = bytes.NewReader(nil)
c.stdout = io.Discard
c.Stdin = bytes.NewReader(nil)
c.Stdout = io.Discard
// Propagate sftp subprocess errors to the parent process.
c.stderr = os.Stderr
c.Stderr = os.Stderr
} else {
// If this is a normal, non-interactive exec session, use the stdio pipes provided as extra files.
c.stdin = os.NewFile(StdinFile, fdName(StdinFile))
if c.stdin == nil {
c.Stdin = os.NewFile(StdinFile, fdName(StdinFile))
if c.Stdin == nil {
return nil, trace.BadParameter("stdin not found")
}
c.stdout = os.NewFile(StdoutFile, fdName(StdoutFile))
if c.stdout == nil {
c.Stdout = os.NewFile(StdoutFile, fdName(StdoutFile))
if c.Stdout == nil {
return nil, trace.BadParameter("stdout not found")
}
c.stderr = os.NewFile(StderrFile, fdName(StderrFile))
if c.stderr == nil {
c.Stderr = os.NewFile(StderrFile, fdName(StderrFile))
if c.Stderr == nil {
return nil, trace.BadParameter("stderr not found")
}
}
@@ -403,9 +424,9 @@ func RunCommand() (exitErr error, err error) {
// account/session.
Env: c.PAMConfig.Environment,
// Connect std{in,out,err} to the TTY if a terminal has been allocated.
Stdin: c.stdin,
Stdout: c.stdout,
Stderr: c.stderr,
Stdin: c.Stdin,
Stdout: c.Stdout,
Stderr: c.Stderr,
}
// Discard std{out,err} for non-interactive requests. Otherwise, things like
@@ -467,7 +488,7 @@ func RunCommand() (exitErr error, err error) {
}
// Build the actual command that will launch the shell.
cmd, err := buildCommand(&c, localUser, pamEnvironment)
cmd, err := BuildCommand(&c, localUser, pamEnvironment)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -475,7 +496,7 @@ func RunCommand() (exitErr error, err error) {
// Wait until the continue signal is received from Teleport signaling that
// Teleport is monitoring this session if Enhanced Session Recording is enabled.
if c.RecordWithBPF {
err = waitForSignal(ctx, contfd, 10*time.Second)
err = WaitForSignal(ctx, contfd, 10*time.Second)
if err != nil {
return nil, trace.Wrap(err)
}
@@ -677,9 +698,9 @@ func (o *osWrapper) startNewParker(ctx context.Context, credential *syscall.Cred
return nil
}
group, err := o.LookupGroup(types.TeleportDropGroup)
group, err := o.LookupGroup(apiconstants.TeleportDropGroup)
if err != nil {
if isUnknownGroupError(err, types.TeleportDropGroup) {
if isUnknownGroupError(err, apiconstants.TeleportDropGroup) {
// The service group doesn't exist. Auto-provision is disabled, do nothing.
return nil
}
@@ -726,21 +747,21 @@ func RunNetworking() (code int, err error) {
// Parent sends the command payload in the third file descriptor.
cmdfd := os.NewFile(CommandFile, fdName(CommandFile))
if cmdfd == nil {
return teleport.RemoteCommandFailure, trace.BadParameter("command pipe not found")
return reexecconstants.RemoteCommandFailure, trace.BadParameter("command pipe not found")
}
logfd := os.NewFile(LogFile, fdName(LogFile))
if logfd == nil {
return teleport.RemoteCommandFailure, trace.BadParameter("log pipe not found")
return reexecconstants.RemoteCommandFailure, trace.BadParameter("log pipe not found")
}
terminatefd := os.NewFile(TerminateFile, fdName(TerminateFile))
if terminatefd == nil {
return teleport.RemoteCommandFailure, trace.BadParameter("terminate pipe not found")
return reexecconstants.RemoteCommandFailure, trace.BadParameter("terminate pipe not found")
}
// Read in the command payload.
var c ExecCommand
if err := json.NewDecoder(cmdfd).Decode(&c); err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err)
return reexecconstants.RemoteCommandFailure, trace.Wrap(err)
}
initLogger("networking", logfd, c.LogConfig)
@@ -763,7 +784,7 @@ func RunNetworking() (code int, err error) {
Env: c.PAMConfig.Environment,
})
if err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err)
return reexecconstants.RemoteCommandFailure, trace.Wrap(err)
}
defer pamContext.Close()
@@ -775,12 +796,12 @@ func RunNetworking() (code int, err error) {
// done with the user's permissions.
localUser, err := user.Lookup(c.Login)
if err != nil {
return teleport.RemoteCommandFailure, trace.NotFound("%s", err)
return reexecconstants.RemoteCommandFailure, trace.NotFound("%s", err)
}
cred, err := host.GetHostUserCredential(localUser)
if err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err)
return reexecconstants.RemoteCommandFailure, trace.Wrap(err)
}
if os.Getuid() != int(cred.Uid) || os.Getgid() != int(cred.Gid) {
@@ -790,14 +811,14 @@ func RunNetworking() (code int, err error) {
groups[i] = int(g)
}
if err := unix.Setgroups(groups); err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err, "failed to set groups for networking process")
return reexecconstants.RemoteCommandFailure, trace.Wrap(err, "failed to set groups for networking process")
}
}
if err := unix.Setgid(int(cred.Gid)); err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err, "failed to set gid for networking process")
return reexecconstants.RemoteCommandFailure, trace.Wrap(err, "failed to set gid for networking process")
}
if err := unix.Setuid(int(cred.Uid)); err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err, "failed to set uid for networking process")
return reexecconstants.RemoteCommandFailure, trace.Wrap(err, "failed to set uid for networking process")
}
}
@@ -816,27 +837,27 @@ func RunNetworking() (code int, err error) {
for _, kv := range pamEnvironment {
key, value, ok := strings.Cut(strings.TrimSpace(kv), "=")
if !ok {
return teleport.RemoteCommandFailure, trace.BadParameter("bad environment variable from PAM, expected format \"key=value\" but got %q", kv)
return reexecconstants.RemoteCommandFailure, trace.BadParameter("bad environment variable from PAM, expected format \"key=value\" but got %q", kv)
}
if err := os.Setenv(key, value); err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err)
return reexecconstants.RemoteCommandFailure, trace.Wrap(err)
}
}
// Ensure that the working directory is one that the local user has access to.
if err := os.Chdir(workingDir); err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err, "failed to set working directory for networking process: %s", workingDir)
return reexecconstants.RemoteCommandFailure, trace.Wrap(err, "failed to set working directory for networking process: %s", workingDir)
}
ffd := os.NewFile(ListenerFile, "listener")
if ffd == nil {
return teleport.RemoteCommandFailure, trace.BadParameter("missing socket fd")
return reexecconstants.RemoteCommandFailure, trace.BadParameter("missing socket fd")
}
parentConn, err := uds.FromFile(ffd)
_ = ffd.Close()
if err != nil {
return teleport.RemoteCommandFailure, trace.Wrap(err)
return reexecconstants.RemoteCommandFailure, trace.Wrap(err)
}
ctx, cancel := context.WithCancel(context.Background())
@@ -865,9 +886,9 @@ func RunNetworking() (code int, err error) {
fbuf := make([]*os.File, 1)
n, fn, err := uds.ReadWithFDs(parentConn, buf, fbuf)
if err != nil {
if utils.IsOKNetworkError(err) {
if isOKNetworkError(err) {
// parent connection closed, process should exit.
return teleport.RemoteCommandSuccess, nil
return reexecconstants.RemoteCommandSuccess, nil
}
slog.ErrorContext(ctx, "error reading networking request from parent", "err", err)
continue
@@ -1055,15 +1076,15 @@ func getConnFile(conn net.Conn) (*os.File, error) {
// runCheckHomeDir checks if the active user's $HOME dir exists and is accessible.
func runCheckHomeDir() (code int) {
code = teleport.RemoteCommandSuccess
code = reexecconstants.RemoteCommandSuccess
if err := hasAccessibleHomeDir(); err != nil {
switch {
case trace.IsNotFound(err), trace.IsBadParameter(err):
code = teleport.HomeDirNotFound
code = reexecconstants.HomeDirNotFound
case trace.IsAccessDenied(err):
code = teleport.HomeDirNotAccessible
code = reexecconstants.HomeDirNotAccessible
default:
code = teleport.RemoteCommandFailure
code = reexecconstants.RemoteCommandFailure
}
}
@@ -1087,27 +1108,27 @@ func RunAndExit(commandType string) {
var err error
switch commandType {
case teleport.ExecSubCommand:
case reexecconstants.ExecSubCommand:
var execErr error
execErr, err = RunCommand()
if err != nil {
code = teleport.RemoteCommandFailure
code = reexecconstants.RemoteCommandFailure
} else {
code = exitCode(execErr)
}
case teleport.NetworkingSubCommand:
case reexecconstants.NetworkingSubCommand:
code, err = RunNetworking()
case teleport.CheckHomeDirSubCommand:
case reexecconstants.CheckHomeDirSubCommand:
code = runCheckHomeDir()
case teleport.ParkSubCommand:
case reexecconstants.ParkSubCommand:
code = runPark()
default:
code, err = teleport.RemoteCommandFailure, fmt.Errorf("unknown command type: %v", commandType)
code, err = reexecconstants.RemoteCommandFailure, fmt.Errorf("unknown command type: %v", commandType)
}
if err != nil {
// Write the error to stderr, where it can be seen by the parent teleport process and
// propagated to the client.
if code == teleport.RemoteCommandFailure {
if code == reexecconstants.RemoteCommandFailure {
fmt.Fprintf(os.Stderr, "Failed to launch: %v.\n", err)
}
@@ -1128,9 +1149,9 @@ func RunAndExit(commandType string) {
func IsReexec() bool {
if len(os.Args) >= 2 {
switch os.Args[1] {
case teleport.ExecSubCommand, teleport.NetworkingSubCommand,
teleport.CheckHomeDirSubCommand,
teleport.ParkSubCommand, teleport.SFTPSubCommand:
case reexecconstants.ExecSubCommand, reexecconstants.NetworkingSubCommand,
reexecconstants.CheckHomeDirSubCommand,
reexecconstants.ParkSubCommand, reexecconstants.SFTPSubCommand:
return true
}
}
@@ -1141,7 +1162,7 @@ func IsReexec() bool {
// openFileAsUser opens a file as the given user to ensure proper access checks. This is unsafe and should not be used outside of
// bootstrapping reexec commands.
func openFileAsUser(localUser *user.User, path string) (file *os.File, err error) {
if os.Args[1] != teleport.ExecSubCommand {
if os.Args[1] != reexecconstants.ExecSubCommand {
return nil, trace.Errorf("opening files as a user is only possible in a reexec context")
}
@@ -1164,7 +1185,7 @@ func openFileAsUser(localUser *user.User, path string) (file *os.File, err error
if uidErr != nil || gidErr != nil {
file.Close()
slog.ErrorContext(context.Background(), "cannot proceed with invalid effective credentials", "uid_err", uidErr, "gid_err", gidErr, "error", err)
os.Exit(teleport.UnexpectedCredentials)
os.Exit(reexecconstants.UnexpectedCredentials)
}
}()
@@ -1176,8 +1197,8 @@ func openFileAsUser(localUser *user.User, path string) (file *os.File, err error
return nil, trace.Wrap(err)
}
file, err = utils.OpenFileNoUnsafeLinks(path)
return file, trace.Wrap(err)
file, err = os.Open(path)
return file, trace.ConvertSystemError(err)
}
func readUserEnv(localUser *user.User, path string) ([]string, error) {
@@ -1191,9 +1212,9 @@ func readUserEnv(localUser *user.User, path string) ([]string, error) {
return envs, trace.Wrap(err)
}
// buildCommand constructs a command that will execute the user's shell. This
// BuildCommand constructs a command that will execute the user's shell. This
// function is run by Teleport while it's re-executing.
func buildCommand(c *ExecCommand, localUser *user.User, pamEnvironment []string) (*exec.Cmd, error) {
func BuildCommand(c *ExecCommand, localUser *user.User, pamEnvironment []string) (*exec.Cmd, error) {
var cmd exec.Cmd
isReexec := false
@@ -1210,15 +1231,15 @@ func buildCommand(c *ExecCommand, localUser *user.User, pamEnvironment []string)
// if it's a normal command execution, and if no command was given,
// configure a shell to run in 'login' mode. Otherwise, execute a command
// through the shell.
if c.RequestType == sshutils.SubsystemRequest {
if c.RequestType == "subsystem" {
switch c.Command {
case teleport.SFTPSubsystem:
case "sftp":
executable, err := os.Executable()
if err != nil {
return nil, trace.Wrap(err)
}
cmd.Path = executable
cmd.Args = []string{executable, teleport.SFTPSubCommand}
cmd.Args = []string{executable, reexecconstants.SFTPSubCommand}
isReexec = true
default:
return nil, trace.BadParameter("unsupported subsystem execution request %q", c.Command)
@@ -1243,7 +1264,7 @@ func buildCommand(c *ExecCommand, localUser *user.User, pamEnvironment []string)
// Create default environment for user.
env := &envutils.SafeEnv{
"LANG=en_US.UTF-8",
getDefaultEnvPath(localUser.Uid, defaultLoginDefsPath),
GetDefaultEnvPath(localUser.Uid),
"HOME=" + localUser.HomeDir,
"USER=" + c.Login,
"SHELL=" + shellPath,
@@ -1274,9 +1295,9 @@ func buildCommand(c *ExecCommand, localUser *user.User, pamEnvironment []string)
cmd.Env = *env
// set stdio. If a terminal was requested, the stdio fields all point to the same tty file.
cmd.Stdin = c.stdin
cmd.Stdout = c.stdout
cmd.Stderr = c.stderr
cmd.Stdin = c.Stdin
cmd.Stdout = c.Stdout
cmd.Stderr = c.Stderr
if c.Terminal {
cmd.SysProcAttr = &syscall.SysProcAttr{
@@ -1292,7 +1313,7 @@ func buildCommand(c *ExecCommand, localUser *user.User, pamEnvironment []string)
}
// Pass extra files for SFTP to grandchild.
if c.RequestType == sshutils.SubsystemRequest && c.Command == teleport.SFTPSubsystem {
if c.RequestType == "subsystem" && c.Command == "sftp" {
out := os.NewFile(FileTransferOutFile, "FileTransferOutFile")
if out == nil {
return nil, trace.NotFound("FileTransferOutFile file not found")
@@ -1344,7 +1365,7 @@ func buildCommand(c *ExecCommand, localUser *user.User, pamEnvironment []string)
// Perform OS-specific tweaks to the command.
if isReexec {
reexecCommandOSTweaks(&cmd)
CommandOSTweaks(&cmd)
} else {
userCommandOSTweaks(&cmd)
}
@@ -1352,98 +1373,6 @@ func buildCommand(c *ExecCommand, localUser *user.User, pamEnvironment []string)
return &cmd, nil
}
// ConfigureCommand creates a command fully configured to execute. This
// function is used by Teleport to re-execute itself and pass whatever data
// is need to the child to actually execute the shell.
func ConfigureCommand(ctx *ServerContext, extraFiles ...*os.File) (*exec.Cmd, error) {
// Create a os.Pipe and start copying over the payload to execute. While the
// pipe buffer is quite large (64k) some users have run into the pipe
// blocking writes on much smaller buffers (7k) leading to Teleport being
// unable to run some exec commands.
//
// To not depend on the OS implementation of a pipe, instead the copy should
// be non-blocking. The io.Copy will be closed when either when the child
// process has fully read in the payload or the process exits with an error
// (and closes all child file descriptors).
//
// See the below for details.
//
// https://man7.org/linux/man-pages/man7/pipe.7.html
cmdmsg, err := ctx.ExecCommand()
if err != nil {
return nil, trace.Wrap(err)
}
go copyCommand(ctx.CancelContext(), ctx.cmdw, cmdmsg)
// Find the Teleport executable and its directory on disk.
executable, err := os.Executable()
if err != nil {
return nil, trace.Wrap(err)
}
// The channel/request type determines the subcommand to execute.
var subCommand string
switch ctx.ExecType {
case teleport.NetworkingSubCommand:
subCommand = teleport.NetworkingSubCommand
default:
subCommand = teleport.ExecSubCommand
}
// Build the list of arguments to have Teleport re-exec itself. The "-d" flag
// is appended if Teleport is running in debug mode.
args := []string{executable, subCommand}
// build env for `teleport exec`
env := &envutils.SafeEnv{}
env.AddExecEnvironment()
// Build the "teleport exec" command.
cmd := &exec.Cmd{
Path: executable,
Args: args,
Env: *env,
ExtraFiles: []*os.File{
ctx.cmdr,
ctx.logw,
ctx.contr,
ctx.readyw,
ctx.killShellr,
},
}
// Add extra files if applicable.
if len(extraFiles) > 0 {
cmd.ExtraFiles = append(cmd.ExtraFiles, extraFiles...)
}
// Perform OS-specific tweaks to the command.
reexecCommandOSTweaks(cmd)
return cmd, nil
}
// copyCommand will copy the provided command to the child process over the
// pipe attached to the context.
func copyCommand(ctx context.Context, cmdw *os.File, cmdmsg *ExecCommand) {
defer func() {
err := cmdw.Close()
if err != nil {
slog.ErrorContext(ctx, "Failed to close command pipe", "error", err)
}
// Set to nil so the close in the context doesn't attempt to re-close.
cmdw = nil
}()
// Write command bytes to pipe. The child process will read the command
// to execute from this pipe.
if err := json.NewEncoder(cmdw).Encode(cmdmsg); err != nil {
slog.ErrorContext(ctx, "Failed to copy command over pipe", "error", err)
return
}
}
func coerceHomeDirError(usr *user.User, err error) error {
if os.IsNotExist(err) {
return trace.NotFound("home directory %q not found for user %q", usr.HomeDir, usr.Name)
@@ -1539,7 +1468,7 @@ func CheckHomeDir(localUser *user.User) (bool, error) {
// Build the "teleport exec" command.
cmd := &exec.Cmd{
Path: executable,
Args: []string{executable, teleport.CheckHomeDirSubCommand},
Args: []string{executable, reexecconstants.CheckHomeDirSubCommand},
Env: []string{"HOME=" + localUser.HomeDir},
Dir: rootDirectory,
SysProcAttr: &syscall.SysProcAttr{
@@ -1549,10 +1478,10 @@ func CheckHomeDir(localUser *user.User) (bool, error) {
}
// Perform OS-specific tweaks to the command.
reexecCommandOSTweaks(cmd)
CommandOSTweaks(cmd)
if err := cmd.Run(); err != nil {
if cmd.ProcessState.ExitCode() == teleport.RemoteCommandFailure {
if cmd.ProcessState.ExitCode() == reexecconstants.RemoteCommandFailure {
return false, trace.Wrap(err)
}
@@ -1569,7 +1498,7 @@ func (o *osWrapper) newParker(ctx context.Context, credential syscall.Credential
return trace.Wrap(err)
}
cmd := o.CommandContext(ctx, executable, teleport.ParkSubCommand)
cmd := o.CommandContext(ctx, executable, reexecconstants.ParkSubCommand)
cmd.SysProcAttr = &syscall.SysProcAttr{
Credential: &credential,
}
@@ -1588,9 +1517,9 @@ func (o *osWrapper) newParker(ctx context.Context, credential syscall.Credential
return nil
}
// waitForSignal will wait for the other side of the pipe to signal, if not
// WaitForSignal will wait for the other side of the pipe to signal, if not
// received, it will stop waiting and exit.
func waitForSignal(ctx context.Context, fd *os.File, timeout time.Duration) error {
func WaitForSignal(ctx context.Context, fd *os.File, timeout time.Duration) error {
waitCh := make(chan error, 1)
go func() {
// Reading from the file descriptor will block until it's closed.
@@ -1615,6 +1544,9 @@ func waitForSignal(ctx context.Context, fd *os.File, timeout time.Duration) erro
}
}
// TODO(espadolini): pass slog records in a fixed format to the parent process
// rather than handling the formatting here, so we can get rid of
// internal/logutils
func initLogger(name string, logWriter *os.File, cfg ExecLogConfig) {
fields, err := logutils.ValidateFields(cfg.ExtraFields)
if err != nil {
@@ -1629,14 +1561,47 @@ func initLogger(name string, logWriter *os.File, cfg ExecLogConfig) {
ConfiguredFields: fields,
Padding: cfg.Padding,
}))
slog.SetDefault(logger.With(teleport.ComponentKey, name))
slog.SetDefault(logger.With(logconstants.ComponentKey, name))
case "json":
logger := slog.New(logutils.NewSlogJSONHandler(logWriter, logutils.SlogJSONHandlerConfig{
Level: cfg.Level,
ConfiguredFields: fields,
}))
slog.SetDefault(logger.With(teleport.ComponentKey, name))
slog.SetDefault(logger.With(logconstants.ComponentKey, name))
default:
return
}
}
// isUseOfClosedNetworkError is [utils.IsUseOfClosedNetworkError].
func isUseOfClosedNetworkError(err error) bool {
if err == nil {
return false
}
return errors.Is(err, net.ErrClosed) || strings.Contains(err.Error(), apiconstants.UseOfClosedNetworkConnection)
}
// isFailedToSendCloseNotifyError is [utils.IsFailedToSendCloseNotifyError].
func isFailedToSendCloseNotifyError(err error) bool {
if err == nil {
return false
}
return strings.Contains(err.Error(), apiconstants.FailedToSendCloseNotify)
}
// isOKNetworkError is [utils.IsOKNetworkError].
func isOKNetworkError(err error) bool {
// trace.Aggregate contains at least one error and all the errors are
// non-nil
var a trace.Aggregate
if errors.As(trace.Unwrap(err), &a) {
for _, err := range a.Errors() {
if !isOKNetworkError(err) {
return false
}
}
return true
}
return errors.Is(err, io.EOF) || isUseOfClosedNetworkError(err) || isFailedToSendCloseNotifyError(err)
}
@@ -18,7 +18,7 @@
* along with this program. If not, see <http://www.gnu.org/licenses/>.
*/
package srv
package reexec
import (
"os"
@@ -64,7 +64,7 @@ func init() {
// passed to reexecCommandOSTweaks, if not empty.
var reexecPath string
func reexecCommandOSTweaks(cmd *exec.Cmd) {
func CommandOSTweaks(cmd *exec.Cmd) {
if cmd.SysProcAttr == nil {
cmd.SysProcAttr = new(syscall.SysProcAttr)
}
@@ -85,7 +85,7 @@ func reexecCommandOSTweaks(cmd *exec.Cmd) {
// we should rework the parker to block on a pipe so it can exit when its parent
// is terminated
func parkerCommandOSTweaks(cmd *exec.Cmd) {
reexecCommandOSTweaks(cmd)
CommandOSTweaks(cmd)
// parker processes can leak if their PDEATHSIG is SIGQUIT, otherwise we
// could just use reexecCommandOSTweaks
@@ -18,13 +18,13 @@
* along with this program. If not, see <http://www.gnu.org/licenses/>.
*/
package srv
package reexec
import (
"os/exec"
)
func reexecCommandOSTweaks(cmd *exec.Cmd) {}
func CommandOSTweaks(cmd *exec.Cmd) {}
func parkerCommandOSTweaks(cmd *exec.Cmd) {}
+359
View File
@@ -0,0 +1,359 @@
// Teleport
// Copyright (C) 2026 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package reexec
import (
"context"
"crypto/rand"
"errors"
"fmt"
"io"
"os"
"os/exec"
"os/user"
"path/filepath"
"strconv"
"syscall"
"testing"
"github.com/gravitational/trace"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
apiconstants "github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/session/host"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
type stubUser struct {
gid string
uid string
groupIDS []string
}
func (s *stubUser) GID() string {
return s.gid
}
func (s *stubUser) UID() string {
return s.uid
}
func (s *stubUser) GroupIds() ([]string, error) {
return s.groupIDS, nil
}
func TestStartNewParker(t *testing.T) {
currentUser, err := user.Current()
require.NoError(t, err)
currentUID, err := strconv.ParseUint(currentUser.Uid, 10, 32)
require.NoError(t, err)
currentGID, err := strconv.ParseUint(currentUser.Gid, 10, 32)
require.NoError(t, err)
t.Parallel()
type args struct {
credential *syscall.Credential
loginAsUser string
localUser *stubUser
}
tests := []struct {
name string
args args
newOsPack func(t *testing.T) (*osWrapper, func())
wantErr require.ErrorAssertionFunc
}{
{
name: "empty credentials does nothing",
wantErr: require.NoError,
newOsPack: func(t *testing.T) (*osWrapper, func()) {
return &osWrapper{}, func() {}
},
},
{
name: "missing Teleport group returns no error",
wantErr: require.NoError,
newOsPack: func(t *testing.T) (*osWrapper, func()) {
return &osWrapper{
LookupGroup: func(name string) (*user.Group, error) {
require.Equal(t, apiconstants.TeleportDropGroup, name)
return nil, user.UnknownGroupError(apiconstants.TeleportDropGroup)
},
}, func() {}
},
},
{
name: "different group doesn't start parker",
wantErr: require.NoError,
newOsPack: func(t *testing.T) (*osWrapper, func()) {
return &osWrapper{
LookupGroup: func(name string) (*user.Group, error) {
require.Equal(t, apiconstants.TeleportDropGroup, name)
return &user.Group{Gid: "1234"}, nil
},
CommandContext: func(ctx context.Context, name string, arg ...string) *exec.Cmd {
require.FailNow(t, "CommandContext should not be called")
return nil
},
}, func() {}
},
args: args{
credential: &syscall.Credential{Gid: 1000},
localUser: &stubUser{
uid: "1001",
gid: "1003",
groupIDS: []string{"1003"},
},
},
},
{
name: "parker is started",
wantErr: require.NoError,
newOsPack: func(t *testing.T) (*osWrapper, func()) {
parkerStarted := false
return &osWrapper{
LookupGroup: func(name string) (*user.Group, error) {
require.Equal(t, apiconstants.TeleportDropGroup, name)
return &user.Group{Gid: currentUser.Gid}, nil
},
CommandContext: func(ctx context.Context, name string, arg ...string) *exec.Cmd {
require.NotNil(t, ctx)
require.Len(t, arg, 1)
require.Equal(t, reexecconstants.ParkSubCommand, arg[0])
parkerStarted = true
return exec.CommandContext(ctx, name, arg...)
},
LookupUser: func(username string) (*user.User, error) {
return &user.User{Uid: currentUser.Uid, Gid: currentUser.Gid}, nil
},
}, func() {
require.True(t, parkerStarted, "parker process didn't start")
}
},
args: args{
credential: &syscall.Credential{
Uid: uint32(currentUID),
Gid: uint32(currentGID),
// Changing to false causes "fork/exec /proc/self/exe: operation not permitted"
// to be returned when creating the park process.
NoSetGroups: true,
},
localUser: &stubUser{
uid: currentUser.Uid,
gid: currentUser.Gid,
groupIDS: []string{currentUser.Gid},
},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
osPack, assertExpected := tt.newOsPack(t)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel) // cancel to stop the park process.
err := osPack.startNewParker(ctx, tt.args.credential, tt.args.loginAsUser, tt.args.localUser)
tt.wantErr(t, err, fmt.Sprintf("startNewParker(%v, %+v, %v, %+v)", ctx, tt.args.credential, tt.args.loginAsUser, tt.args.localUser))
assertExpected()
})
}
}
func TestRootCheckHomeDir(t *testing.T) {
requireRoot(t)
// this test manipulates global state, ensure we're not going to run it in
// parallel with something else
t.Setenv("foo", "bar")
tmp := t.TempDir()
require.NoError(t, os.Chmod(filepath.Dir(tmp), 0777))
require.NoError(t, os.Chmod(tmp, 0777))
home := filepath.Join(tmp, "home")
noAccess := filepath.Join(tmp, "no_access")
file := filepath.Join(tmp, "file")
notFound := filepath.Join(tmp, "not_found")
require.NoError(t, os.Mkdir(home, 0700))
require.NoError(t, os.Mkdir(noAccess, 0700))
_, err := os.Create(file)
require.NoError(t, err)
login := generateLocalUsername(t)
_, err = host.UserAdd(login, nil, host.UserOpts{Home: home})
require.NoError(t, err)
t.Cleanup(func() {
// change back to accessible home so deletion works
changeHomeDir(t, login, home)
_, err := host.UserDel(login)
require.NoError(t, err)
})
testUser, err := user.Lookup(login)
require.NoError(t, err)
uid, err := strconv.Atoi(testUser.Uid)
require.NoError(t, err)
gid, err := strconv.Atoi(testUser.Gid)
require.NoError(t, err)
require.NoError(t, os.Chown(home, uid, gid))
require.NoError(t, os.Chown(file, uid, gid))
hasAccess, err := CheckHomeDir(testUser)
require.NoError(t, err)
require.True(t, hasAccess)
changeHomeDir(t, login, file)
hasAccess, err = CheckHomeDir(testUser)
require.NoError(t, err)
require.False(t, hasAccess)
changeHomeDir(t, login, notFound)
hasAccess, err = CheckHomeDir(testUser)
require.NoError(t, err)
require.False(t, hasAccess)
changeHomeDir(t, login, noAccess)
hasAccess, err = CheckHomeDir(testUser)
require.NoError(t, err)
require.False(t, hasAccess)
}
func changeHomeDir(t *testing.T, username, home string) {
usermodBin, err := exec.LookPath("usermod")
assert.NoError(t, err, "usermod binary must be present")
cmd := exec.Command(usermodBin, "--home", home, username)
_, err = cmd.CombinedOutput()
assert.NoError(t, err, "changing home should not error")
assert.Equal(t, 0, cmd.ProcessState.ExitCode(), "changing home should exit 0")
}
func TestRootOpenFileAsUser(t *testing.T) {
requireRoot(t)
euid := os.Geteuid()
egid := os.Getegid()
username := "processing-user"
arg := os.Args[1]
os.Args[1] = reexecconstants.ExecSubCommand
defer func() {
os.Args[1] = arg
}()
_, err := host.UserAdd(username, nil, host.UserOpts{})
require.NoError(t, err)
t.Cleanup(func() {
_, err := host.UserDel(username)
require.NoError(t, err)
})
tmp := t.TempDir()
testFile := filepath.Join(tmp, "testfile")
fileContent := "one does not simply open without permission"
err = os.WriteFile(testFile, []byte(fileContent), 0777)
require.NoError(t, err)
testUser, err := user.Lookup(username)
require.NoError(t, err)
// no access
file, err := openFileAsUser(testUser, testFile)
require.True(t, trace.IsAccessDenied(err))
require.Nil(t, file)
// ensure we fallback to root after
file, err = os.Open(testFile)
require.NoError(t, err)
require.NotNil(t, file)
file.Close()
// has access
uid, err := strconv.Atoi(testUser.Uid)
require.NoError(t, err)
gid, err := strconv.Atoi(testUser.Gid)
require.NoError(t, err)
err = os.Chown(filepath.Dir(tmp), uid, gid)
require.NoError(t, err)
err = os.Chown(tmp, uid, gid)
require.NoError(t, err)
err = os.Chown(testFile, uid, gid)
require.NoError(t, err)
file, err = openFileAsUser(testUser, testFile)
require.NoError(t, err)
require.NotNil(t, file)
data, err := io.ReadAll(file)
file.Close()
require.NoError(t, err)
require.Equal(t, fileContent, string(data))
// not exist
file, err = openFileAsUser(testUser, filepath.Join(tmp, "no_exist"))
require.ErrorIs(t, err, os.ErrNotExist)
require.Nil(t, file)
require.Equal(t, euid, os.Geteuid())
require.Equal(t, egid, os.Getegid())
}
// requireRoot is [testutils.RequireRoot] but inlined.
func requireRoot(tb testing.TB) {
tb.Helper()
if os.Geteuid() != 0 {
tb.Skip("This test will be skipped because tests are not being run as root.")
}
}
func generateUsername(tb testing.TB) string {
suffix := make([]byte, 8)
_, err := rand.Read(suffix)
require.NoError(tb, err)
return fmt.Sprintf("teleport-%x", suffix)
}
// generateLocalUsername is [testutils.GenerateLocalUsername] but inlined.
func generateLocalUsername(tb testing.TB) string {
const maxAttempts = 10
for range maxAttempts {
login := generateUsername(tb)
_, err := user.Lookup(login)
if errors.Is(err, user.UnknownUserError(login)) {
return login
}
require.NoError(tb, err)
}
tb.Fatalf("Unable to generate unused username after %v attempts", maxAttempts)
return ""
}
@@ -0,0 +1,59 @@
// Teleport
// Copyright (C) 2026 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package reexecconstants
const (
// ExecSubCommand is the sub-command Teleport uses to re-exec itself for
// command execution (exec and shells).
ExecSubCommand = "exec"
// NetworkingSubCommand is the sub-command Teleport uses to re-exec itself
// for networking operations. e.g. local/remote port forwarding, agent forwarding,
// or x11 forwarding.
NetworkingSubCommand = "networking"
// CheckHomeDirSubCommand is the sub-command Teleport uses to re-exec itself
// to check if the user's home directory exists.
CheckHomeDirSubCommand = "checkhomedir"
// ParkSubCommand is the sub-command Teleport uses to re-exec itself as a
// specific UID to prevent the matching user from being deleted before
// spawning the intended child process.
ParkSubCommand = "park"
// SFTPSubCommand is the sub-command Teleport uses to re-exec itself to
// handle SFTP connections.
SFTPSubCommand = "sftp"
)
const (
// RemoteCommandSuccess is returned when a command has successfully executed.
RemoteCommandSuccess = 0
// RemoteCommandFailure is returned when a command has failed to execute and
// we don't have another status code for it.
RemoteCommandFailure = 255
// HomeDirNotFound is returned when the "teleport checkhomedir" command cannot
// find the user's home directory.
HomeDirNotFound = 254
// HomeDirNotAccessible is returned when the "teleport checkhomedir" command has
// found the user's home directory, but the user does NOT have permissions to
// access it.
HomeDirNotAccessible = 253
// UnexpectedCredentials is returned when a command is no longer running with the expected
// credentials.
UnexpectedCredentials = 252
)
+37
View File
@@ -0,0 +1,37 @@
// Teleport
// Copyright (C) 2026 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package reexec
import (
"errors"
"os/user"
"strings"
"syscall"
)
// isUnknownGroupError returns whether the error from LookupGroup is an unknown group error.
//
// LookupGroup is supposed to return an UnknownGroupError, but due to an existing issue
// may instead return a generic "no such file or directory" error when sssd is installed
// or "no such process" as Go std library just forwards errors returned by getgrpnam_r.
// See github issue - https://github.com/golang/go/issues/40334
func isUnknownGroupError(err error, groupName string) bool {
return errors.Is(err, user.UnknownGroupError(groupName)) ||
errors.Is(err, user.UnknownGroupIdError(groupName)) ||
strings.HasSuffix(err.Error(), syscall.ENOENT.Error()) ||
strings.HasSuffix(err.Error(), syscall.ESRCH.Error())
}
+11 -10
View File
@@ -52,12 +52,13 @@ import (
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/service"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/sshutils/scp"
"github.com/gravitational/teleport/lib/tpm"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/versioncontrol"
"github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
"github.com/gravitational/teleport/session/selinux"
)
@@ -110,11 +111,11 @@ func Run(options Options) (app *kingpin.Application, executedCommand string, con
join := app.Command("join", "Join a Teleport cluster without running the Teleport daemon.")
joinOpenSSH := join.Command("openssh", "Join an SSH server to a Teleport cluster.")
scpc := app.Command("scp", "Server-side implementation of SCP.").Hidden()
sftp := app.Command(teleport.SFTPSubCommand, "Server-side implementation of SFTP.").Hidden()
exec := app.Command(teleport.ExecSubCommand, "Used internally by Teleport to re-exec itself to run a command.").Hidden()
networking := app.Command(teleport.NetworkingSubCommand, "Used internally by Teleport to re-exec itself to handle networking requests.").Hidden()
checkHomeDir := app.Command(teleport.CheckHomeDirSubCommand, "Used internally by Teleport to re-exec itself to check access to a directory.").Hidden()
park := app.Command(teleport.ParkSubCommand, "Used internally by Teleport to re-exec itself to do nothing.").Hidden()
sftp := app.Command(reexecconstants.SFTPSubCommand, "Server-side implementation of SFTP.").Hidden()
exec := app.Command(reexecconstants.ExecSubCommand, "Used internally by Teleport to re-exec itself to run a command.").Hidden()
networking := app.Command(reexecconstants.NetworkingSubCommand, "Used internally by Teleport to re-exec itself to handle networking requests.").Hidden()
checkHomeDir := app.Command(reexecconstants.CheckHomeDirSubCommand, "Used internally by Teleport to re-exec itself to check access to a directory.").Hidden()
park := app.Command(reexecconstants.ParkSubCommand, "Used internally by Teleport to re-exec itself to do nothing.").Hidden()
app.HelpFlag.Short('h')
// define start flags:
@@ -747,13 +748,13 @@ Examples:
dumpFlags.Roles = defaults.RoleNode
err = onConfigDump(dumpFlags)
case exec.FullCommand():
srv.RunAndExit(teleport.ExecSubCommand)
reexec.RunAndExit(reexecconstants.ExecSubCommand)
case networking.FullCommand():
srv.RunAndExit(teleport.NetworkingSubCommand)
reexec.RunAndExit(reexecconstants.NetworkingSubCommand)
case checkHomeDir.FullCommand():
srv.RunAndExit(teleport.CheckHomeDirSubCommand)
reexec.RunAndExit(reexecconstants.CheckHomeDirSubCommand)
case park.FullCommand():
srv.RunAndExit(teleport.ParkSubCommand)
reexec.RunAndExit(reexecconstants.ParkSubCommand)
case waitNoResolveCmd.FullCommand():
err = onWaitNoResolve(waitFlags)
case waitDurationCmd.FullCommand():
+2 -2
View File
@@ -50,10 +50,10 @@ import (
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/service"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/log/logtest"
"github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/tool/teleport/common"
)
@@ -68,7 +68,7 @@ const StaticToken = "test-static-token"
func init() {
// If the test is re-executing itself, execute the command that comes over
// the pipe. Used to test tsh ssh and tsh scp commands.
if srv.IsReexec() {
if reexec.IsReexec() {
common.Run(common.Options{Args: os.Args[1:]})
return
}
+6 -5
View File
@@ -94,13 +94,14 @@ import (
"github.com/gravitational/teleport/lib/service"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/log/logtest"
"github.com/gravitational/teleport/lib/utils/testutils"
"github.com/gravitational/teleport/lib/utils/testutils/golden"
"github.com/gravitational/teleport/session/networking/x11"
"github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
"github.com/gravitational/teleport/tool/common"
testserver "github.com/gravitational/teleport/tool/teleport/testenv"
)
@@ -239,8 +240,8 @@ func handleReexec() {
}
// Re-exec teleport commands. Used to test tsh ssh command.
if srv.IsReexec() {
srv.RunAndExit(os.Args[1])
if reexec.IsReexec() {
reexec.RunAndExit(os.Args[1])
}
}
@@ -8452,7 +8453,7 @@ func TestReexecErrorPropagation(t *testing.T) {
var exitCodeErr *common.ExitCodeError
require.ErrorAs(t, err, &exitCodeErr)
require.Equal(t, teleport.RemoteCommandFailure, exitCodeErr.Code)
require.Equal(t, reexecconstants.RemoteCommandFailure, exitCodeErr.Code)
// Check for exact match to catch regressions with new lines.
require.Equal(t, unknownUserReexecError, stdout)
@@ -8499,7 +8500,7 @@ func TestReexecErrorPropagation(t *testing.T) {
var exitCodeErr *common.ExitCodeError
require.ErrorAs(t, err, &exitCodeErr)
require.Equal(t, teleport.RemoteCommandFailure, exitCodeErr.Code)
require.Equal(t, reexecconstants.RemoteCommandFailure, exitCodeErr.Code)
// Check for exact match to catch regressions with new lines.
require.Equal(t, contextualReexecErrorMessage, stdout)