diff --git a/internal/flows/proxy_flow.go b/internal/flows/proxy_flow.go index a962897..78ae992 100644 --- a/internal/flows/proxy_flow.go +++ b/internal/flows/proxy_flow.go @@ -3,11 +3,9 @@ package flows import ( "context" "fmt" - "io" "os" "os/exec" "path/filepath" - "sync" "time" "github.com/safedep/dry/log" @@ -15,17 +13,12 @@ import ( "github.com/safedep/pmg/config" "github.com/safedep/pmg/guard" "github.com/safedep/pmg/internal/audit" - "github.com/safedep/pmg/internal/pty" "github.com/safedep/pmg/internal/runner" - "github.com/safedep/pmg/internal/shim" "github.com/safedep/pmg/internal/ui" "github.com/safedep/pmg/packagemanager" "github.com/safedep/pmg/proxy" "github.com/safedep/pmg/proxy/certmanager" "github.com/safedep/pmg/proxy/interceptors" - "github.com/safedep/pmg/sandbox" - "github.com/safedep/pmg/sandbox/executor" - "github.com/safedep/pmg/usefulerror" ) type proxyFlow struct { @@ -208,26 +201,67 @@ func (f *proxyFlow) Run(ctx context.Context, args []string, parsedCmd *packagema log.Infof("Proxy server started on %s", proxyAddr) log.Infof("Running %s with proxy protection enabled", f.pm.Name()) - proxyEnv := f.setupEnvForProxy(proxyAddr, caCertPath) + executionError := runner.ExecuteWithOptions(ctx, parsedCmd, runner.ExecuteOptions{ + PackageManagerName: f.pm.Name(), + DryRun: cfg.DryRun, + Mode: runner.ExecutionModeAuto, + EnvOverrides: f.setupEnvForProxy(proxyAddr, caCertPath), + DirectEnvOverrides: []string{"CI=true"}, + BeforeDirectRun: func() error { + log.Debugf("Executing proxy for non interactive TTY") - // Resolve the real package manager binary by searching PATH with ~/.pmg/bin - // stripped out. Without this, exec.CommandContext resolves to the shim script - // (because ~/.pmg/bin is still in the current process's PATH), causing - // infinite recursion: shim → pmg → shim → pmg → ... - realBinary, err := shim.ResolveRealBinary(parsedCmd.Command.Exe) - if err != nil { - return fmt.Errorf("failed to resolve real %s binary: %w", parsedCmd.Command.Exe, err) - } - parsedCmd.Command.Exe = realBinary + interaction.GetConfirmationOnMalware = func(_ []*analyzer.PackageVersionAnalysisResult) (bool, error) { + return false, nil + } - var executionError error - if pty.IsInteractiveTerminal() { - // Execute the package manager command with proxy environment variables - executionError = f.executeWithProxy(ctx, parsedCmd, proxyEnv, confirmationChan, interaction) - } else { - // Execute the package manager command with proxy environment variables for non PTY or non-interactive TTY - executionError = f.executeWithProxyForNonInteractiveTTY(ctx, parsedCmd, proxyEnv, confirmationChan, interaction) - } + go interceptors.HandleConfirmationRequests(confirmationChan, interaction, nil) + return nil + }, + PreparePTYSession: func(runtime *runner.PTYRuntime) error { + log.Debugf("Executing proxy for interactive TTY") + + interaction.GetConfirmationOnMalware = func(malwarePackages []*analyzer.PackageVersionAnalysisResult) (bool, error) { + return ui.GetConfirmationOnMalwareWithReader(malwarePackages, interaction.Reader()) + } + + go interceptors.HandleConfirmationRequests( + confirmationChan, + interaction, + &interceptors.ConfirmationHook{ + BeforeInteraction: func(_ []*analyzer.PackageVersionAnalysisResult) error { + runtime.OutputRouter.Pause() + + if err := runtime.Session.SetCookedMode(); err != nil { + return fmt.Errorf("failed to set cooked mode: %w", err) + } + + if _, err := fmt.Fprint(os.Stdout, "\033[?25h"); err != nil { + log.Warnf("failed to force cursor visible: %v", err) + } + + runtime.InputRouter.RouteToPrompt(runtime.PromptWriter) + interaction.SetInput(runtime.PromptReader) + + return nil + }, + AfterInteraction: func(_ []*analyzer.PackageVersionAnalysisResult, _ bool) error { + runtime.InputRouter.RouteToPTY() + + if err := runtime.Session.SetRawMode(); err != nil { + return fmt.Errorf("failed to set raw mode: %w", err) + } + + interaction.SetInput(nil) + runtime.OutputRouter.Resume() + + return nil + }, + }, + ) + + return nil + }, + }) // Populate report data from stats collector stats := statsCollector.GetStats() @@ -341,8 +375,7 @@ func (f *proxyFlow) setupEnvForProxy(proxyAddr, caCertPath string) []string { noProxyList := "localhost,127.0.0.1,[::1]" - env := shim.FilterPMGFromEnv(os.Environ()) - env = append(env, + return []string{ "NODE_USE_ENV_PROXY=1", fmt.Sprintf("HTTP_PROXY=%s", proxyURL), fmt.Sprintf("HTTPS_PROXY=%s", proxyURL), @@ -356,223 +389,5 @@ func (f *proxyFlow) setupEnvForProxy(proxyAddr, caCertPath string) []string { fmt.Sprintf("PIP_CERT=%s", caCertPath), fmt.Sprintf("PIP_PROXY=%s", proxyURL), "PIP_RETRIES=0", - ) - - return env -} - -// executeWithProxyForNonInteractiveTTY runs the command without PTY (for CI/non-interactive environments) -func (f *proxyFlow) executeWithProxyForNonInteractiveTTY( - ctx context.Context, - parsedCmd *packagemanager.ParsedCommand, - env []string, - confirmationChan chan *interceptors.ConfirmationRequest, - interaction *guard.PackageManagerGuardInteraction, -) error { - log.Debugf("Executing proxy for non interactive TTY") - - // For non-interactive terminals, we enforce suspicious packages as malicious - interaction.GetConfirmationOnMalware = func(malwarePackages []*analyzer.PackageVersionAnalysisResult) (bool, error) { - return false, nil } - - cmd := exec.CommandContext(ctx, parsedCmd.Command.Exe, parsedCmd.Command.Args...) - cmd.Env = append(env, "CI=true") - cmd.Stdin = os.Stdin - cmd.Stdout = os.Stdout - cmd.Stderr = os.Stderr - - go interceptors.HandleConfirmationRequests( - confirmationChan, - interaction, - nil, - ) - - result, err := executor.ApplySandbox(ctx, cmd, f.pm.Name()) - if err != nil { - return fmt.Errorf("failed to apply sandbox: %w", err) - } - - defer func() { - err := result.Close() - if err != nil { - log.Errorf("failed to close sandbox: %v", err) - } - }() - - // Only run the command if the sandbox didn't already execute it - if result.ShouldRun() { - log.Debugf("Running command with args: %s: %v", cmd.Path, cmd.Args[1:]) - - err = cmd.Run() - if err != nil { - return f.handlePackageManagerExecutionError(err, result) - } - } - - log.Debugf("Command completed successfully") - return nil -} - -// executeWithProxy executes the package manager command with proxy environment variables. -func (f *proxyFlow) executeWithProxy( - ctx context.Context, - parsedCmd *packagemanager.ParsedCommand, - env []string, - confirmationChan chan *interceptors.ConfirmationRequest, - interaction *guard.PackageManagerGuardInteraction, -) error { - log.Debugf("Executing proxy for interactive TTY") - - // Set the confirmation handler to use the interaction's reader - // This allows PTY input routing during proxy mode - interaction.GetConfirmationOnMalware = func(malwarePackages []*analyzer.PackageVersionAnalysisResult) (bool, error) { - return ui.GetConfirmationOnMalwareWithReader(malwarePackages, interaction.Reader()) - } - - cmd := exec.CommandContext(ctx, parsedCmd.Command.Exe, parsedCmd.Command.Args...) - result, err := executor.ApplySandbox(ctx, cmd, f.pm.Name()) - if err != nil { - return fmt.Errorf("failed to apply sandbox: %w", err) - } - - defer func() { - if err := result.Close(); err != nil { - log.Errorf("failed to close sandbox: %v", err) - } - }() - - if !result.ShouldRun() { - return usefulerror.Useful(). - Wrap(fmt.Errorf("sandbox not supported for PTY sessions")). - WithHumanError("Sandbox executed command cannot be used with PTY session. Please use non-interactive TTY mode instead.") - } - - // Extract the command executable and arguments from the sandboxed command - // for use to create the PTY session. - cmdExe := cmd.Path - cmdArgs := cmd.Args[1:] - - log.Debugf("Running command with args: %s: %v", cmdExe, cmdArgs) - - // Create the PTY session with the sandbox command - // This is not compatible with sandbox that executes the command directly within the sandbox - // because internally we use ptyx.Spawn() to create the process with PTY support. - sessionConfig := pty.NewSessionConfig(cmdExe, cmdArgs, env) - - sess, err := pty.NewSession(ctx, sessionConfig) - if err != nil { - return fmt.Errorf("failed to create pty session: %w", err) - } - defer sess.Close() - - outputRouter, err := pty.NewOutputRouter(os.Stdout) - if err != nil { - return fmt.Errorf("failed to create output router: %w", err) - } - - var wg sync.WaitGroup - wg.Go(func() { - if _, err := io.Copy(outputRouter, sess.PtyReader()); err != nil { - log.Errorf("failed to copy output: %v", err) - } - }) - - inputRouter, err := pty.NewInputRouter(sess.PtyWriter()) - if err != nil { - return fmt.Errorf("failed to create input router: %w", err) - } - - promptReader, promptWriter := io.Pipe() - defer func() { - promptWriter.Close() - promptReader.Close() - }() - - // Note: This goroutine cannot be cleanly cancelled because os.Stdin.Read() is - // a blocking syscall that doesn't support timeouts or cancellation. This is a - // known limitation. The goroutine will exit when the process terminates, which - // is acceptable for a CLI tool. For long-running servers, stdin reading should - // be handled differently. - go inputRouter.ReadLoop(os.Stdin) - - go interceptors.HandleConfirmationRequests( - confirmationChan, - interaction, - &interceptors.ConfirmationHook{ - BeforeInteraction: func(_ []*analyzer.PackageVersionAnalysisResult) error { - // Pause printing the child output - outputRouter.Pause() - - // Restore "Cooked" mode so user can type normally with echo - if err := sess.SetCookedMode(); err != nil { - return fmt.Errorf("failed to set cooked mode: %w", err) - } - - // Force cursor visible (ANSI escape sequence) - fmt.Fprint(os.Stdout, "\033[?25h") - - // Switch Input: Route keystrokes to the Prompt Pipe - inputRouter.RouteToPrompt(promptWriter) - - // Inject the Reader into the Interaction for the confirmation prompt - interaction.SetInput(promptReader) - - return nil - }, - AfterInteraction: func(_ []*analyzer.PackageVersionAnalysisResult, _ bool) error { - // Switch input back to PTY - inputRouter.RouteToPTY() - - // Restore "Raw" mode for the PTY - if err := sess.SetRawMode(); err != nil { - return fmt.Errorf("failed to set raw mode: %w", err) - } - - // Clear the interaction input (back to default) - interaction.SetInput(nil) - - // Flush buffered output and resume live output - outputRouter.Resume() - - return nil - }, - }, - ) - - // sessionError may contain the exit code of the command if the command exited with a non-zero code. - sessionError := sess.Wait() - - // Wait for the routers to copy all the remaining data - wg.Wait() - - if err := promptReader.Close(); err != nil { - log.Errorf("failed to close prompt reader: %v", err) - } - - if err := promptWriter.Close(); err != nil { - log.Errorf("failed to close prompt writer: %v", err) - } - - if err := sess.Close(); err != nil { - log.Errorf("failed to close session: %v", err) - } - - if sessionError != nil { - return f.handlePackageManagerExecutionError(sessionError, result) - } - - return nil -} - -func (f *proxyFlow) handlePackageManagerExecutionError(err error, result *sandbox.ExecutionResult) error { - if exitErr, ok := err.(*exec.ExitError); ok { - return executor.WrapCommandExecutionError(err, result, exitErr.ExitCode()) - } - - if sessionError, ok := err.(*pty.ExitError); ok { - return executor.WrapCommandExecutionError(sessionError, result, sessionError.Code) - } - - return executor.WrapCommandExecutionError(err, result, -1) } diff --git a/internal/pty/input_read_unix.go b/internal/pty/input_read_unix.go new file mode 100644 index 0000000..2d9eb0f --- /dev/null +++ b/internal/pty/input_read_unix.go @@ -0,0 +1,57 @@ +//go:build !windows + +package pty + +import ( + "context" + "errors" + "io" + "os" + + "golang.org/x/sys/unix" +) + +func readInput(ctx context.Context, src io.Reader, buf []byte) (int, error) { + file, ok := src.(*os.File) + if !ok { + select { + case <-ctx.Done(): + return 0, ctx.Err() + default: + return src.Read(buf) + } + } + + fd := int32(file.Fd()) + pollFds := []unix.PollFd{{Fd: fd, Events: unix.POLLIN}} + + for { + select { + case <-ctx.Done(): + return 0, ctx.Err() + default: + } + + n, err := unix.Poll(pollFds, 100) + if err != nil { + if errors.Is(err, unix.EINTR) { + continue + } + + return 0, err + } + + if n == 0 { + continue + } + + revents := pollFds[0].Revents + if revents&(unix.POLLERR|unix.POLLHUP|unix.POLLNVAL) != 0 { + return 0, io.EOF + } + + if revents&unix.POLLIN != 0 { + return file.Read(buf) + } + } +} diff --git a/internal/pty/input_read_windows.go b/internal/pty/input_read_windows.go new file mode 100644 index 0000000..2a0f163 --- /dev/null +++ b/internal/pty/input_read_windows.go @@ -0,0 +1,17 @@ +//go:build windows + +package pty + +import ( + "context" + "io" +) + +func readInput(ctx context.Context, src io.Reader, buf []byte) (int, error) { + select { + case <-ctx.Done(): + return 0, ctx.Err() + default: + return src.Read(buf) + } +} diff --git a/internal/pty/router.go b/internal/pty/router.go index b877a10..59bb9e6 100644 --- a/internal/pty/router.go +++ b/internal/pty/router.go @@ -2,6 +2,7 @@ package pty import ( "bytes" + "context" "io" "sync" "sync/atomic" @@ -86,20 +87,31 @@ func NewInputRouter(ptyWriter io.Writer) (*InputRouter, error) { // // This function blocks until src returns an error (e.g., EOF). func (r *InputRouter) ReadLoop(src io.Reader) { + r.ReadLoopContext(context.Background(), src) +} + +// ReadLoopContext continuously reads from src and routes data to the current +// destination until src returns an error or ctx is cancelled. +func (r *InputRouter) ReadLoopContext(ctx context.Context, src io.Reader) { buf := make([]byte, 1024) for { - nr, err := src.Read(buf) + nr, err := readInput(ctx, src, buf) if err != nil { return } - // Check where to route the data + if nr == 0 { + continue + } + if dest := r.dest.Load(); dest != nil { - // Send confirmation prompt response to the pipe. (PMG) - _, _ = dest.w.Write(buf[:nr]) + if _, err := dest.w.Write(buf[:nr]); err != nil { + return + } } else { - // Send response to the child PTY. - _, _ = r.defaultDst.Write(buf[:nr]) + if _, err := r.defaultDst.Write(buf[:nr]); err != nil { + return + } } } } diff --git a/internal/pty/router_test.go b/internal/pty/router_test.go new file mode 100644 index 0000000..ca8cb83 --- /dev/null +++ b/internal/pty/router_test.go @@ -0,0 +1,84 @@ +//go:build !windows + +package pty + +import ( + "bytes" + "context" + "errors" + "os" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestInputRouterReadLoopContextStopsOnCancelForFileReader(t *testing.T) { + reader, writer, err := os.Pipe() + require.NoError(t, err) + defer func() { + require.NoError(t, reader.Close()) + }() + defer func() { + require.NoError(t, writer.Close()) + }() + + router, err := NewInputRouter(&bytes.Buffer{}) + require.NoError(t, err) + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + defer close(done) + router.ReadLoopContext(ctx, reader) + }() + + cancel() + + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} + +func TestInputRouterReadLoopContextStopsOnWriteError(t *testing.T) { + reader, writer, err := os.Pipe() + require.NoError(t, err) + defer func() { + require.NoError(t, reader.Close()) + }() + defer func() { + require.NoError(t, writer.Close()) + }() + + router, err := NewInputRouter(errorWriter{}) + require.NoError(t, err) + + done := make(chan struct{}) + go func() { + defer close(done) + router.ReadLoopContext(context.Background(), reader) + }() + + _, err = writer.Write([]byte("x")) + require.NoError(t, err) + + require.Eventually(t, func() bool { + select { + case <-done: + return true + default: + return false + } + }, time.Second, 10*time.Millisecond) +} + +type errorWriter struct{} + +func (errorWriter) Write(_ []byte) (int, error) { + return 0, errors.New("write failed") +} diff --git a/internal/runner/execute.go b/internal/runner/execute.go index 5f6ffdc..59942d4 100644 --- a/internal/runner/execute.go +++ b/internal/runner/execute.go @@ -3,23 +3,77 @@ package runner import ( "context" "fmt" + "io" "os" "os/exec" + "strings" + "sync" "github.com/safedep/dry/log" + "github.com/safedep/pmg/internal/pty" "github.com/safedep/pmg/internal/shim" "github.com/safedep/pmg/packagemanager" + "github.com/safedep/pmg/sandbox" "github.com/safedep/pmg/sandbox/executor" + "github.com/safedep/pmg/usefulerror" ) +type ExecutionMode int + +const ( + ExecutionModeDirect ExecutionMode = iota + ExecutionModePTY + ExecutionModeAuto +) + +type ExecuteOptions struct { + PackageManagerName string + DryRun bool + EnvOverrides []string + DirectEnvOverrides []string + PTYEnvOverrides []string + Mode ExecutionMode + + // BeforeDirectRun runs after command/env construction and before sandbox + // application for non-PTY execution. Use this for setup that must exist even + // when a sandbox implementation executes the child inside ApplySandbox. + BeforeDirectRun func() error + + // PreparePTYSession runs after the PTY session and routers are created, but + // before waiting for the child process. Use this to wire interactive routing, + // prompts, or output buffering around an already-started PTY child. + PreparePTYSession func(*PTYRuntime) error + + IsInteractive func() bool +} + +type PTYRuntime struct { + Session pty.InteractiveSession + OutputRouter *pty.OutputRouter + InputRouter *pty.InputRouter + PromptReader *io.PipeReader + PromptWriter *io.PipeWriter +} + // Execute runs a package manager command without proxy or guard analysis. // It applies sandbox policy if configured, then executes the command directly. func Execute(ctx context.Context, pc *packagemanager.ParsedCommand, pmName string, dryRun bool) error { + return ExecuteWithOptions(ctx, pc, ExecuteOptions{ + PackageManagerName: pmName, + DryRun: dryRun, + Mode: ExecutionModeDirect, + }) +} + +// ExecuteWithOptions runs a package manager command through PMG's shared +// execution path: real binary resolution, environment setup, sandbox +// application, command launch, sandbox cleanup, and exit error wrapping. +func ExecuteWithOptions(ctx context.Context, pc *packagemanager.ParsedCommand, opts ExecuteOptions) error { if len(pc.Command.Exe) == 0 { return fmt.Errorf("no command to execute") } - if dryRun { + if opts.DryRun { log.Debugf("Dry run, skipping command execution") return nil } @@ -29,13 +83,21 @@ func Execute(ctx context.Context, pc *packagemanager.ParsedCommand, pmName strin return fmt.Errorf("failed to resolve real %s binary: %w", pc.Command.Exe, err) } + mode := executionMode(opts) + cmd := exec.CommandContext(ctx, realBinary, pc.Command.Args...) cmd.Stdin = os.Stdin cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr - cmd.Env = shim.FilterPMGFromEnv(os.Environ()) + cmd.Env = commandEnv(modeEnvOverrides(opts, mode)) - result, err := executor.ApplySandbox(ctx, cmd, pmName) + if mode != ExecutionModePTY && opts.BeforeDirectRun != nil { + if err := opts.BeforeDirectRun(); err != nil { + return err + } + } + + result, err := executor.ApplySandbox(ctx, cmd, opts.PackageManagerName) if err != nil { return fmt.Errorf("failed to apply sandbox: %w", err) } @@ -46,15 +108,193 @@ func Execute(ctx context.Context, pc *packagemanager.ParsedCommand, pmName strin } }() - if result.ShouldRun() { - if err := cmd.Run(); err != nil { - exitCode := -1 - if exitErr, ok := err.(*exec.ExitError); ok { - exitCode = exitErr.ExitCode() - } - return executor.WrapCommandExecutionError(err, result, exitCode) + switch mode { + case ExecutionModePTY: + return runPTY(ctx, cmd, cmd.Env, result, opts.PreparePTYSession) + default: + return runDirect(cmd, result) + } +} + +func runDirect(cmd *exec.Cmd, result *sandbox.ExecutionResult) error { + if !result.ShouldRun() { + return nil + } + + log.Debugf("Running command with args: %s: %v", cmd.Path, cmd.Args[1:]) + + if err := cmd.Run(); err != nil { + return wrapCommandExecutionError(err, result) + } + + log.Debugf("Command completed successfully") + return nil +} + +func runPTY( + ctx context.Context, + cmd *exec.Cmd, + env []string, + result *sandbox.ExecutionResult, + beforeWait func(*PTYRuntime) error, +) error { + if !result.ShouldRun() { + return usefulerror.Useful(). + Wrap(fmt.Errorf("sandbox not supported for PTY sessions")). + WithHumanError("Sandbox executed command cannot be used with PTY session. Please use non-interactive TTY mode instead.") + } + + cmdExe := cmd.Path + cmdArgs := cmd.Args[1:] + + log.Debugf("Running command with args: %s: %v", cmdExe, cmdArgs) + + sessionConfig := pty.NewSessionConfig(cmdExe, cmdArgs, env) + sess, err := pty.NewSession(ctx, sessionConfig) + if err != nil { + return fmt.Errorf("failed to create pty session: %w", err) + } + defer func() { + if err := sess.Close(); err != nil { + log.Warnf("failed to close pty session: %v", err) } + }() + + outputRouter, err := pty.NewOutputRouter(os.Stdout) + if err != nil { + return fmt.Errorf("failed to create output router: %w", err) + } + + var wg sync.WaitGroup + wg.Go(func() { + if _, err := io.Copy(outputRouter, sess.PtyReader()); err != nil { + log.Errorf("failed to copy output: %v", err) + } + }) + + inputRouter, err := pty.NewInputRouter(sess.PtyWriter()) + if err != nil { + return fmt.Errorf("failed to create input router: %w", err) + } + + promptReader, promptWriter := io.Pipe() + defer func() { + if err := promptWriter.Close(); err != nil { + log.Warnf("failed to close prompt writer: %v", err) + } + }() + defer func() { + if err := promptReader.Close(); err != nil { + log.Warnf("failed to close prompt reader: %v", err) + } + }() + + inputCtx, cancelInput := context.WithCancel(ctx) + inputDone := make(chan struct{}) + go func() { + defer close(inputDone) + inputRouter.ReadLoopContext(inputCtx, os.Stdin) + }() + defer func() { + cancelInput() + <-inputDone + }() + + if beforeWait != nil { + runtime := &PTYRuntime{ + Session: sess, + OutputRouter: outputRouter, + InputRouter: inputRouter, + PromptReader: promptReader, + PromptWriter: promptWriter, + } + + if err := beforeWait(runtime); err != nil { + return err + } + } + + sessionError := sess.Wait() + wg.Wait() + + if sessionError != nil { + return wrapCommandExecutionError(sessionError, result) } return nil } + +func executionMode(opts ExecuteOptions) ExecutionMode { + if opts.Mode != ExecutionModeAuto { + return opts.Mode + } + + isInteractive := pty.IsInteractiveTerminal + if opts.IsInteractive != nil { + isInteractive = opts.IsInteractive + } + + if isInteractive() { + return ExecutionModePTY + } + + return ExecutionModeDirect +} + +func commandEnv(overrides []string) []string { + return mergeEnv(shim.FilterPMGFromEnv(os.Environ()), overrides) +} + +func modeEnvOverrides(opts ExecuteOptions, mode ExecutionMode) []string { + overrides := append([]string{}, opts.EnvOverrides...) + switch mode { + case ExecutionModePTY: + overrides = append(overrides, opts.PTYEnvOverrides...) + default: + overrides = append(overrides, opts.DirectEnvOverrides...) + } + + return overrides +} + +func mergeEnv(base, overrides []string) []string { + env := append([]string{}, base...) + indexByKey := make(map[string]int, len(env)) + + for i, entry := range env { + key, _, ok := strings.Cut(entry, "=") + if ok { + indexByKey[key] = i + } + } + + for _, entry := range overrides { + key, _, ok := strings.Cut(entry, "=") + if !ok { + env = append(env, entry) + continue + } + + if idx, exists := indexByKey[key]; exists { + env[idx] = entry + continue + } + + indexByKey[key] = len(env) + env = append(env, entry) + } + + return env +} + +func wrapCommandExecutionError(err error, result *sandbox.ExecutionResult) error { + if exitErr, ok := err.(*exec.ExitError); ok { + return executor.WrapCommandExecutionError(err, result, exitErr.ExitCode()) + } + + if sessionError, ok := err.(*pty.ExitError); ok { + return executor.WrapCommandExecutionError(sessionError, result, sessionError.Code) + } + + return executor.WrapCommandExecutionError(err, result, -1) +} diff --git a/internal/runner/execute_test.go b/internal/runner/execute_test.go new file mode 100644 index 0000000..9c879ec --- /dev/null +++ b/internal/runner/execute_test.go @@ -0,0 +1,95 @@ +package runner + +import ( + "context" + "os" + "testing" + + "github.com/safedep/pmg/config" + "github.com/safedep/pmg/packagemanager" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestMergeEnvOverridesExistingValues(t *testing.T) { + env := mergeEnv( + []string{ + "PATH=/usr/bin", + "HTTP_PROXY=http://old-proxy", + "NO_PROXY=localhost", + }, + []string{ + "HTTP_PROXY=http://pmg-proxy", + "HTTPS_PROXY=http://pmg-proxy", + "NO_PROXY=localhost,127.0.0.1", + }, + ) + + assert.Equal(t, []string{ + "PATH=/usr/bin", + "HTTP_PROXY=http://pmg-proxy", + "NO_PROXY=localhost,127.0.0.1", + "HTTPS_PROXY=http://pmg-proxy", + }, env) +} + +func TestModeEnvOverrides(t *testing.T) { + opts := ExecuteOptions{ + EnvOverrides: []string{"HTTP_PROXY=http://pmg-proxy"}, + DirectEnvOverrides: []string{"CI=true"}, + PTYEnvOverrides: []string{"TERM=xterm-256color"}, + } + + assert.Equal(t, + []string{"HTTP_PROXY=http://pmg-proxy", "CI=true"}, + modeEnvOverrides(opts, ExecutionModeDirect), + ) + + assert.Equal(t, + []string{"HTTP_PROXY=http://pmg-proxy", "TERM=xterm-256color"}, + modeEnvOverrides(opts, ExecutionModePTY), + ) +} + +func TestExecutionModeAuto(t *testing.T) { + assert.Equal(t, ExecutionModePTY, executionMode(ExecuteOptions{ + Mode: ExecutionModeAuto, + IsInteractive: func() bool { return true }, + })) + + assert.Equal(t, ExecutionModeDirect, executionMode(ExecuteOptions{ + Mode: ExecutionModeAuto, + IsInteractive: func() bool { return false }, + })) +} + +func TestExecuteWithOptionsRunsDirectHookBeforeSandbox(t *testing.T) { + cfg := config.Get() + previous := *cfg + t.Cleanup(func() { + *cfg = previous + }) + + cfg.Config.Sandbox.Enabled = true + cfg.Config.Sandbox.Policies = map[string]config.SandboxPolicyRef{} + + exe, err := os.Executable() + require.NoError(t, err) + + hookCalled := false + err = ExecuteWithOptions(context.Background(), &packagemanager.ParsedCommand{ + Command: packagemanager.Command{ + Exe: exe, + }, + }, ExecuteOptions{ + PackageManagerName: "npm", + Mode: ExecutionModeDirect, + BeforeDirectRun: func() error { + hookCalled = true + return nil + }, + }) + + require.Error(t, err) + assert.True(t, hookCalled) +}