diff --git a/internal/configvalue/configvalue.go b/internal/configvalue/configvalue.go index 4b0c2b4edc0f016488518de6d9fb6e84db5edd15..17d87cdf40a0b237ad7f25245817a9b5fed7dce3 100644 --- a/internal/configvalue/configvalue.go +++ b/internal/configvalue/configvalue.go @@ -18,6 +18,8 @@ import ( "time" ) +const maxCommandStdoutBytes = 64 * 1024 + var ( envNameRE = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`) envNamePrefixRE = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*`) @@ -44,14 +46,16 @@ type result struct { type Resolver struct { CommandTimeout time.Duration - mu sync.Mutex - cache map[string]result + shellCommand func(string) (string, []string) + mu sync.Mutex + cache map[string]result } // NewResolver returns a resolver with the default command timeout. func NewResolver() *Resolver { return &Resolver{ CommandTimeout: 10 * time.Second, + shellCommand: platformShellCommand, cache: make(map[string]result), } } @@ -130,12 +134,12 @@ func (r *Resolver) executeCommand(ctx context.Context, command string) (string, ctx, cancel := context.WithTimeout(ctx, timeout) defer cancel() - name, args := shellCommand(command) + name, args := r.shellCommand(command) cmd := exec.CommandContext(ctx, name, args...) cmd.Stdin = nil cmd.Stderr = io.Discard - output, err := cmd.Output() + output, err := commandOutput(cmd) if err != nil { return "", false } @@ -148,7 +152,38 @@ func (r *Resolver) executeCommand(ctx context.Context, command string) (string, return value, true } -func shellCommand(command string) (string, []string) { +func commandOutput(cmd *exec.Cmd) ([]byte, error) { + stdout, err := cmd.StdoutPipe() + if err != nil { + return nil, fmt.Errorf("capture command stdout: %w", err) + } + + if err := cmd.Start(); err != nil { + return nil, fmt.Errorf("start command: %w", err) + } + + output, err := io.ReadAll(io.LimitReader(stdout, maxCommandStdoutBytes+1)) + oversized := len(output) > maxCommandStdoutBytes + if err != nil || oversized { + if cmd.Process != nil { + _ = cmd.Process.Kill() + } + _ = cmd.Wait() + if oversized { + return nil, fmt.Errorf("command stdout exceeds limit") + } + + return nil, fmt.Errorf("read command stdout: %w", err) + } + + if err := cmd.Wait(); err != nil { + return nil, fmt.Errorf("wait for command: %w", err) + } + + return output, nil +} + +func platformShellCommand(command string) (string, []string) { if runtime.GOOS == "windows" { return "cmd.exe", []string{"/d", "/s", "/c", command} } diff --git a/internal/configvalue/configvalue_test.go b/internal/configvalue/configvalue_test.go index 422c9ad4af841cee6af19bddd8084e2ec758ddb9..4f58c10fcde37caf08c2a035ba7c9c44a3cd2ed1 100644 --- a/internal/configvalue/configvalue_test.go +++ b/internal/configvalue/configvalue_test.go @@ -6,7 +6,10 @@ package configvalue import ( "context" + "os" "runtime" + "strconv" + "strings" "testing" "time" ) @@ -78,3 +81,37 @@ func TestResolveACIDConfigValuesCommands3CachesCommandResults(t *testing.T) { t.Fatalf("cached command value = %q, %v; want cooked, true", third, ok) } } + +func TestResolveACIDAuthenticationCredentials17RejectsOversizedCommandStdout(t *testing.T) { + t.Setenv("COOKED_CONFIGVALUE_FAKE_SHELL", "1") + t.Setenv("COOKED_CONFIGVALUE_FAKE_SHELL_STDOUT_BYTES", strconv.Itoa(maxCommandStdoutBytes+1)) + + resolver := NewResolver() + resolver.shellCommand = func(string) (string, []string) { + return os.Args[0], []string{"-test.run=TestConfigValueFakeShell"} + } + + got, ok := resolver.Resolve(context.Background(), "!oversized") + if ok { + t.Fatalf("oversized command stdout resolved with length %d, want unresolved", len(got)) + } + if got != "" { + t.Fatalf("oversized command stdout returned non-empty value length %d, want empty", len(got)) + } +} + +func TestConfigValueFakeShell(t *testing.T) { + if os.Getenv("COOKED_CONFIGVALUE_FAKE_SHELL") != "1" { + return + } + + byteCount, err := strconv.Atoi(os.Getenv("COOKED_CONFIGVALUE_FAKE_SHELL_STDOUT_BYTES")) + if err != nil { + os.Exit(2) + } + if _, err := os.Stdout.WriteString(strings.Repeat("x", byteCount)); err != nil { + os.Exit(2) + } + + os.Exit(0) +}