configvalue: bound command stdout

Amolith created

Change summary

internal/configvalue/configvalue.go      | 45 +++++++++++++++++++++++--
internal/configvalue/configvalue_test.go | 37 +++++++++++++++++++++
2 files changed, 77 insertions(+), 5 deletions(-)

Detailed changes

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}
 	}

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)
+}