fix(auth): harden diagnostic cleanup races
This commit is contained in:
@@ -0,0 +1,41 @@
|
||||
package compose
|
||||
|
||||
import "time"
|
||||
|
||||
// terminateWindowsProcess reports a reap failure only when the process has not exited within the
|
||||
// final bound. taskkill and Process.Kill can both race with a child that has already been reaped.
|
||||
func terminateWindowsProcess(
|
||||
done <-chan error,
|
||||
terminateTree func() error,
|
||||
terminateDirect func() error,
|
||||
bound time.Duration,
|
||||
) error {
|
||||
if processReaped(done) {
|
||||
return nil
|
||||
}
|
||||
_ = terminateTree()
|
||||
if processReaped(done) {
|
||||
return nil
|
||||
}
|
||||
_ = terminateDirect()
|
||||
if processReaped(done) {
|
||||
return nil
|
||||
}
|
||||
timer := time.NewTimer(bound)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-done:
|
||||
return nil
|
||||
case <-timer.C:
|
||||
return ErrProcessReap
|
||||
}
|
||||
}
|
||||
|
||||
func processReaped(done <-chan error) bool {
|
||||
select {
|
||||
case <-done:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
package compose
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestWindowsCancellationTreatsProcessReapedBeforeTreeTerminationAsSuccess(t *testing.T) {
|
||||
done := make(chan error, 1)
|
||||
done <- nil
|
||||
treeCalls := 0
|
||||
directCalls := 0
|
||||
|
||||
err := terminateWindowsProcess(done, func() error {
|
||||
treeCalls++
|
||||
return errors.New("tree termination should not run")
|
||||
}, func() error {
|
||||
directCalls++
|
||||
return errors.New("direct termination should not run")
|
||||
}, 0)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("terminateWindowsProcess() error = %v, want nil", err)
|
||||
}
|
||||
if treeCalls != 0 || directCalls != 0 {
|
||||
t.Fatalf("termination calls = tree:%d direct:%d, want none", treeCalls, directCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWindowsCancellationTreatsProcessReapedAfterTreeTerminationAsSuccess(t *testing.T) {
|
||||
done := make(chan error, 1)
|
||||
directCalls := 0
|
||||
|
||||
err := terminateWindowsProcess(done, func() error {
|
||||
done <- nil
|
||||
return ErrProcessReap
|
||||
}, func() error {
|
||||
directCalls++
|
||||
return errors.New("direct termination should not run")
|
||||
}, 0)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("terminateWindowsProcess() error = %v, want nil", err)
|
||||
}
|
||||
if directCalls != 0 {
|
||||
t.Fatalf("direct termination calls = %d, want 0", directCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWindowsCancellationTreatsProcessReapedAfterDirectTerminationAsSuccess(t *testing.T) {
|
||||
done := make(chan error, 1)
|
||||
directCalls := 0
|
||||
|
||||
err := terminateWindowsProcess(done, func() error {
|
||||
return ErrProcessReap
|
||||
}, func() error {
|
||||
directCalls++
|
||||
done <- nil
|
||||
return ErrProcessReap
|
||||
}, 0)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("terminateWindowsProcess() error = %v, want nil", err)
|
||||
}
|
||||
if directCalls != 1 {
|
||||
t.Fatalf("direct termination calls = %d, want 1", directCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWindowsCancellationReportsReapFailureWhenProcessSurvives(t *testing.T) {
|
||||
done := make(chan error)
|
||||
treeCalls := 0
|
||||
directCalls := 0
|
||||
|
||||
err := terminateWindowsProcess(done, func() error {
|
||||
treeCalls++
|
||||
return ErrProcessReap
|
||||
}, func() error {
|
||||
directCalls++
|
||||
return ErrProcessReap
|
||||
}, 0)
|
||||
|
||||
if !errors.Is(err, ErrProcessReap) {
|
||||
t.Fatalf("terminateWindowsProcess() error = %v, want ErrProcessReap", err)
|
||||
}
|
||||
if treeCalls != 1 || directCalls != 1 {
|
||||
t.Fatalf("termination calls = tree:%d direct:%d, want one each", treeCalls, directCalls)
|
||||
}
|
||||
}
|
||||
@@ -6,15 +6,20 @@ import (
|
||||
"context"
|
||||
"io"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
const finalTerminationBound = 2 * time.Second
|
||||
const processTreeTerminationBound = 2 * time.Second
|
||||
const createNewProcessGroup = 0x00000200
|
||||
|
||||
var windowsSystemDirectory = windows.GetSystemDirectory
|
||||
|
||||
func configureProcess(command *exec.Cmd) {
|
||||
command.SysProcAttr = &syscall.SysProcAttr{CreationFlags: createNewProcessGroup}
|
||||
}
|
||||
@@ -23,23 +28,22 @@ func terminateProcess(command *exec.Cmd, done <-chan error) error {
|
||||
if command.Process == nil {
|
||||
return nil
|
||||
}
|
||||
treeErr := terminateWindowsProcessTree(command.Process.Pid)
|
||||
_ = command.Process.Kill()
|
||||
select {
|
||||
case <-done:
|
||||
if treeErr != nil {
|
||||
return ErrProcessReap
|
||||
}
|
||||
return nil
|
||||
case <-time.After(finalTerminationBound):
|
||||
return ErrProcessReap
|
||||
}
|
||||
return terminateWindowsProcess(
|
||||
done,
|
||||
func() error { return terminateWindowsProcessTree(command.Process.Pid) },
|
||||
command.Process.Kill,
|
||||
finalTerminationBound,
|
||||
)
|
||||
}
|
||||
|
||||
func terminateWindowsProcessTree(pid int) error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), processTreeTerminationBound)
|
||||
defer cancel()
|
||||
command := windowsTreeKillCommand(ctx, pid)
|
||||
taskkillPath, err := systemTaskkillPath()
|
||||
if err != nil {
|
||||
return ErrProcessReap
|
||||
}
|
||||
command := windowsTreeKillCommand(ctx, taskkillPath, pid)
|
||||
command.Stdout = io.Discard
|
||||
command.Stderr = io.Discard
|
||||
if err := command.Run(); err != nil {
|
||||
@@ -51,6 +55,23 @@ func terminateWindowsProcessTree(pid int) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func windowsTreeKillCommand(ctx context.Context, pid int) *exec.Cmd {
|
||||
return exec.CommandContext(ctx, "taskkill.exe", "/PID", strconv.Itoa(pid), "/T", "/F")
|
||||
// systemTaskkillPath resolves taskkill through Windows' protected system-directory API, never PATH.
|
||||
func systemTaskkillPath() (string, error) {
|
||||
systemDirectory, err := windowsSystemDirectory()
|
||||
if err != nil {
|
||||
return "", ErrProcessReap
|
||||
}
|
||||
systemDirectory = filepath.Clean(systemDirectory)
|
||||
if systemDirectory == "." || !filepath.IsAbs(systemDirectory) {
|
||||
return "", ErrProcessReap
|
||||
}
|
||||
taskkillPath := filepath.Join(systemDirectory, "taskkill.exe")
|
||||
if !filepath.IsAbs(taskkillPath) || filepath.Dir(taskkillPath) != systemDirectory {
|
||||
return "", ErrProcessReap
|
||||
}
|
||||
return taskkillPath, nil
|
||||
}
|
||||
|
||||
func windowsTreeKillCommand(ctx context.Context, taskkillPath string, pid int) *exec.Cmd {
|
||||
return exec.CommandContext(ctx, taskkillPath, "/PID", strconv.Itoa(pid), "/T", "/F")
|
||||
}
|
||||
|
||||
@@ -13,14 +13,19 @@ func TestWindowsTerminationUsesBoundedExactPIDTreeKillWithoutAShell(t *testing.T
|
||||
}
|
||||
text := string(source)
|
||||
for _, required := range []string{
|
||||
"context.WithTimeout", "exec.CommandContext", `"taskkill.exe"`, `"/PID"`,
|
||||
"strconv.Itoa", `"/T"`, `"/F"`, "io.Discard", "ErrProcessReap",
|
||||
"context.WithTimeout", "windows.GetSystemDirectory", "filepath.Clean", "filepath.IsAbs",
|
||||
`filepath.Join(systemDirectory, "taskkill.exe")`, "systemTaskkillPath",
|
||||
`exec.CommandContext(ctx, taskkillPath, "/PID", strconv.Itoa(pid), "/T", "/F")`,
|
||||
"io.Discard", "ErrProcessReap",
|
||||
} {
|
||||
if !strings.Contains(text, required) {
|
||||
t.Errorf("process_windows.go is missing %q", required)
|
||||
}
|
||||
}
|
||||
for _, forbidden := range []string{"cmd.exe", "powershell", `exec.Command("taskkill.exe"`} {
|
||||
for _, forbidden := range []string{
|
||||
"cmd.exe", "powershell", `exec.Command("taskkill.exe"`,
|
||||
`exec.CommandContext(ctx, "taskkill.exe"`, `exec.LookPath("taskkill.exe")`,
|
||||
} {
|
||||
if strings.Contains(strings.ToLower(text), strings.ToLower(forbidden)) {
|
||||
t.Errorf("process_windows.go contains unsafe command form %q", forbidden)
|
||||
}
|
||||
|
||||
@@ -4,14 +4,37 @@ package compose
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestWindowsTreeKillCommandUsesExactPIDArgumentArray(t *testing.T) {
|
||||
command := windowsTreeKillCommand(context.Background(), 4242)
|
||||
want := []string{"taskkill.exe", "/PID", "4242", "/T", "/F"}
|
||||
func TestWindowsTreeKillCommandUsesTrustedAbsoluteSystemPathAndExactPIDArgumentArray(t *testing.T) {
|
||||
original := windowsSystemDirectory
|
||||
windowsSystemDirectory = func() (string, error) { return `C:\\Windows\\System32`, nil }
|
||||
t.Cleanup(func() { windowsSystemDirectory = original })
|
||||
|
||||
taskkillPath, err := systemTaskkillPath()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !filepath.IsAbs(taskkillPath) {
|
||||
t.Fatalf("taskkill path = %q, want absolute trusted system path", taskkillPath)
|
||||
}
|
||||
command := windowsTreeKillCommand(context.Background(), taskkillPath, 4242)
|
||||
want := []string{taskkillPath, "/PID", "4242", "/T", "/F"}
|
||||
if !reflect.DeepEqual(command.Args, want) {
|
||||
t.Fatalf("taskkill args = %#v, want %#v", command.Args, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWindowsSystemTaskkillPathRejectsRelativeSystemDirectory(t *testing.T) {
|
||||
original := windowsSystemDirectory
|
||||
windowsSystemDirectory = func() (string, error) { return `System32`, nil }
|
||||
t.Cleanup(func() { windowsSystemDirectory = original })
|
||||
|
||||
if _, err := systemTaskkillPath(); !errors.Is(err, ErrProcessReap) {
|
||||
t.Fatalf("systemTaskkillPath() error = %v, want ErrProcessReap", err)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user