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