package compose import ( "context" "errors" "os" "os/exec" "path/filepath" "runtime" "strconv" "strings" "testing" "time" ) func TestRunnerPassesEachArgumentWithoutShellSplitting(t *testing.T) { t.Parallel() runner := NewRunner(writeExecutable(t, "#!/bin/sh\nprintf '<%s>\\n' \"$@\"\ncat\n")) result, err := runner.Run(context.Background(), []string{"compose", "--project-directory", "/tmp/a project with spaces", "config"}, strings.NewReader("stdin value\n")) if err != nil { t.Fatalf("Run() error = %v", err) } want := "\n<--project-directory>\n\n\nstdin value\n" if result.Stdout != want { t.Errorf("stdout = %q, want %q", result.Stdout, want) } if result.ExitCode != 0 { t.Errorf("ExitCode = %d, want 0", result.ExitCode) } } func TestRunnerReturnsTheChildExitCode(t *testing.T) { t.Parallel() runner := NewRunner(writeExecutable(t, "#!/bin/sh\necho unavailable >&2\nexit 42\n")) result, err := runner.Run(context.Background(), []string{"compose", "ps"}, nil) if err == nil { t.Fatal("Run() error = nil, want child exit error") } if result.ExitCode != 42 { t.Errorf("ExitCode = %d, want 42", result.ExitCode) } if result.Stderr != "unavailable\n" { t.Errorf("stderr = %q, want unavailable output", result.Stderr) } } func TestRunnerReportsMissingDocker(t *testing.T) { t.Parallel() runner := NewRunner(filepath.Join(t.TempDir(), "docker-does-not-exist")) result, err := runner.Run(context.Background(), []string{"compose", "version"}, nil) if !errors.Is(err, exec.ErrNotFound) { t.Fatalf("Run() error = %v, want exec.ErrNotFound", err) } if result.ExitCode != 127 { t.Errorf("ExitCode = %d, want 127", result.ExitCode) } } func writeExecutable(t *testing.T, contents string) string { t.Helper() path := filepath.Join(t.TempDir(), "fake-docker") if err := os.WriteFile(path, []byte(contents), 0o700); err != nil { t.Fatal(err) } return path } func TestRunnerBoundedCapturesExactlyAtLimit(t *testing.T) { runner := NewRunner(writeExecutable(t, "#!/bin/sh\nprintf 1234567890\nprintf error12345 >&2\n")) result, err := runner.RunBounded(context.Background(), nil, nil, CaptureLimits{StdoutBytes: 10, StderrBytes: 10}) if err != nil { t.Fatal(err) } if result.Stdout != "1234567890" || result.Stderr != "error12345" { t.Fatalf("bounded result = %#v", result) } } func TestRunnerBoundedRejectsOneByteOverEachStream(t *testing.T) { for _, script := range []string{"#!/bin/sh\nprintf 12345678901\n", "#!/bin/sh\nprintf 12345678901 >&2\n"} { runner := NewRunner(writeExecutable(t, script)) result, err := runner.RunBounded(context.Background(), nil, nil, CaptureLimits{StdoutBytes: 10, StderrBytes: 10}) if !errors.Is(err, ErrOutputLimit) { t.Fatalf("err = %v result=%#v", err, result) } if len(result.Stdout) > 10 || len(result.Stderr) > 10 { t.Fatalf("retained overflow: %#v", result) } } } func TestRunnerBoundedTerminatesInfiniteStreams(t *testing.T) { for _, script := range []string{"#!/bin/sh\nwhile :; do printf x; done", "#!/bin/sh\nwhile :; do printf x >&2; done"} { runner := NewRunner(writeExecutable(t, script)) result, err := runner.RunBounded(context.Background(), nil, nil, CaptureLimits{StdoutBytes: 1024, StderrBytes: 1024}) if !errors.Is(err, ErrOutputLimit) { t.Fatalf("err = %v result = %#v", err, result) } if len(result.Stdout) > 1024 || len(result.Stderr) > 1024 { t.Fatalf("retained overflow: %#v", result) } } } func TestRunnerBoundedCancellationWinsOverChildKill(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) runner := NewRunner(writeExecutable(t, "#!/bin/sh\nsleep 30\n")) done := make(chan error, 1) go func() { _, err := runner.RunBounded(ctx, nil, nil, CaptureLimits{StdoutBytes: 32, StderrBytes: 32}) done <- err }() cancel() select { case err := <-done: if !errors.Is(err, context.Canceled) { t.Fatalf("err = %v", err) } case <-time.After(3 * time.Second): t.Fatal("cancellation did not terminate process") } } func TestRunnerBoundedKillsGrandchildProcess(t *testing.T) { if runtime.GOOS == "windows" { t.Skip("shell process-group fixture is Unix-specific") } pidFile := filepath.Join(t.TempDir(), "grandchild.pid") script := "#!/bin/sh\n(sleep 30) &\nprintf '%s' $! > '" + pidFile + "'\nwhile :; do printf x; done\n" runner := NewRunner(writeExecutable(t, script)) _, err := runner.RunBounded(context.Background(), nil, nil, CaptureLimits{StdoutBytes: 64, StderrBytes: 64}) if !errors.Is(err, ErrOutputLimit) { t.Fatalf("err = %v", err) } pidBytes, readErr := os.ReadFile(pidFile) if readErr != nil { t.Fatal(readErr) } pid, err := strconv.Atoi(strings.TrimSpace(string(pidBytes))) if err != nil { t.Fatal(err) } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { if !processExists(pid) { return } time.Sleep(20 * time.Millisecond) } t.Fatalf("grandchild pid %d remained alive after owned-process termination", pid) } func TestRunnerBoundedTerminatesGrandchildHoldingPipe(t *testing.T) { if runtime.GOOS == "windows" { t.Skip("shell process-group fixture is Unix-specific") } runner := NewRunner(writeExecutable(t, "#!/bin/sh\n(sleep 30; printf survivor) &\nwhile :; do printf x; done\n")) done := make(chan error, 1) go func() { _, err := runner.RunBounded(context.Background(), nil, nil, CaptureLimits{StdoutBytes: 64, StderrBytes: 64}) done <- err }() select { case err := <-done: if !errors.Is(err, ErrOutputLimit) { t.Fatalf("err = %v", err) } case <-time.After(3 * time.Second): t.Fatal("grandchild kept owned pipe alive") } }