fix: close workspace preprocessing contract gaps

This commit is contained in:
2026-08-11 01:18:29 +02:00
parent 05a6e8cc2d
commit de27275bb1
15 changed files with 516 additions and 87 deletions
+58
View File
@@ -105,6 +105,59 @@ func TestUsageDocumentsClosedConfigureUpdateSourcesAndMaintenanceRecovery(t *tes
}
}
func TestRunWorkspacePublicDispatchExitMatrix(t *testing.T) {
cases := []struct {
name, status, code string
exit int
}{
{"succeeded", "succeeded", "ok", 0},
{"blocked", "blocked", "manual_review_required", 3},
{"failed", "failed", "workspace_not_found", 1},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
fixture := newCLIFixture(t, "")
fixture.setEnvironment(t)
t.Setenv("THOTHCTL_FAKE_WORKSPACE_RESULT", fmt.Sprintf(`{"schemaVersion":1,"status":%q,"code":%q,"workspaceId":"psd","workspaceRevision":"%s","descriptorBlob":"%s","operation":"inspect","completedStages":[]}`, tc.status, tc.code, strings.Repeat("0", 40), strings.Repeat("a", 40)))
t.Setenv("THOTHCTL_FAKE_WORKSPACE_EXIT", strconv.Itoa(tc.exit))
var stdout, stderr bytes.Buffer
got := run(context.Background(), []string{"--installation", fixture.installationPath, "workspace", "inspect", "--workspace", "psd", "--json"}, &stdout, &stderr)
if got != tc.exit || stderr.Len() != 0 {
t.Fatalf("exit=%d stderr=%q; want exit %d", got, stderr.String(), tc.exit)
}
if len(fixture.invocations(t)) != 1 {
t.Fatalf("workspace dispatch invocations = %#v, want one", fixture.invocations(t))
}
})
}
}
func TestRunWorkspaceOperationalFailureExitsOne(t *testing.T) {
fixture := newCLIFixture(t, "")
fixture.setEnvironment(t)
t.Setenv("THOTHCTL_FAKE_WORKSPACE_RESULT", "not-json")
var stdout, stderr bytes.Buffer
got := run(context.Background(), []string{"--installation", fixture.installationPath, "workspace", "inspect", "--workspace", "psd"}, &stdout, &stderr)
if got != 1 || !strings.Contains(stderr.String(), "invalid workspace result") {
t.Fatalf("exit=%d stderr=%q; want operational exit 1", got, stderr.String())
}
}
func TestRunWorkspaceUnsafeOutputDoesNotInvokeCompose(t *testing.T) {
fixture := newCLIFixture(t, "")
fixture.setEnvironment(t)
output := filepath.Join(fixture.root, "existing.yaml")
if err := os.WriteFile(output, []byte("existing"), 0o600); err != nil {
t.Fatal(err)
}
var stdout, stderr bytes.Buffer
got := run(context.Background(), []string{"--installation", fixture.installationPath, "workspace", "schema", "suggest-fks", "--workspace", "psd", "--output", output}, &stdout, &stderr)
if got != 2 || !strings.Contains(stderr.String(), "unsafe output file") {
t.Fatalf("exit=%d stderr=%q; want unsafe host failure", got, stderr.String())
}
assertDockerNotInvoked(t, fixture)
}
func TestRunWorkspaceRejectsInvalidCommandBeforeDocker(t *testing.T) {
fixture := newCLIFixture(t, "")
fixture.setEnvironment(t)
@@ -750,6 +803,9 @@ case " $* " in
else
printf '%s\n' '{"volumes":{"settings":{}},"services":{"core":{"image":"thothii-core:local","environment":{"THT_LLM_URL":"https://llm.example.invalid"}}}}'
fi ;;
*" run --rm --no-deps --no-TTY workspace-maintenance "*)
printf '%s\n' "${THOTHCTL_FAKE_WORKSPACE_RESULT:-}"
exit "${THOTHCTL_FAKE_WORKSPACE_EXIT:-0}" ;;
*" run --rm --no-deps --no-TTY session-migrate "*)
if [ "${THOTHCTL_FAKE_MIGRATION_EXIT:-0}" -ne 0 ]; then
printf '%s\n' "$THOTHCTL_FAKE_MIGRATION_FAILURE" >&2
@@ -805,6 +861,8 @@ func (f cliFixture) setEnvContents(t *testing.T, env string) {
t.Setenv("THOTHCTL_FAKE_CONFIG", "")
t.Setenv("THOTHCTL_FAKE_MIGRATION_FAILURE", "")
t.Setenv("THOTHCTL_FAKE_MIGRATION_EXIT", "0")
t.Setenv("THOTHCTL_FAKE_WORKSPACE_RESULT", "")
t.Setenv("THOTHCTL_FAKE_WORKSPACE_EXIT", "0")
}
func (f cliFixture) setProfile(t *testing.T, profile string) {
@@ -0,0 +1,7 @@
//go:build !windows
package compose
import "syscall"
func processExists(pid int) bool { return syscall.Kill(pid, 0) == nil }
@@ -0,0 +1,5 @@
//go:build windows
package compose
func processExists(int) bool { return true }
@@ -3,6 +3,8 @@
package compose
import (
"context"
"io"
"os/exec"
"syscall"
)
@@ -16,3 +18,7 @@ func terminateOwnedProcess(c *exec.Cmd) {
_ = c.Process.Kill()
}
}
func runBoundedPlatform(context.Context, Runner, []string, io.Reader, CaptureLimits) (Result, bool, error) {
return Result{}, false, nil
}
@@ -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)
}
@@ -8,11 +8,14 @@ import (
"testing"
)
func TestOwnedProcessConfigurationStartsSuspended(t *testing.T) {
func TestOwnedProcessUsesCreationTimeJobAttribute(t *testing.T) {
if procThreadAttributeJobList == 0 {
t.Fatal("owned Windows process must carry a creation-time job-list attribute")
}
cmd := exec.Command("cmd.exe")
configureOwnedProcess(cmd)
if cmd.SysProcAttr == nil || cmd.SysProcAttr.CreationFlags&windows.CREATE_SUSPENDED == 0 {
t.Fatal("owned Windows process must start suspended until job assignment")
if cmd.SysProcAttr != nil && cmd.SysProcAttr.CreationFlags&windows.CREATE_SUSPENDED != 0 {
t.Fatal("owned Windows process must not rely on a post-start assignment race")
}
}
@@ -57,6 +57,9 @@ func (r Runner) Run(ctx context.Context, args []string, stdin io.Reader) (Result
// RunBounded streams each pipe into a preallocated fixed-capacity collector. It owns the
// process group and tears it down on the first overflow or cancellation.
func (r Runner) RunBounded(ctx context.Context, args []string, stdin io.Reader, limits CaptureLimits) (Result, error) {
if result, handled, err := runBoundedPlatform(ctx, r, args, stdin, limits); handled {
return result, err
}
if limits.StdoutBytes < 0 || limits.StderrBytes < 0 {
return Result{}, ErrOutputLimit
}
@@ -7,6 +7,7 @@ import (
"os/exec"
"path/filepath"
"runtime"
"strconv"
"strings"
"testing"
"time"
@@ -123,6 +124,35 @@ func TestRunnerBoundedCancellationWinsOverChildKill(t *testing.T) {
}
}
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")
@@ -75,3 +75,51 @@ func TestWriteCanonicalExclusiveRejectsExistingAndCreatesPrivateFile(t *testing.
t.Fatalf("replacement = %v", err)
}
}
func TestReadCanonicalRegularRejectsHardlinkAndDirectory(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
original := filepath.Join(root, "original.sql")
if err := os.WriteFile(original, []byte("select 1"), 0o600); err != nil {
t.Fatal(err)
}
hardlink := filepath.Join(root, "hardlink.sql")
if err := os.Link(original, hardlink); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalRegular(hardlink, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("hardlink error = %v, want ErrUnsafeFile", err)
}
directory := filepath.Join(root, "directory.sql")
if err := os.Mkdir(directory, 0o700); err != nil {
t.Fatal(err)
}
if _, err := ReadCanonicalRegular(directory, 1024); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("directory error = %v, want ErrUnsafeFile", err)
}
}
func TestValidateCanonicalOutputPathRejectsExistingDirectoryAndSymlink(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
directory := filepath.Join(root, "existing.yaml")
if err := os.Mkdir(directory, 0o700); err != nil {
t.Fatal(err)
}
if err := ValidateCanonicalOutputPath(directory); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("existing directory error = %v, want ErrUnsafeFile", err)
}
target := filepath.Join(root, "target.yaml")
if err := os.WriteFile(target, []byte("target"), 0o600); err != nil {
t.Fatal(err)
}
link := filepath.Join(root, "link.yaml")
testsupport.SymlinkOrSkip(t, target, link)
if err := ValidateCanonicalOutputPath(link); !errors.Is(err, ErrUnsafeFile) {
t.Fatalf("output symlink error = %v, want ErrUnsafeFile", err)
}
}
+1 -1
View File
@@ -79,7 +79,7 @@ func closeUnixDescriptors(descriptors []int) {
}
func writeCanonicalExclusive(path string, contents []byte, mode fs.FileMode) error {
if err := ValidateCanonicalPath(path); err != nil || len(contents) > 16<<20 || mode.Perm() == 0 || mode.Perm()&0o077 != 0 {
if err := ValidateCanonicalPath(path); err != nil || len(contents) > 16<<20 || mode.Perm() != 0o600 {
return ErrUnsafeFile
}
components := strings.Split(strings.TrimPrefix(path, string(os.PathSeparator)), string(os.PathSeparator))
@@ -6,8 +6,11 @@ import (
"io/fs"
"os"
"path/filepath"
"runtime"
"strings"
"unsafe"
"golang.org/x/sys/windows"
)
@@ -119,7 +122,7 @@ func closeWindowsHandles(handles []windows.Handle) {
}
func writeCanonicalExclusive(path string, contents []byte, mode fs.FileMode) error {
if err := ValidateCanonicalPath(path); err != nil || len(contents) > 16<<20 || mode.Perm() == 0 || mode.Perm()&0o077 != 0 {
if err := ValidateCanonicalPath(path); err != nil || len(contents) > 16<<20 || mode.Perm() != 0o600 {
return ErrUnsafeFile
}
parent, retainedParents, err := openWindowsParents(path)
@@ -129,7 +132,11 @@ func writeCanonicalExclusive(path string, contents []byte, mode fs.FileMode) err
defer closeWindowsHandles(retainedParents)
// Parent handles stay open with delete sharing denied until publication and
// identity recheck complete; this is the Windows equivalent of retained dirfds.
h, err := windows.CreateFile(windows.StringToUTF16Ptr(filepath.Join(parent, filepath.Base(path))), windows.GENERIC_WRITE, 0, nil, windows.CREATE_NEW, windows.FILE_ATTRIBUTE_NORMAL|windows.FILE_FLAG_OPEN_REPARSE_POINT, 0)
securityDescriptor, securityAttributes, err := ownerOnlySecurityAttributes()
if err != nil {
return ErrUnsafeFile
}
h, err := windows.CreateFile(windows.StringToUTF16Ptr(filepath.Join(parent, filepath.Base(path))), windows.GENERIC_WRITE, 0, securityAttributes, windows.CREATE_NEW, windows.FILE_ATTRIBUTE_NORMAL|windows.FILE_FLAG_OPEN_REPARSE_POINT, 0)
if err != nil {
return ErrUnsafeFile
}
@@ -139,7 +146,11 @@ func writeCanonicalExclusive(path string, contents []byte, mode fs.FileMode) err
return ErrUnsafeFile
}
defer f.Close()
if err := f.Chmod(mode); err != nil {
// Go's Windows Chmod only toggles the read-only attribute; it cannot enforce
// owner-only permissions. The restrictive DACL is installed at creation time
// through SECURITY_ATTRIBUTES above.
_ = securityDescriptor
if mode.Perm() != 0o600 {
return ErrUnsafeFile
}
if _, err := f.Write(contents); err != nil {
@@ -202,3 +213,47 @@ func validateCanonicalOutputPath(path string) error {
func sameWindowsFile(a, b windows.ByHandleFileInformation) bool {
return a.VolumeSerialNumber == b.VolumeSerialNumber && a.FileIndexHigh == b.FileIndexHigh && a.FileIndexLow == b.FileIndexLow
}
// ownerOnlySecurityAttributes builds a non-inheriting DACL granting the current
// token's user read/write access and no access to inherited or other trustees.
// The descriptor is retained by the returned Go value for the duration of CreateFile.
func ownerOnlySecurityAttributes() (*windows.SECURITY_DESCRIPTOR, *windows.SecurityAttributes, error) {
user, err := windows.GetCurrentProcessToken().GetTokenUser()
if err != nil || user == nil || user.User.Sid == nil {
return nil, nil, ErrUnsafeFile
}
var pinner runtime.Pinner
pinner.Pin(user.User.Sid)
defer pinner.Unpin()
trustee := windows.TRUSTEE{
TrusteeForm: windows.TRUSTEE_IS_SID,
TrusteeType: windows.TRUSTEE_IS_USER,
TrusteeValue: windows.TrusteeValueFromSID(user.User.Sid),
}
entries := []windows.EXPLICIT_ACCESS{{
AccessPermissions: windows.FILE_GENERIC_READ | windows.FILE_GENERIC_WRITE,
AccessMode: windows.SET_ACCESS,
Inheritance: windows.NO_INHERITANCE,
Trustee: trustee,
}}
descriptor, err := windows.BuildSecurityDescriptor(nil, nil, entries, nil, nil)
if err != nil || descriptor == nil {
return nil, nil, ErrUnsafeFile
}
acl, _, err := descriptor.DACL()
if err != nil || acl == nil {
return nil, nil, ErrUnsafeFile
}
if err := descriptor.SetControl(windows.SE_DACL_PROTECTED, windows.SE_DACL_PROTECTED); err != nil {
return nil, nil, ErrUnsafeFile
}
// BuildSecurityDescriptor returns a self-relative descriptor. Re-using its
// DACL as the creation descriptor is valid, and the explicit DACL has no
// inheritable ACEs; the protected flag is applied by the kernel on creation.
_ = acl
attrs := &windows.SecurityAttributes{
Length: uint32(unsafe.Sizeof(windows.SecurityAttributes{})),
SecurityDescriptor: descriptor,
}
return descriptor, attrs, nil
}
@@ -72,3 +72,29 @@ func TestWriteCanonicalExclusiveRequiresRestrictiveMode(t *testing.T) {
t.Fatal("accepted non-restrictive output mode")
}
}
func TestWriteCanonicalExclusiveCreatesProtectedOwnerOnlyDACL(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
path := filepath.Join(root, "candidate.yaml")
if err := writeCanonicalExclusive(path, []byte("x"), 0o600); err != nil {
t.Fatal(err)
}
sd, err := windows.GetNamedSecurityInfo(path, windows.SE_FILE_OBJECT, windows.DACL_SECURITY_INFORMATION)
if err != nil {
t.Fatal(err)
}
control, _, err := sd.Control()
if err != nil {
t.Fatal(err)
}
if control&windows.SE_DACL_PROTECTED == 0 {
t.Fatalf("output DACL control = %#x, want protected", control)
}
acl, _, err := sd.DACL()
if err != nil || acl == nil || acl.AceCount != 1 {
t.Fatalf("output DACL = %#v, err=%v; want one owner ACE", acl, err)
}
}
@@ -29,8 +29,11 @@ const (
maxAssumptions = 256
maxAssumptionBytes = 256
maxResult = 1 << 20
maxCandidate = 700 << 10
maxAnnotations = 16 << 20
// Requests carry up to 16 MiB of raw SQL or annotation bytes encoded as base64.
// This is deliberately independent from maxResult, which bounds child stdout.
maxRequest = 24 << 20
maxCandidate = 700 << 10
maxAnnotations = 16 << 20
)
var workspaceIDPattern = regexp.MustCompile(`^[a-z][a-z0-9-]{2,62}$`)
@@ -392,7 +395,7 @@ func makeInput(c Command) (inputEnvelope, string, error) {
if e != nil {
return env, "", e
}
if len(payload) > maxResult {
if len(payload) > maxRequest {
return env, "", errors.New("request exceeds limit")
}
return env, string(payload), nil
@@ -461,8 +464,8 @@ func Run(ctx context.Context, installation config.Installation, runner compose.R
}
payload := []byte(generated)
if stdin != nil {
payload, e = io.ReadAll(io.LimitReader(stdin, maxResult+1))
if e != nil || len(payload) > maxResult {
payload, e = io.ReadAll(io.LimitReader(stdin, maxRequest+1))
if e != nil || len(payload) > maxRequest {
return Result{}, errors.New("request exceeds limit")
}
if e = validateIngress(payload, env); e != nil {
@@ -83,9 +83,10 @@ func TestRunUsesBase64BasenameIngressAndNoTTY(t *testing.T) {
if err := os.WriteFile(sqlPath, []byte("select 1"), 0o600); err != nil {
t.Fatal(err)
}
out := filepath.Join(root, "candidate.yaml")
stdinCapture := filepath.Join(root, "stdin.json")
argsCapture := filepath.Join(root, "args.txt")
resultJSON := `{"schemaVersion":1,"status":"succeeded","code":"ok","workspaceId":"psd","workspaceRevision":"0123456789012345678901234567890123456789","descriptorBlob":"abcdefabcdefabcdefabcdefabcdefabcdefabcd","operation":"suggest-fks","runId":"0123456789abcdef0123456789abcdef","completedStages":[]}`
script := "#!/bin/sh\ncat >/dev/null\nprintf '%s' '" + resultJSON + "'\n"
script := "#!/bin/sh\nprintf '%s\n' \"$@\" > '" + argsCapture + "'\ncat > '" + stdinCapture + "'\nprintf '%s' '" + resultJSON + "'\n"
fake := filepath.Join(root, "docker")
if err := os.WriteFile(fake, []byte(script), 0o700); err != nil {
t.Fatal(err)
@@ -101,7 +102,20 @@ func TestRunUsesBase64BasenameIngressAndNoTTY(t *testing.T) {
if got.Code != "ok" {
t.Fatalf("result = %#v", got)
}
_ = out
stdinBytes, err := os.ReadFile(stdinCapture)
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(stdinBytes), sqlPath) || !strings.Contains(string(stdinBytes), "contentBase64") || !strings.Contains(string(stdinBytes), "schema.sql") {
t.Fatalf("stdin envelope = %q; want basename/base64 without host path", stdinBytes)
}
argsBytes, err := os.ReadFile(argsCapture)
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(argsBytes), sqlPath) || !strings.Contains(string(argsBytes), "--operation\nsuggest-fks") || !strings.Contains(string(argsBytes), "--workspace\npsd") {
t.Fatalf("child args = %q; want closed operation/workspace args", argsBytes)
}
}
func TestRunPublishesOnlyVerifiedCandidateExport(t *testing.T) {
@@ -295,3 +309,43 @@ func TestParseWorkspaceAcceptsAssumptionAndEnforcesLimits(t *testing.T) {
t.Fatal("accepted 33 SQL files")
}
}
func TestMakeInputAcceptsMaximumSingleSQLFile(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
path := filepath.Join(root, "schema.sql")
if err := os.WriteFile(path, bytes.Repeat([]byte("x"), maxSQLFile), 0o600); err != nil {
t.Fatal(err)
}
command := SuggestFksRequest{WorkspaceID: "psd", FromSQL: []string{path}}
if _, payload, err := makeInput(command); err != nil {
t.Fatalf("makeInput() error = %v", err)
} else if len(payload) <= maxResult {
t.Fatalf("payload length = %d, want larger than result cap %d", len(payload), maxResult)
}
}
func TestMakeInputAcceptsMaximumAnnotationWithReviewDigestPair(t *testing.T) {
root, err := filepath.EvalSymlinks(t.TempDir())
if err != nil {
t.Fatal(err)
}
path := filepath.Join(root, "annotations.md")
contents := bytes.Repeat([]byte("a"), maxAnnotations)
if err := os.WriteFile(path, contents, 0o600); err != nil {
t.Fatal(err)
}
command := CheckSchemaRequest{WorkspaceID: "psd", Resume: strings.Repeat("0", 32), Annotations: path, ReviewedCandidates: DigestBytes([]byte("reviewed"))}
env, payload, err := makeInput(command)
if err != nil {
t.Fatal(err)
}
if env.Annotations == nil || env.ReviewedCandidates == "" || len(payload) <= maxResult {
t.Fatalf("annotation envelope = %#v payload=%d; want pair and request larger than result cap", env, len(payload))
}
if err := validateIngress([]byte(payload), env); err != nil {
t.Fatalf("validateIngress() = %v", err)
}
}