fix: close workspace preprocessing contract gaps
This commit is contained in:
@@ -3,32 +3,57 @@
|
||||
package compose
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"sync"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
var ownedJobs = struct {
|
||||
sync.Mutex
|
||||
m map[*exec.Cmd]windows.Handle
|
||||
}{m: make(map[*exec.Cmd]windows.Handle)}
|
||||
// PROC_THREAD_ATTRIBUTE_JOB_LIST is intentionally kept here because x/sys does
|
||||
// not expose this newer SDK constant yet. Supplying it to CreateProcess makes
|
||||
// job ownership atomic with process creation; CREATE_SUSPENDED plus a later
|
||||
// AssignProcessToJobObject is not sufficient (the parent can die in between).
|
||||
const procThreadAttributeJobList = 0x0002000d
|
||||
|
||||
// CREATE_SUSPENDED closes the registration race: no child code can create a
|
||||
// descendant until the real process HANDLE has been assigned to the job.
|
||||
func configureOwnedProcess(c *exec.Cmd) {
|
||||
c.SysProcAttr = &windows.SysProcAttr{CreationFlags: windows.CREATE_NEW_PROCESS_GROUP | windows.CREATE_SUSPENDED}
|
||||
// The generic exec.Cmd path is not used on Windows: runBoundedPlatform launches
|
||||
// with CreateProcess and a creation-time job-list attribute. These hooks remain
|
||||
// for the platform-neutral runner's compile-time shape only.
|
||||
func configureOwnedProcess(*exec.Cmd) {}
|
||||
func registerOwnedProcess(*exec.Cmd) error {
|
||||
return errors.New("Windows owned processes require creation-time job ownership")
|
||||
}
|
||||
func releaseOwnedProcess(*exec.Cmd) {}
|
||||
func terminateOwnedProcess(c *exec.Cmd) {
|
||||
if c.Process != nil {
|
||||
_ = c.Process.Kill()
|
||||
}
|
||||
}
|
||||
|
||||
func registerOwnedProcess(c *exec.Cmd) error {
|
||||
if c.Process == nil {
|
||||
return errors.New("owned process has no process handle")
|
||||
func runBoundedPlatform(ctx context.Context, runner Runner, args []string, stdin io.Reader, limits CaptureLimits) (Result, bool, error) {
|
||||
if limits.StdoutBytes < 0 || limits.StderrBytes < 0 {
|
||||
return Result{}, true, ErrOutputLimit
|
||||
}
|
||||
result, err := runBoundedWindows(ctx, runner.binary, args, stdin, limits)
|
||||
return result, true, err
|
||||
}
|
||||
|
||||
type windowsCollector struct {
|
||||
buf []byte
|
||||
n int64
|
||||
overflow bool
|
||||
}
|
||||
|
||||
func runBoundedWindows(ctx context.Context, binary string, args []string, stdin io.Reader, limits CaptureLimits) (Result, error) {
|
||||
job, err := windows.CreateJobObject(nil, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
return Result{}, err
|
||||
}
|
||||
closeJob := true
|
||||
defer func() {
|
||||
@@ -36,72 +61,178 @@ func registerOwnedProcess(c *exec.Cmd) error {
|
||||
windows.CloseHandle(job)
|
||||
}
|
||||
}()
|
||||
limits := windows.JOBOBJECT_EXTENDED_LIMIT_INFORMATION{}
|
||||
limits.BasicLimitInformation.LimitFlags = windows.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE
|
||||
if _, err = windows.SetInformationJobObject(job, windows.JobObjectExtendedLimitInformation, uintptr(unsafe.Pointer(&limits)), uint32(unsafe.Sizeof(limits))); err != nil {
|
||||
return err
|
||||
limitsInfo := windows.JOBOBJECT_EXTENDED_LIMIT_INFORMATION{}
|
||||
limitsInfo.BasicLimitInformation.LimitFlags = windows.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE
|
||||
if _, err = windows.SetInformationJobObject(job, windows.JobObjectExtendedLimitInformation, uintptr(unsafe.Pointer(&limitsInfo)), uint32(unsafe.Sizeof(limitsInfo))); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
// c.Process.Pid is used only to obtain a genuine process HANDLE; a PID is
|
||||
// never passed to AssignProcessToJobObject.
|
||||
process, err := windows.OpenProcess(windows.PROCESS_SET_QUOTA|windows.PROCESS_TERMINATE|windows.PROCESS_QUERY_LIMITED_INFORMATION, false, uint32(c.Process.Pid))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer windows.CloseHandle(process)
|
||||
if err = windows.AssignProcessToJobObject(job, process); err != nil {
|
||||
return err
|
||||
}
|
||||
thread, err := suspendedPrimaryThread(uint32(c.Process.Pid))
|
||||
if err != nil {
|
||||
_ = windows.TerminateJobObject(job, 1)
|
||||
return err
|
||||
}
|
||||
_, resumeErr := windows.ResumeThread(thread)
|
||||
windows.CloseHandle(thread)
|
||||
if resumeErr != nil {
|
||||
_ = windows.TerminateJobObject(job, 1)
|
||||
return resumeErr
|
||||
}
|
||||
ownedJobs.Lock()
|
||||
ownedJobs.m[c] = job
|
||||
ownedJobs.Unlock()
|
||||
closeJob = false
|
||||
return nil
|
||||
}
|
||||
|
||||
func suspendedPrimaryThread(pid uint32) (windows.Handle, error) {
|
||||
snapshot, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPTHREAD, 0)
|
||||
stdinRead, stdinWrite, err := os.Pipe()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return Result{}, err
|
||||
}
|
||||
defer windows.CloseHandle(snapshot)
|
||||
entry := windows.ThreadEntry32{Size: uint32(unsafe.Sizeof(windows.ThreadEntry32{}))}
|
||||
err = windows.Thread32First(snapshot, &entry)
|
||||
for err == nil {
|
||||
if entry.OwnerProcessID == pid {
|
||||
return windows.OpenThread(windows.THREAD_SUSPEND_RESUME, false, entry.ThreadID)
|
||||
stdoutRead, stdoutWrite, err := os.Pipe()
|
||||
if err != nil {
|
||||
_ = stdinRead.Close()
|
||||
_ = stdinWrite.Close()
|
||||
return Result{}, err
|
||||
}
|
||||
stderrRead, stderrWrite, err := os.Pipe()
|
||||
if err != nil {
|
||||
_ = stdinRead.Close()
|
||||
_ = stdinWrite.Close()
|
||||
_ = stdoutRead.Close()
|
||||
_ = stdoutWrite.Close()
|
||||
return Result{}, err
|
||||
}
|
||||
closeAll := func() {
|
||||
_ = stdinRead.Close()
|
||||
_ = stdinWrite.Close()
|
||||
_ = stdoutRead.Close()
|
||||
_ = stdoutWrite.Close()
|
||||
_ = stderrRead.Close()
|
||||
_ = stderrWrite.Close()
|
||||
}
|
||||
defer closeAll()
|
||||
childHandles := []windows.Handle{windows.Handle(stdinRead.Fd()), windows.Handle(stdoutWrite.Fd()), windows.Handle(stderrWrite.Fd())}
|
||||
for _, h := range childHandles {
|
||||
if err := windows.SetHandleInformation(h, windows.HANDLE_FLAG_INHERIT, windows.HANDLE_FLAG_INHERIT); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
err = windows.Thread32Next(snapshot, &entry)
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
// The attribute list limits inheritance to exactly these standard handles.
|
||||
attributes, err := windows.NewProcThreadAttributeList(2)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
defer attributes.Delete()
|
||||
jobList := []windows.Handle{job}
|
||||
if err := attributes.Update(procThreadAttributeJobList, unsafe.Pointer(&jobList[0]), unsafe.Sizeof(jobList[0])); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
if err := attributes.Update(windows.PROC_THREAD_ATTRIBUTE_HANDLE_LIST, unsafe.Pointer(&childHandles[0]), uintptr(len(childHandles))*unsafe.Sizeof(childHandles[0])); err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
|
||||
func releaseOwnedProcess(c *exec.Cmd) {
|
||||
ownedJobs.Lock()
|
||||
h := ownedJobs.m[c]
|
||||
delete(ownedJobs.m, c)
|
||||
ownedJobs.Unlock()
|
||||
if h != 0 {
|
||||
windows.CloseHandle(h)
|
||||
command := syscall.EscapeArg(binary)
|
||||
for _, arg := range args {
|
||||
command += " " + syscall.EscapeArg(arg)
|
||||
}
|
||||
}
|
||||
func terminateOwnedProcess(c *exec.Cmd) {
|
||||
ownedJobs.Lock()
|
||||
h := ownedJobs.m[c]
|
||||
ownedJobs.Unlock()
|
||||
if h != 0 {
|
||||
_ = windows.TerminateJobObject(h, 1)
|
||||
} else if c.Process != nil {
|
||||
_ = c.Process.Kill()
|
||||
commandLine, err := windows.UTF16FromString(command)
|
||||
if err != nil {
|
||||
return Result{}, err
|
||||
}
|
||||
startup := windows.StartupInfoEx{
|
||||
StartupInfo: windows.StartupInfo{
|
||||
Cb: uint32(unsafe.Sizeof(windows.StartupInfoEx{})),
|
||||
Flags: windows.STARTF_USESTDHANDLES,
|
||||
StdInput: childHandles[0],
|
||||
StdOutput: childHandles[1],
|
||||
StdErr: childHandles[2],
|
||||
},
|
||||
ProcThreadAttributeList: attributes.List(),
|
||||
}
|
||||
var process windows.ProcessInformation
|
||||
creationFlags := uint32(windows.CREATE_UNICODE_ENVIRONMENT | windows.CREATE_NEW_PROCESS_GROUP | windows.EXTENDED_STARTUPINFO_PRESENT)
|
||||
if err := windows.CreateProcess(nil, &commandLine[0], nil, nil, true, creationFlags, nil, nil, (*windows.StartupInfo)(unsafe.Pointer(&startup)), &process); err != nil {
|
||||
return Result{ExitCode: 127}, err
|
||||
}
|
||||
// The child owns duplicate pipe handles; parent copies use the opposite ends.
|
||||
_ = stdinRead.Close()
|
||||
_ = stdoutWrite.Close()
|
||||
_ = stderrWrite.Close()
|
||||
_ = windows.CloseHandle(process.Thread)
|
||||
|
||||
var inputDone sync.WaitGroup
|
||||
inputDone.Add(1)
|
||||
go func() {
|
||||
defer inputDone.Done()
|
||||
defer stdinWrite.Close()
|
||||
if stdin != nil {
|
||||
_, _ = io.Copy(stdinWrite, stdin)
|
||||
}
|
||||
}()
|
||||
out := windowsCollector{buf: make([]byte, limits.StdoutBytes)}
|
||||
er := windowsCollector{buf: make([]byte, limits.StderrBytes)}
|
||||
overflow := make(chan struct{}, 1)
|
||||
var readers sync.WaitGroup
|
||||
readers.Add(2)
|
||||
read := func(file *os.File, maximum int64, collector *windowsCollector) {
|
||||
defer readers.Done()
|
||||
chunk := make([]byte, 32*1024)
|
||||
for {
|
||||
n, readErr := file.Read(chunk)
|
||||
if n > 0 {
|
||||
remain := maximum - collector.n
|
||||
if remain > 0 {
|
||||
take := int64(n)
|
||||
if take > remain {
|
||||
take = remain
|
||||
}
|
||||
copy(collector.buf[collector.n:collector.n+take], chunk[:take])
|
||||
collector.n += take
|
||||
}
|
||||
if int64(n) > remain {
|
||||
collector.overflow = true
|
||||
select {
|
||||
case overflow <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
if readErr != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
go read(stdoutRead, limits.StdoutBytes, &out)
|
||||
go read(stderrRead, limits.StderrBytes, &er)
|
||||
waited := make(chan uint32, 1)
|
||||
go func() {
|
||||
_, waitErr := windows.WaitForSingleObject(process.Process, windows.INFINITE)
|
||||
if waitErr != nil {
|
||||
waited <- uint32(0x103)
|
||||
} else {
|
||||
var code uint32
|
||||
if windows.GetExitCodeProcess(process.Process, &code) != nil {
|
||||
code = 1
|
||||
}
|
||||
waited <- code
|
||||
}
|
||||
}()
|
||||
var exitCode uint32
|
||||
var cancelled error
|
||||
select {
|
||||
case exitCode = <-waited:
|
||||
// Closing the job handle kills any descendants before waiting on pipe readers.
|
||||
_ = windows.CloseHandle(job)
|
||||
closeJob = false
|
||||
case <-overflow:
|
||||
_ = windows.TerminateJobObject(job, 1)
|
||||
exitCode = <-waited
|
||||
_ = windows.CloseHandle(job)
|
||||
closeJob = false
|
||||
case <-ctx.Done():
|
||||
cancelled = ctx.Err()
|
||||
_ = windows.TerminateJobObject(job, 1)
|
||||
exitCode = <-waited
|
||||
_ = windows.CloseHandle(job)
|
||||
closeJob = false
|
||||
}
|
||||
_ = windows.CloseHandle(process.Process)
|
||||
_ = stdinWrite.Close()
|
||||
readers.Wait()
|
||||
inputDone.Wait()
|
||||
result := Result{Stdout: string(out.buf[:out.n]), Stderr: string(er.buf[:er.n]), ExitCode: int(exitCode)}
|
||||
if out.overflow || er.overflow {
|
||||
return result, ErrOutputLimit
|
||||
}
|
||||
if cancelled != nil {
|
||||
return result, cancelled
|
||||
}
|
||||
if exitCode == 0 {
|
||||
return result, nil
|
||||
}
|
||||
return result, fmt.Errorf("child exited with code %d", exitCode)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user