From 3a9e0c7887ac5e47a062a579f3ae80ec70e0d4a8 Mon Sep 17 00:00:00 2001 From: naiba Date: Mon, 20 Jul 2026 04:43:28 +0000 Subject: [PATCH] test(agentcompat): add process supervision harness Co-authored-by: naiba/CloudCode --- .../agentcompat/internal/process/cleanup.go | 46 +++ .../agentcompat/internal/process/fd_path.go | 37 +++ .../internal/process/fd_path_test.go | 30 ++ .../internal/process/helper_test.go | 231 +++++++++++++++ .../agentcompat/internal/process/log.go | 127 ++++++++ .../agentcompat/internal/process/proc.go | 159 ++++++++++ .../agentcompat/internal/process/sampler.go | 109 +++++++ .../internal/process/sampler_test.go | 209 +++++++++++++ .../process/sqlite_journal_identity.go | 76 +++++ .../internal/process/sqlite_journal_watch.go | 195 ++++++++++++ .../process/sqlite_journal_watch_test.go | 186 ++++++++++++ .../internal/process/supervisor.go | 278 ++++++++++++++++++ .../process/supervisor_agentcompat.go | 7 + .../internal/process/supervisor_test.go | 195 ++++++++++++ .../agentcompat/internal/workspace/build.go | 56 ++++ .../internal/workspace/listener.go | 90 ++++++ .../agentcompat/internal/workspace/log.go | 117 ++++++++ .../internal/workspace/residue_test.go | 72 +++++ .../supervisor_exit_integration_test.go | 250 ++++++++++++++++ .../workspace/supervisor_integration_test.go | 104 +++++++ .../internal/workspace/workspace.go | 265 +++++++++++++++++ .../internal/workspace/workspace_test.go | 211 +++++++++++++ 22 files changed, 3050 insertions(+) create mode 100644 integration/agentcompat/internal/process/cleanup.go create mode 100644 integration/agentcompat/internal/process/fd_path.go create mode 100644 integration/agentcompat/internal/process/fd_path_test.go create mode 100644 integration/agentcompat/internal/process/helper_test.go create mode 100644 integration/agentcompat/internal/process/log.go create mode 100644 integration/agentcompat/internal/process/proc.go create mode 100644 integration/agentcompat/internal/process/sampler.go create mode 100644 integration/agentcompat/internal/process/sampler_test.go create mode 100644 integration/agentcompat/internal/process/sqlite_journal_identity.go create mode 100644 integration/agentcompat/internal/process/sqlite_journal_watch.go create mode 100644 integration/agentcompat/internal/process/sqlite_journal_watch_test.go create mode 100644 integration/agentcompat/internal/process/supervisor.go create mode 100644 integration/agentcompat/internal/process/supervisor_agentcompat.go create mode 100644 integration/agentcompat/internal/process/supervisor_test.go create mode 100644 integration/agentcompat/internal/workspace/build.go create mode 100644 integration/agentcompat/internal/workspace/listener.go create mode 100644 integration/agentcompat/internal/workspace/log.go create mode 100644 integration/agentcompat/internal/workspace/residue_test.go create mode 100644 integration/agentcompat/internal/workspace/supervisor_exit_integration_test.go create mode 100644 integration/agentcompat/internal/workspace/supervisor_integration_test.go create mode 100644 integration/agentcompat/internal/workspace/workspace.go create mode 100644 integration/agentcompat/internal/workspace/workspace_test.go diff --git a/integration/agentcompat/internal/process/cleanup.go b/integration/agentcompat/internal/process/cleanup.go new file mode 100644 index 00000000..36afb9b3 --- /dev/null +++ b/integration/agentcompat/internal/process/cleanup.go @@ -0,0 +1,46 @@ +//go:build linux + +package process + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" +) + +type CleanupRecord struct { + Name string `json:"name"` + PID int `json:"pid"` + Forced bool `json:"forced"` + Error string `json:"error,omitempty"` +} + +type CleanupReceipt struct { + Passed bool `json:"passed"` + Forced bool `json:"forced"` + Processes []CleanupRecord `json:"processes"` +} + +func NewCleanupReceipt(records []CleanupRecord) CleanupReceipt { + receipt := CleanupReceipt{Passed: true, Processes: append([]CleanupRecord(nil), records...)} + for _, record := range records { + receipt.Forced = receipt.Forced || record.Forced + receipt.Passed = receipt.Passed && record.Error == "" + } + return receipt +} + +func WriteCleanupReceipt(path string, receipt CleanupReceipt) error { + data, err := json.MarshalIndent(receipt, "", " ") + if err != nil { + return fmt.Errorf("marshal cleanup receipt: %w", err) + } + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return fmt.Errorf("create cleanup receipt directory: %w", err) + } + if err := os.WriteFile(path, append(data, '\n'), 0o600); err != nil { + return fmt.Errorf("write cleanup receipt: %w", err) + } + return nil +} diff --git a/integration/agentcompat/internal/process/fd_path.go b/integration/agentcompat/internal/process/fd_path.go new file mode 100644 index 00000000..15675608 --- /dev/null +++ b/integration/agentcompat/internal/process/fd_path.go @@ -0,0 +1,37 @@ +//go:build linux + +package process + +import ( + "errors" + "os" + "path/filepath" + "strconv" + "strings" +) + +func ProcessHasOpenPath(pid int, path string) (bool, error) { + if pid < 1 || !filepath.IsAbs(path) { + return false, errors.New("invalid process path query") + } + directory := filepath.Join("/proc", strconv.Itoa(pid), "fd") + entries, err := os.ReadDir(directory) + if err != nil { + return false, err + } + wanted := filepath.Clean(path) + for _, entry := range entries { + target, err := os.Readlink(filepath.Join(directory, entry.Name())) + if err != nil { + if os.IsNotExist(err) { + continue + } + return false, err + } + target = strings.TrimSuffix(target, " (deleted)") + if filepath.IsAbs(target) && filepath.Clean(target) == wanted { + return true, nil + } + } + return false, nil +} diff --git a/integration/agentcompat/internal/process/fd_path_test.go b/integration/agentcompat/internal/process/fd_path_test.go new file mode 100644 index 00000000..f1e9d7c2 --- /dev/null +++ b/integration/agentcompat/internal/process/fd_path_test.go @@ -0,0 +1,30 @@ +//go:build linux + +package process + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestProcessHasOpenPathTogglesWithDescriptorLifecycle(t *testing.T) { + // Given + path := filepath.Join(t.TempDir(), "dashboard.sqlite-journal") + require.NoError(t, os.WriteFile(path, []byte("journal"), 0o600)) + file, err := os.Open(path) + require.NoError(t, err) + + // When + held, err := ProcessHasOpenPath(os.Getpid(), path) + require.NoError(t, err) + require.NoError(t, file.Close()) + released, err := ProcessHasOpenPath(os.Getpid(), path) + + // Then + require.NoError(t, err) + require.True(t, held) + require.False(t, released) +} diff --git a/integration/agentcompat/internal/process/helper_test.go b/integration/agentcompat/internal/process/helper_test.go new file mode 100644 index 00000000..1c219ce4 --- /dev/null +++ b/integration/agentcompat/internal/process/helper_test.go @@ -0,0 +1,231 @@ +//go:build linux + +package process + +import ( + "bufio" + "context" + "errors" + "fmt" + "io" + "net" + "os" + "os/exec" + "os/signal" + "path/filepath" + "strconv" + "strings" + "syscall" + "testing" + "time" +) + +const ( + helperModeEnv = "NEZHA_AGENTCOMPAT_PROCESS_HELPER" + helperMarkerEnv = "NEZHA_AGENTCOMPAT_PROCESS_MARKER" + helperFDEnv = "NEZHA_AGENTCOMPAT_PROCESS_FD" +) + +func TestProcessHelper(t *testing.T) { + switch os.Getenv(helperModeEnv) { + case "": + return + case "clean": + fmt.Println("READY") + case "credential": + marker := os.Getenv(helperMarkerEnv) + if err := os.WriteFile(marker, []byte(fmt.Sprintf("%d:%d", os.Getuid(), os.Getgid())), 0o600); err != nil { + t.Fatal(err) + } + case "block": + fmt.Println("READY") + _, _ = io.Copy(io.Discard, os.Stdin) + case "tree": + runTreeHelper(t, false) + case "force-tree": + runTreeHelper(t, true) + case "grandchild": + runGrandchildHelper(t) + case "ignore-term-grandchild": + signal.Ignore(syscall.SIGTERM) + runGrandchildHelper(t) + case "ignore-term": + signal.Ignore(syscall.SIGTERM) + fmt.Println("READY") + waitForSignal(syscall.SIGINT) + case "listener": + runListenerHelper(t) + case "logs": + fmt.Println("READY") + fmt.Println("Authorization: Bearer eyJsecret.secret.secret password=top-secret") + fmt.Println(strings.Repeat("x", 1024)) + case "interrupt-probe": + runInterruptProbeHelper(t) + default: + t.Fatalf("unknown helper mode %q", os.Getenv(helperModeEnv)) + } +} + +func runInterruptProbeHelper(t *testing.T) { + t.Helper() + marker := os.Getenv(helperMarkerEnv) + ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGTERM) + defer stop() + supervisor := newHelperSupervisor(ctx, "tree", []string{helperMarkerEnv + "=" + marker}) + requireNoError(t, supervisor.Start()) + requireNoError(t, supervisor.WaitReady(t.Context())) + grandchildPID := readPID(t, marker) + if err := os.WriteFile(marker+".leader", []byte(strconv.Itoa(supervisor.PID())), 0o600); err != nil { + t.Fatal(err) + } + fmt.Println("PROBE_READY") + <-ctx.Done() + select { + case <-supervisor.cleanupDone: + case <-time.After(2 * time.Second): + t.Fatal("context cancellation did not complete process-tree cleanup") + } + requirePIDGone(t, supervisor.PID()) + requirePIDGone(t, grandchildPID) +} + +func runTreeHelper(t *testing.T, ignoreTermination bool) { + t.Helper() + child := exec.Command(os.Args[0], "-test.run=^TestProcessHelper$") + childMode := "grandchild" + if ignoreTermination { + childMode = "ignore-term-grandchild" + signal.Ignore(syscall.SIGTERM) + } + child.Env = append(os.Environ(), helperModeEnv+"="+childMode) + child.Stdout = os.Stdout + child.Stderr = os.Stderr + if err := child.Start(); err != nil { + t.Fatal(err) + } + if ignoreTermination { + waitForSignal(syscall.SIGINT) + _ = child.Wait() + return + } + waitForSignal(syscall.SIGTERM) + _ = child.Wait() +} + +func runGrandchildHelper(t *testing.T) { + t.Helper() + marker := os.Getenv(helperMarkerEnv) + if marker == "" { + t.Fatal("helper marker is empty") + } + if err := os.WriteFile(marker, []byte(strconv.Itoa(os.Getpid())), 0o600); err != nil { + t.Fatal(err) + } + fmt.Println("READY") + waitForSignal(syscall.SIGTERM) +} + +func runListenerHelper(t *testing.T) { + t.Helper() + descriptor, err := strconv.Atoi(os.Getenv(helperFDEnv)) + if err != nil { + t.Fatal(err) + } + file := os.NewFile(uintptr(descriptor), "inherited-listener") + listener, err := net.FileListener(file) + if err != nil { + t.Fatal(err) + } + _ = file.Close() + defer listener.Close() + fmt.Println("READY") + waitForSignal(syscall.SIGTERM) +} + +func waitForSignal(expected os.Signal) { + signals := make(chan os.Signal, 1) + signal.Notify(signals, expected) + defer signal.Stop(signals) + <-signals +} + +func newHelperSupervisor(ctx context.Context, mode string, environment []string) *Supervisor { + return NewSupervisor(ctx, Spec{ + Name: "helper-" + mode, + Path: os.Args[0], + Args: []string{"-test.run=^TestProcessHelper$"}, + Env: append(append(os.Environ(), helperModeEnv+"="+mode), environment...), + MaxLogBytes: 1024, + TerminateTimeout: 100 * time.Millisecond, + KillTimeout: time.Second, + Stdout: os.Stdout, + Stderr: os.Stderr, + Readiness: func(_ Stream, line string) bool { + return strings.Contains(line, "READY") + }, + }) +} + +func startBlockingHelper(t *testing.T) (*exec.Cmd, func()) { + t.Helper() + command := exec.Command(os.Args[0], "-test.run=^TestProcessHelper$") + command.Env = append(os.Environ(), helperModeEnv+"=block") + input, err := command.StdinPipe() + requireNoError(t, err) + output, err := command.StdoutPipe() + requireNoError(t, err) + requireNoError(t, command.Start()) + scanner := bufio.NewScanner(output) + if !scanner.Scan() || scanner.Text() != "READY" { + t.Fatalf("helper readiness = %q, err = %v", scanner.Text(), scanner.Err()) + } + return command, func() { _ = input.Close() } +} + +func startCleanHelper(t *testing.T) *exec.Cmd { + t.Helper() + command := exec.Command(os.Args[0], "-test.run=^TestProcessHelper$") + command.Env = append(os.Environ(), helperModeEnv+"=clean") + requireNoError(t, command.Start()) + return command +} + +func reapHelper(command *exec.Cmd) { + if command.ProcessState == nil { + _ = command.Process.Kill() + _ = command.Wait() + } +} + +func readPID(t *testing.T, path string) int { + t.Helper() + data, err := os.ReadFile(path) + requireNoError(t, err) + pid, err := strconv.Atoi(strings.TrimSpace(string(data))) + requireNoError(t, err) + return pid +} + +func requirePIDGone(t *testing.T, pid int) { + t.Helper() + _, err := os.Stat(filepath.Join("/proc", strconv.Itoa(pid))) + if !errors.Is(err, os.ErrNotExist) { + t.Fatalf("PID %d remains: %v", pid, err) + } +} + +func containsPID(pids []int, target int) bool { + for _, pid := range pids { + if pid == target { + return true + } + } + return false +} + +func requireNoError(t *testing.T, err error) { + t.Helper() + if err != nil { + t.Fatal(err) + } +} diff --git a/integration/agentcompat/internal/process/log.go b/integration/agentcompat/internal/process/log.go new file mode 100644 index 00000000..38fb2648 --- /dev/null +++ b/integration/agentcompat/internal/process/log.go @@ -0,0 +1,127 @@ +//go:build linux + +package process + +import ( + "bytes" + "errors" + "fmt" + "io" + "sync" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/evidence" +) + +const truncationMarker = "[TRUNCATED]\n" + +type boundedLog struct { + mu sync.Mutex + destination io.Writer + maxBytes int + written int + pending []byte + dropLine bool + truncated bool + closed bool + onLine func(string) + writeErr error +} + +func newBoundedLog(destination io.Writer, maxBytes int, onLine func(string)) *boundedLog { + return &boundedLog{destination: destination, maxBytes: maxBytes, onLine: onLine} +} + +func (log *boundedLog) Write(data []byte) (int, error) { + log.mu.Lock() + defer log.mu.Unlock() + if log.closed { + return 0, errors.New("write closed process log") + } + inputLength := len(data) + for len(data) > 0 { + newline := bytes.IndexByte(data, '\n') + if newline < 0 { + log.appendFragment(data) + break + } + log.appendFragment(data[:newline+1]) + if err := log.flushLine(); err != nil { + return 0, err + } + data = data[newline+1:] + } + return inputLength, nil +} + +func (log *boundedLog) appendFragment(fragment []byte) { + if log.dropLine { + return + } + if len(log.pending)+len(fragment) > log.maxBytes { + log.pending = nil + log.dropLine = true + log.truncated = true + return + } + log.pending = append(log.pending, fragment...) +} + +func (log *boundedLog) flushLine() error { + if log.dropLine { + log.dropLine = false + return log.writeMarker() + } + redacted := evidence.Redact(string(log.pending)) + log.pending = nil + if log.onLine != nil { + log.onLine(redacted) + } + if len(redacted) > log.maxBytes-log.written { + log.truncated = true + return log.writeMarker() + } + if log.destination != nil && redacted != "" { + written, err := io.WriteString(log.destination, redacted) + log.written += written + if err != nil { + log.writeErr = fmt.Errorf("write process log: %w", err) + return log.writeErr + } + } + return nil +} + +func (log *boundedLog) writeMarker() error { + if log.destination == nil || log.written >= log.maxBytes { + return nil + } + marker := truncationMarker + if len(marker) > log.maxBytes-log.written { + marker = marker[:log.maxBytes-log.written] + } + written, err := io.WriteString(log.destination, marker) + log.written += written + if err != nil { + log.writeErr = fmt.Errorf("write process log marker: %w", err) + return log.writeErr + } + return nil +} + +func (log *boundedLog) Close() { + log.mu.Lock() + defer log.mu.Unlock() + if log.closed { + return + } + if len(log.pending) > 0 || log.dropLine { + _ = log.flushLine() + } + log.closed = true +} + +func (log *boundedLog) Truncated() bool { + log.mu.Lock() + defer log.mu.Unlock() + return log.truncated +} diff --git a/integration/agentcompat/internal/process/proc.go b/integration/agentcompat/internal/process/proc.go new file mode 100644 index 00000000..f028a247 --- /dev/null +++ b/integration/agentcompat/internal/process/proc.go @@ -0,0 +1,159 @@ +//go:build linux + +package process + +import ( + "bufio" + "errors" + "fmt" + "os" + "path/filepath" + "sort" + "strconv" + "strings" + "syscall" +) + +func readRSSBytes(pid int) (uint64, error) { + path := filepath.Join("/proc", strconv.Itoa(pid), "status") + file, err := os.Open(path) + if err != nil { + return 0, err + } + defer file.Close() + scanner := bufio.NewScanner(file) + for scanner.Scan() { + fields := strings.Fields(scanner.Text()) + if len(fields) == 3 && fields[0] == "VmRSS:" && fields[2] == "kB" { + kilobytes, err := strconv.ParseUint(fields[1], 10, 64) + if err != nil { + return 0, fmt.Errorf("parse VmRSS: %w", err) + } + return kilobytes * 1024, nil + } + } + if err := scanner.Err(); err != nil { + return 0, fmt.Errorf("read %s: %w", path, err) + } + return 0, errors.New("VmRSS not found") +} + +func descendantPIDs(rootPID int) ([]int, error) { + entries, err := os.ReadDir("/proc") + if err != nil { + return nil, fmt.Errorf("read /proc: %w", err) + } + children := make(map[int][]int) + for _, entry := range entries { + pid, err := strconv.Atoi(entry.Name()) + if err != nil || !entry.IsDir() { + continue + } + parentPID, err := readParentPID(pid) + if err != nil { + // /proc is a live snapshot: an unrelated process can disappear + // between ReadDir and reading stat. Root PID reads stay strict. + if os.IsNotExist(err) || errors.Is(err, syscall.ESRCH) { + continue + } + return nil, err + } + children[parentPID] = append(children[parentPID], pid) + } + descendants := make([]int, 0) + queue := append([]int(nil), children[rootPID]...) + for len(queue) > 0 { + pid := queue[0] + queue = queue[1:] + descendants = append(descendants, pid) + queue = append(queue, children[pid]...) + } + sort.Ints(descendants) + return descendants, nil +} + +func readParentPID(pid int) (int, error) { + path := filepath.Join("/proc", strconv.Itoa(pid), "stat") + data, err := os.ReadFile(path) + if err != nil { + return 0, err + } + closingParenthesis := strings.LastIndexByte(string(data), ')') + if closingParenthesis < 0 { + return 0, fmt.Errorf("parse %s: missing command terminator", path) + } + fields := strings.Fields(string(data[closingParenthesis+1:])) + if len(fields) < 2 { + return 0, fmt.Errorf("parse %s: missing parent PID", path) + } + parentPID, err := strconv.Atoi(fields[1]) + if err != nil { + return 0, fmt.Errorf("parse %s parent PID: %w", path, err) + } + return parentPID, nil +} + +func processFDs(pid int) (int, map[uint64]struct{}, error) { + directory := filepath.Join("/proc", strconv.Itoa(pid), "fd") + entries, err := os.ReadDir(directory) + if err != nil { + return 0, nil, err + } + count := 0 + sockets := make(map[uint64]struct{}) + for _, entry := range entries { + descriptor, err := strconv.Atoi(entry.Name()) + if err != nil || descriptor < 3 { + continue + } + target, err := os.Readlink(filepath.Join(directory, entry.Name())) + if err != nil { + if os.IsNotExist(err) { + continue + } + return 0, nil, err + } + count++ + if inode, exists := parseSocketInode(target); exists { + sockets[inode] = struct{}{} + } + } + return count, sockets, nil +} + +func parseSocketInode(target string) (uint64, bool) { + if !strings.HasPrefix(target, "socket:[") || !strings.HasSuffix(target, "]") { + return 0, false + } + inode, err := strconv.ParseUint(strings.TrimSuffix(strings.TrimPrefix(target, "socket:["), "]"), 10, 64) + return inode, err == nil +} + +func listeningSocketInodes(pid int, protocol string) (map[uint64]struct{}, error) { + path := filepath.Join("/proc", strconv.Itoa(pid), "net", protocol) + file, err := os.Open(path) + if err != nil { + return nil, err + } + defer file.Close() + listeners := make(map[uint64]struct{}) + scanner := bufio.NewScanner(file) + if scanner.Scan() { + // Skip the stable kernel table header. + } + for scanner.Scan() { + fields := strings.Fields(scanner.Text()) + if len(fields) < 10 || fields[3] != "0A" { + continue + } + inode, err := strconv.ParseUint(fields[9], 10, 64) + if err != nil { + return nil, fmt.Errorf("parse %s listener inode: %w", path, err) + } + listeners[inode] = struct{}{} + } + if err := scanner.Err(); err != nil { + return nil, fmt.Errorf("read %s: %w", path, err) + } + return listeners, nil +} diff --git a/integration/agentcompat/internal/process/sampler.go b/integration/agentcompat/internal/process/sampler.go new file mode 100644 index 00000000..1c24de43 --- /dev/null +++ b/integration/agentcompat/internal/process/sampler.go @@ -0,0 +1,109 @@ +//go:build linux + +package process + +import ( + "context" + "errors" + "fmt" + "os" + "time" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/contract" +) + +type Sample struct { + PID int `json:"pid"` + RSSBytes uint64 `json:"rss_bytes"` + DescendantPIDs []int `json:"descendant_pids"` + DescendantCount int `json:"descendant_count"` + NonStdioFDCount int `json:"non_stdio_fd_count"` + TCPListenerCount int `json:"tcp_listener_count"` + TCP6ListenerCount int `json:"tcp6_listener_count"` +} + +type Window struct { + PID int `json:"pid"` + Samples []Sample `json:"samples"` +} + +type WindowSpec struct { + PID int + Interval time.Duration + AllowTerminated bool + ObserveSample func(context.Context, Sample) error +} + +func SampleProcess(pid int) (Sample, error) { + rssBytes, err := readRSSBytes(pid) + if err != nil { + return Sample{}, err + } + descendants, err := descendantPIDs(pid) + if err != nil { + return Sample{}, err + } + fdCount, socketInodes, err := processFDs(pid) + if err != nil { + return Sample{}, err + } + tcpListeners, err := listeningSocketInodes(pid, "tcp") + if err != nil { + return Sample{}, err + } + tcp6Listeners, err := listeningSocketInodes(pid, "tcp6") + if err != nil { + return Sample{}, err + } + return Sample{ + PID: pid, + RSSBytes: rssBytes, + DescendantPIDs: descendants, + DescendantCount: len(descendants), + NonStdioFDCount: fdCount, + TCPListenerCount: intersectionCount(socketInodes, tcpListeners), + TCP6ListenerCount: intersectionCount(socketInodes, tcp6Listeners), + }, nil +} + +func SampleWindow(ctx context.Context, spec WindowSpec) (Window, error) { + if spec.PID < 1 || spec.Interval <= 0 { + return Window{}, errors.New("invalid sample window specification") + } + window := Window{PID: spec.PID, Samples: make([]Sample, 0, contract.ResourceSampleCount)} + for index := 0; index < contract.ResourceSampleCount; index++ { + if index > 0 { + timer := time.NewTimer(spec.Interval) + select { + case <-timer.C: + case <-ctx.Done(): + timer.Stop() + return Window{}, ctx.Err() + } + } + sample, err := SampleProcess(spec.PID) + if err != nil { + if spec.AllowTerminated && os.IsNotExist(err) { + return window, nil + } + return Window{}, fmt.Errorf("sample %d of PID %d: %w", index+1, spec.PID, err) + } + window.Samples = append(window.Samples, sample) + if spec.ObserveSample != nil { + if err := spec.ObserveSample(ctx, sample); err != nil { + return window, err + } + } + } + return window, nil +} + +func intersectionCount(left, right map[uint64]struct{}) int { + count := 0 + for value := range left { + if _, exists := right[value]; exists { + count++ + } + } + return count +} diff --git a/integration/agentcompat/internal/process/sampler_test.go b/integration/agentcompat/internal/process/sampler_test.go new file mode 100644 index 00000000..7268636d --- /dev/null +++ b/integration/agentcompat/internal/process/sampler_test.go @@ -0,0 +1,209 @@ +//go:build linux + +package process + +import ( + "context" + "errors" + "net" + "os" + "os/exec" + "reflect" + "testing" + "time" +) + +func TestSampler_ReadsRSS(t *testing.T) { + // Given / When + sample, err := SampleProcess(os.Getpid()) + + // Then + requireNoError(t, err) + if sample.RSSBytes == 0 { + t.Fatal("RSS is zero") + } +} + +func TestSampler_CountsDescendants(t *testing.T) { + // Given + child, closeInput := startBlockingHelper(t) + defer closeInput() + defer reapHelper(child) + + // When + sample, err := SampleProcess(os.Getpid()) + + // Then + requireNoError(t, err) + if !containsPID(sample.DescendantPIDs, child.Process.Pid) { + t.Fatalf("descendants = %v, want PID %d", sample.DescendantPIDs, child.Process.Pid) + } +} + +func TestSampler_CountsNonStdioFDs(t *testing.T) { + // Given + baseline, err := SampleProcess(os.Getpid()) + requireNoError(t, err) + file, err := os.Open("/proc/self/status") + requireNoError(t, err) + t.Cleanup(func() { _ = file.Close() }) + + // When + sample, err := SampleProcess(os.Getpid()) + + // Then + requireNoError(t, err) + if sample.NonStdioFDCount != baseline.NonStdioFDCount+1 { + t.Fatalf("non-stdio FDs = %d, baseline = %d", sample.NonStdioFDCount, baseline.NonStdioFDCount) + } +} + +func TestSampler_CountsTCPListeners(t *testing.T) { + // Given + baseline, err := SampleProcess(os.Getpid()) + requireNoError(t, err) + listener, err := net.Listen("tcp4", "127.0.0.1:0") + requireNoError(t, err) + t.Cleanup(func() { _ = listener.Close() }) + + // When + sample, err := SampleProcess(os.Getpid()) + + // Then + requireNoError(t, err) + if sample.TCPListenerCount != baseline.TCPListenerCount+1 { + t.Fatalf("TCP listeners = %d, baseline = %d", sample.TCPListenerCount, baseline.TCPListenerCount) + } +} + +func TestSampler_CountsTCP6Listeners(t *testing.T) { + // Given + baseline, err := SampleProcess(os.Getpid()) + requireNoError(t, err) + listener, err := net.Listen("tcp6", "[::1]:0") + if err != nil { + t.Skipf("IPv6 loopback listener unavailable: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + + // When + sample, err := SampleProcess(os.Getpid()) + + // Then + requireNoError(t, err) + if sample.TCP6ListenerCount != baseline.TCP6ListenerCount+1 { + t.Fatalf("TCP6 listeners = %d, baseline = %d", sample.TCP6ListenerCount, baseline.TCP6ListenerCount) + } +} + +func TestSampler_CollectsFiveSampleWindow(t *testing.T) { + // Given + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + + // When + window, err := SampleWindow(ctx, WindowSpec{PID: os.Getpid(), Interval: time.Millisecond}) + + // Then + requireNoError(t, err) + if len(window.Samples) != 5 { + t.Fatalf("samples = %d, want 5", len(window.Samples)) + } +} + +func TestSampleWindow_InvokesObserverAfterEachSuccessfulAppend(t *testing.T) { + // Given + observed := make([]Sample, 0, 5) + + // When + window, err := SampleWindow(t.Context(), WindowSpec{PID: os.Getpid(), Interval: time.Millisecond, ObserveSample: func(_ context.Context, sample Sample) error { + observed = append(observed, sample) + return nil + }}) + + // Then + requireNoError(t, err) + if len(window.Samples) != 5 || len(observed) != 5 { + t.Fatalf("samples = %d, observed = %d, want 5", len(window.Samples), len(observed)) + } + for index := range window.Samples { + if !reflect.DeepEqual(window.Samples[index], observed[index]) { + t.Fatalf("sample %d was not observed after append", index+1) + } + } +} + +func TestSampleWindow_ReturnsAppendedSampleWhenObserverFails(t *testing.T) { + // Given + observerErr := errors.New("observer failed") + calls := 0 + + // When + window, err := SampleWindow(t.Context(), WindowSpec{PID: os.Getpid(), Interval: time.Millisecond, ObserveSample: func(context.Context, Sample) error { + calls++ + return observerErr + }}) + + // Then + if !errors.Is(err, observerErr) { + t.Fatalf("observer error = %v, want %v", err, observerErr) + } + if calls != 1 || len(window.Samples) != 1 { + t.Fatalf("calls = %d, samples = %d, want one appended sample", calls, len(window.Samples)) + } +} + +func TestSampler_RejectsVanishedPIDDuringWindow(t *testing.T) { + // Given + child := startCleanHelper(t) + requireNoError(t, child.Wait()) + + // When + _, err := SampleWindow(t.Context(), WindowSpec{PID: child.Process.Pid, Interval: time.Millisecond}) + + // Then + if err == nil { + t.Fatal("vanished PID was accepted") + } +} + +func TestSampler_AllowsExplicitlyTerminatedPID(t *testing.T) { + // Given + child := startCleanHelper(t) + requireNoError(t, child.Wait()) + + // When + window, err := SampleWindow(t.Context(), WindowSpec{PID: child.Process.Pid, Interval: time.Millisecond, AllowTerminated: true}) + + // Then + requireNoError(t, err) + if len(window.Samples) != 0 { + t.Fatalf("samples = %d, want 0", len(window.Samples)) + } +} + +func TestSampler_ToleratesVanishedUnrelatedProcEntries(t *testing.T) { + // Given + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + churnDone := make(chan struct{}) + go func() { + defer close(churnDone) + for ctx.Err() == nil { + command := exec.Command(os.Args[0], "-test.run=^TestProcessHelper$") + command.Env = append(os.Environ(), helperModeEnv+"=clean") + if err := command.Run(); err != nil { + return + } + } + }() + + // When + for range 100 { + if _, err := SampleProcess(os.Getpid()); err != nil { + t.Fatalf("sample during /proc churn: %v", err) + } + } + cancel() + <-churnDone +} diff --git a/integration/agentcompat/internal/process/sqlite_journal_identity.go b/integration/agentcompat/internal/process/sqlite_journal_identity.go new file mode 100644 index 00000000..b95b285f --- /dev/null +++ b/integration/agentcompat/internal/process/sqlite_journal_identity.go @@ -0,0 +1,76 @@ +//go:build linux + +package process + +import ( + "errors" + "sync" + + "golang.org/x/sys/unix" +) + +var ( + ErrSQLiteJournalUnsupported = errors.New("sqlite journal identity unsupported") + ErrSQLiteJournalIdentityMismatch = errors.New("sqlite journal identity mismatch") +) + +type SQLiteJournalUnsupportedError struct{ Missing uint32 } + +func (err *SQLiteJournalUnsupportedError) Error() string { return ErrSQLiteJournalUnsupported.Error() } +func (err *SQLiteJournalUnsupportedError) Unwrap() error { return ErrSQLiteJournalUnsupported } + +type SQLiteJournalIdentity struct { + MountID uint64 + DeviceMajor uint32 + DeviceMinor uint32 + Inode uint64 + BirthTime unix.StatxTimestamp +} + +func (identity SQLiteJournalIdentity) equal(other SQLiteJournalIdentity) bool { + return identity == other +} + +func sqliteJournalIdentity(stat unix.Statx_t) (SQLiteJournalIdentity, error) { + required := uint32(unix.STATX_MNT_ID | unix.STATX_BTIME) + if stat.Mask&required != required { + return SQLiteJournalIdentity{}, &SQLiteJournalUnsupportedError{Missing: required &^ stat.Mask} + } + return SQLiteJournalIdentity{ + MountID: stat.Mnt_id, + DeviceMajor: stat.Dev_major, + DeviceMinor: stat.Dev_minor, + Inode: stat.Ino, + BirthTime: stat.Btime, + }, nil +} + +func readSQLiteJournalIdentity(fd int) (SQLiteJournalIdentity, error) { + var stat unix.Statx_t + err := unix.Statx(fd, "", unix.AT_EMPTY_PATH, unix.STATX_BASIC_STATS|unix.STATX_MNT_ID|unix.STATX_BTIME, &stat) + if err != nil { + if errors.Is(err, unix.ENOSYS) || errors.Is(err, unix.EINVAL) || errors.Is(err, unix.EPERM) { + return SQLiteJournalIdentity{}, &SQLiteJournalUnsupportedError{Missing: unix.STATX_MNT_ID | unix.STATX_BTIME} + } + return SQLiteJournalIdentity{}, err + } + if stat.Mode&unix.S_IFMT != unix.S_IFREG { + return SQLiteJournalIdentity{}, &SQLiteJournalUnsupportedError{} + } + return sqliteJournalIdentity(stat) +} + +type SQLiteJournalIdentityError struct { + Expected SQLiteJournalIdentity + Actual SQLiteJournalIdentity +} + +func (err *SQLiteJournalIdentityError) Error() string { + return ErrSQLiteJournalIdentityMismatch.Error() +} +func (err *SQLiteJournalIdentityError) Unwrap() error { return ErrSQLiteJournalIdentityMismatch } + +type sqliteJournalCloser struct { + once sync.Once + err error +} diff --git a/integration/agentcompat/internal/process/sqlite_journal_watch.go b/integration/agentcompat/internal/process/sqlite_journal_watch.go new file mode 100644 index 00000000..e5fbef52 --- /dev/null +++ b/integration/agentcompat/internal/process/sqlite_journal_watch.go @@ -0,0 +1,195 @@ +//go:build linux + +package process + +import ( + "bytes" + "context" + "errors" + "fmt" + "path/filepath" + "sync" + "unsafe" + + "golang.org/x/sys/unix" +) + +var ErrSQLiteJournalLifecycle = errors.New("invalid sqlite journal lifecycle") + +type SQLiteJournalLifecycleError struct{ Event uint32 } + +func (err *SQLiteJournalLifecycleError) Error() string { return ErrSQLiteJournalLifecycle.Error() } +func (err *SQLiteJournalLifecycleError) Unwrap() error { return ErrSQLiteJournalLifecycle } + +type SQLiteJournalWatch struct { + path string + journalFD int + inotifyFD int + identity SQLiteJournalIdentity + journalWD int + directoryWD int + journalName []byte + closed sqliteJournalCloser + mu sync.Mutex + closeSeen bool + deleted bool +} + +func OpenSQLiteJournalWatch(path string) (*SQLiteJournalWatch, error) { + journalFD, err := unix.Open(path, unix.O_PATH|unix.O_NOFOLLOW|unix.O_CLOEXEC, 0) + if err != nil { + return nil, err + } + identity, err := readSQLiteJournalIdentity(journalFD) + if err != nil { + _ = unix.Close(journalFD) + return nil, err + } + inotifyFD, err := unix.InotifyInit1(unix.IN_CLOEXEC | unix.IN_NONBLOCK) + if err != nil { + _ = unix.Close(journalFD) + return nil, err + } + journalWD, err := unix.InotifyAddWatch(inotifyFD, fmt.Sprintf("/proc/self/fd/%d", journalFD), unix.IN_CLOSE_WRITE|unix.IN_DELETE_SELF|unix.IN_MOVE_SELF|unix.IN_UNMOUNT) + if err != nil { + _ = unix.Close(inotifyFD) + _ = unix.Close(journalFD) + return nil, err + } + directoryWD, err := unix.InotifyAddWatch(inotifyFD, filepath.Dir(path), unix.IN_DELETE|unix.IN_UNMOUNT) + if err != nil { + _ = unix.Close(inotifyFD) + _ = unix.Close(journalFD) + return nil, err + } + watch := &SQLiteJournalWatch{path: path, journalFD: journalFD, inotifyFD: inotifyFD, identity: identity, journalWD: journalWD, directoryWD: directoryWD, journalName: []byte(filepath.Base(path))} + if err := watch.Verify(); err != nil { + _ = watch.Close() + return nil, err + } + return watch, nil +} + +func (watch *SQLiteJournalWatch) Identity() SQLiteJournalIdentity { return watch.identity } + +func (watch *SQLiteJournalWatch) ObserveSample(context.Context, Sample) error { return watch.Verify() } + +func (watch *SQLiteJournalWatch) Verify() error { + fd, err := unix.Open(watch.path, unix.O_PATH|unix.O_NOFOLLOW|unix.O_CLOEXEC, 0) + if err != nil { + return &SQLiteJournalIdentityError{Expected: watch.identity} + } + defer unix.Close(fd) + actual, err := readSQLiteJournalIdentity(fd) + if err != nil { + return err + } + if !watch.identity.equal(actual) { + return &SQLiteJournalIdentityError{Expected: watch.identity, Actual: actual} + } + return nil +} + +func (watch *SQLiteJournalWatch) Wait(ctx context.Context) error { + cancelFD, err := unix.Eventfd(0, unix.EFD_CLOEXEC|unix.EFD_NONBLOCK) + if err != nil { + return err + } + defer unix.Close(cancelFD) + stop := make(chan struct{}) + var done sync.WaitGroup + done.Add(1) + go func() { + defer done.Done() + select { + case <-ctx.Done(): + _, _ = unix.Write(cancelFD, []byte{1, 0, 0, 0, 0, 0, 0, 0}) + case <-stop: + } + }() + defer func() { close(stop); done.Wait() }() + for { + fds := []unix.PollFd{{Fd: int32(watch.inotifyFD), Events: unix.POLLIN}, {Fd: int32(cancelFD), Events: unix.POLLIN}} + if _, err := unix.Poll(fds, -1); err != nil { + if errors.Is(err, unix.EINTR) { + continue + } + return err + } + if fds[1].Revents&unix.POLLIN != 0 { + return ctx.Err() + } + if err := watch.readEvents(); err != nil { + return err + } + watch.mu.Lock() + completed := watch.deleted + watch.mu.Unlock() + if completed { + return nil + } + } +} + +func (watch *SQLiteJournalWatch) readEvents() error { + var buffer [unix.SizeofInotifyEvent * 8]byte + count, err := unix.Read(watch.inotifyFD, buffer[:]) + if errors.Is(err, unix.EAGAIN) || errors.Is(err, unix.EINTR) { + return nil + } + if err != nil { + return err + } + for offset := 0; offset+unix.SizeofInotifyEvent <= count; { + event := (*unix.InotifyEvent)(unsafe.Pointer(&buffer[offset])) + next := offset + unix.SizeofInotifyEvent + int(event.Len) + if next > count { + return &SQLiteJournalLifecycleError{} + } + nameStart := offset + unix.SizeofInotifyEvent + name := bytes.TrimRight(buffer[nameStart:next], "\x00") + if err := watch.observeEvent(event.Wd, event.Mask, name); err != nil { + return err + } + offset = next + } + return nil +} + +func (watch *SQLiteJournalWatch) observeEvent(watchDescriptor int32, mask uint32, name []byte) error { + if int(watchDescriptor) == watch.directoryWD && mask&unix.IN_DELETE != 0 && bytes.Equal(name, watch.journalName) { + return watch.observe(unix.IN_DELETE_SELF) + } + if int(watchDescriptor) != watch.journalWD { + return nil + } + return watch.observe(mask) +} + +func (watch *SQLiteJournalWatch) observe(mask uint32) error { + watch.mu.Lock() + defer watch.mu.Unlock() + if mask&(unix.IN_Q_OVERFLOW|unix.IN_MOVE_SELF|unix.IN_UNMOUNT) != 0 || mask&unix.IN_IGNORED != 0 && !watch.deleted { + return &SQLiteJournalLifecycleError{Event: mask} + } + if mask&unix.IN_CLOSE_WRITE != 0 { + if watch.closeSeen || watch.deleted { + return &SQLiteJournalLifecycleError{Event: mask} + } + watch.closeSeen = true + } + if mask&unix.IN_DELETE_SELF != 0 { + if !watch.closeSeen || watch.deleted { + return &SQLiteJournalLifecycleError{Event: mask} + } + watch.deleted = true + } + return nil +} + +func (watch *SQLiteJournalWatch) Close() error { + watch.closed.once.Do(func() { + watch.closed.err = errors.Join(unix.Close(watch.inotifyFD), unix.Close(watch.journalFD)) + }) + return watch.closed.err +} diff --git a/integration/agentcompat/internal/process/sqlite_journal_watch_test.go b/integration/agentcompat/internal/process/sqlite_journal_watch_test.go new file mode 100644 index 00000000..ffb0c31e --- /dev/null +++ b/integration/agentcompat/internal/process/sqlite_journal_watch_test.go @@ -0,0 +1,186 @@ +//go:build linux + +package process + +import ( + "context" + "errors" + "os" + "path/filepath" + "testing" + "time" + + "golang.org/x/sys/unix" +) + +func TestSQLiteJournalWatch_CapturesExactIdentityAndCloses(t *testing.T) { + // Given + path := writeJournal(t, "dashboard.sqlite-journal") + watch, err := OpenSQLiteJournalWatch(path) + requireNoError(t, err) + t.Cleanup(func() { requireNoError(t, watch.Close()) }) + + // When + identity := watch.Identity() + + // Then + if identity.MountID == 0 || identity.Inode == 0 || identity.BirthTime.Sec == 0 { + t.Fatalf("identity = %#v, want complete statx identity", identity) + } + if err := watch.Verify(); err != nil { + t.Fatalf("verify identity: %v", err) + } + journalFD, inotifyFD := watch.journalFD, watch.inotifyFD + requireNoError(t, watch.Close()) + requireNoError(t, watch.Close()) + if _, err := unix.FcntlInt(uintptr(journalFD), unix.F_GETFD, 0); !errors.Is(err, unix.EBADF) { + t.Fatalf("journal descriptor remains open: %v", err) + } + if _, err := unix.FcntlInt(uintptr(inotifyFD), unix.F_GETFD, 0); !errors.Is(err, unix.EBADF) { + t.Fatalf("inotify descriptor remains open: %v", err) + } +} + +func TestSQLiteJournalWatch_RejectsReplacementPathDrift(t *testing.T) { + // Given + path := writeJournal(t, "dashboard.sqlite-journal") + watch, err := OpenSQLiteJournalWatch(path) + requireNoError(t, err) + t.Cleanup(func() { requireNoError(t, watch.Close()) }) + replacement := filepath.Join(filepath.Dir(path), "replacement") + requireNoError(t, os.WriteFile(replacement, []byte("replacement"), 0o600)) + requireNoError(t, os.Rename(replacement, path)) + + // When + err = watch.Verify() + + // Then + if !errors.Is(err, ErrSQLiteJournalIdentityMismatch) { + t.Fatalf("verify error = %v, want identity mismatch", err) + } +} + +func TestSQLiteJournalWatch_VerifiesIdentityForEveryWindowSample(t *testing.T) { + // Given + path := writeJournal(t, "dashboard.sqlite-journal") + watch, err := OpenSQLiteJournalWatch(path) + requireNoError(t, err) + t.Cleanup(func() { requireNoError(t, watch.Close()) }) + verified := 0 + + // When + window, err := SampleWindow(t.Context(), WindowSpec{PID: os.Getpid(), Interval: time.Millisecond, ObserveSample: func(ctx context.Context, sample Sample) error { + verified++ + return watch.ObserveSample(ctx, sample) + }}) + + // Then + requireNoError(t, err) + if len(window.Samples) != 5 || verified != 5 { + t.Fatalf("samples = %d, verified = %d, want 5", len(window.Samples), verified) + } +} + +func TestSQLiteJournalIdentity_RejectsMissingRequiredStatxMask(t *testing.T) { + // Given + stat := unix.Statx_t{Mask: unix.STATX_MNT_ID, Mnt_id: 1, Ino: 2} + + // When / Then + if _, err := sqliteJournalIdentity(stat); !errors.Is(err, ErrSQLiteJournalUnsupported) { + t.Fatalf("birth-time error = %v, want unsupported", err) + } + stat.Mask = unix.STATX_BTIME + if _, err := sqliteJournalIdentity(stat); !errors.Is(err, ErrSQLiteJournalUnsupported) { + t.Fatalf("mount-ID error = %v, want unsupported", err) + } +} + +func TestSQLiteJournalWatch_RejectsInvalidLifecycleEvents(t *testing.T) { + for name, mask := range map[string]uint32{ + "overflow": unix.IN_Q_OVERFLOW, + "move self": unix.IN_MOVE_SELF, + "unmount": unix.IN_UNMOUNT, + "ignored": unix.IN_IGNORED, + "missing close": unix.IN_DELETE_SELF, + } { + t.Run(name, func(t *testing.T) { + // Given + watch := &SQLiteJournalWatch{} + // When + err := watch.observe(mask) + + // Then + if err == nil { + t.Fatal("invalid lifecycle event was accepted") + } + }) + } +} + +func TestSQLiteJournalWatch_RejectsDuplicateTerminalEvent(t *testing.T) { + // Given + watch := &SQLiteJournalWatch{} + requireNoError(t, watch.observe(unix.IN_CLOSE_WRITE)) + requireNoError(t, watch.observe(unix.IN_DELETE_SELF)) + + // When + err := watch.observe(unix.IN_DELETE_SELF) + + // Then + if !errors.Is(err, ErrSQLiteJournalLifecycle) { + t.Fatalf("duplicate terminal error = %v, want lifecycle error", err) + } +} + +func TestSQLiteJournalWatch_WaitsForCloseThenDeleteAndCancellation(t *testing.T) { + // Given + path := writeJournal(t, "dashboard.sqlite-journal") + watch, err := OpenSQLiteJournalWatch(path) + requireNoError(t, err) + t.Cleanup(func() { requireNoError(t, watch.Close()) }) + ctx, cancel := context.WithCancel(t.Context()) + result := make(chan error, 1) + go func() { result <- watch.Wait(ctx) }() + + // When + cancel() + err = <-result + + // Then + if !errors.Is(err, context.Canceled) { + t.Fatalf("wait error = %v, want cancellation", err) + } + if err := watch.observe(unix.IN_CLOSE_WRITE); err != nil { + t.Fatalf("close write: %v", err) + } + if err := watch.observe(unix.IN_DELETE_SELF); err != nil { + t.Fatalf("delete self: %v", err) + } +} + +func TestSQLiteJournalWatch_WaitsForExactCloseDeleteLifecycle(t *testing.T) { + // Given + path := writeJournal(t, "dashboard.sqlite-journal") + watch, err := OpenSQLiteJournalWatch(path) + requireNoError(t, err) + t.Cleanup(func() { requireNoError(t, watch.Close()) }) + result := make(chan error, 1) + go func() { result <- watch.Wait(t.Context()) }() + journal, err := os.OpenFile(path, os.O_WRONLY|os.O_APPEND, 0) + requireNoError(t, err) + requireNoError(t, journal.Close()) + requireNoError(t, os.Remove(path)) + + // When + err = <-result + + // Then + requireNoError(t, err) +} + +func writeJournal(t *testing.T, name string) string { + t.Helper() + path := filepath.Join(t.TempDir(), name) + requireNoError(t, os.WriteFile(path, []byte("journal"), 0o600)) + return path +} diff --git a/integration/agentcompat/internal/process/supervisor.go b/integration/agentcompat/internal/process/supervisor.go new file mode 100644 index 00000000..edef60f0 --- /dev/null +++ b/integration/agentcompat/internal/process/supervisor.go @@ -0,0 +1,278 @@ +//go:build linux + +package process + +import ( + "context" + "errors" + "fmt" + "io" + "os" + "os/exec" + "sync" + "syscall" + "time" +) + +type Stream string + +const ( + Stdout Stream = "stdout" + Stderr Stream = "stderr" +) + +type Spec struct { + Name string + Path string + Args []string + Dir string + Env []string + ExtraFiles []*os.File + Stdout io.Writer + Stderr io.Writer + MaxLogBytes int + TerminateTimeout time.Duration + KillTimeout time.Duration + Readiness func(Stream, string) bool + Credential *syscall.Credential +} + +type Supervisor struct { + ctx context.Context + spec Spec + cmd *exec.Cmd + pid int + pgid int + ready chan struct{} + readyOnce sync.Once + exited chan struct{} + waitErr error + waitMu sync.Mutex + cleanupOnce sync.Once + cleanupDone chan struct{} + cleanupErr error + forced bool + stateMu sync.Mutex + stdoutLog *boundedLog + stderrLog *boundedLog +} + +func NewSupervisor(ctx context.Context, spec Spec) *Supervisor { + return &Supervisor{ctx: ctx, spec: spec, ready: make(chan struct{}), exited: make(chan struct{}), cleanupDone: make(chan struct{})} +} + +func (supervisor *Supervisor) Start() error { + if supervisor.spec.Name == "" || supervisor.spec.Path == "" || supervisor.spec.MaxLogBytes < 1 || supervisor.spec.TerminateTimeout <= 0 || supervisor.spec.KillTimeout <= 0 { + return errors.New("invalid process specification") + } + command := exec.Command(supervisor.spec.Path, supervisor.spec.Args...) + command.Dir = supervisor.spec.Dir + command.Env = supervisor.spec.Env + if command.Env == nil { + command.Env = os.Environ() + } + command.ExtraFiles = supervisor.spec.ExtraFiles + command.SysProcAttr = &syscall.SysProcAttr{Setpgid: true, Pdeathsig: syscall.SIGKILL, Credential: supervisor.spec.Credential} + supervisor.stdoutLog = newBoundedLog(supervisor.spec.Stdout, supervisor.spec.MaxLogBytes, supervisor.lineObserver(Stdout)) + supervisor.stderrLog = newBoundedLog(supervisor.spec.Stderr, supervisor.spec.MaxLogBytes, supervisor.lineObserver(Stderr)) + command.Stdout = supervisor.stdoutLog + command.Stderr = supervisor.stderrLog + if err := command.Start(); err != nil { + supervisor.closeExtraFiles() + return fmt.Errorf("start %s: %w", supervisor.spec.Name, err) + } + supervisor.closeExtraFiles() + supervisor.cmd = command + supervisor.pid = command.Process.Pid + supervisor.pgid = command.Process.Pid + go supervisor.reap() + go supervisor.watchContext() + return nil +} + +func (supervisor *Supervisor) lineObserver(stream Stream) func(string) { + return func(line string) { + if supervisor.spec.Readiness != nil && supervisor.spec.Readiness(stream, line) { + supervisor.SignalReady() + } + } +} + +func (supervisor *Supervisor) closeExtraFiles() { + for _, file := range supervisor.spec.ExtraFiles { + if file != nil { + _ = file.Close() + } + } +} + +func (supervisor *Supervisor) reap() { + err := supervisor.cmd.Wait() + supervisor.stdoutLog.Close() + supervisor.stderrLog.Close() + supervisor.waitMu.Lock() + supervisor.waitErr = err + supervisor.waitMu.Unlock() + close(supervisor.exited) +} + +func (supervisor *Supervisor) watchContext() { + select { + case <-supervisor.ctx.Done(): + _ = supervisor.Stop(context.WithoutCancel(supervisor.ctx)) + case <-supervisor.exited: + } +} + +func (supervisor *Supervisor) SignalReady() { + supervisor.readyOnce.Do(func() { close(supervisor.ready) }) +} + +func (supervisor *Supervisor) Ready() <-chan struct{} { return supervisor.ready } + +func (supervisor *Supervisor) Exited() <-chan struct{} { return supervisor.exited } + +func (supervisor *Supervisor) WaitReady(ctx context.Context) error { + select { + case <-supervisor.ready: + return nil + default: + } + select { + case <-supervisor.ready: + return nil + case <-supervisor.exited: + select { + case <-supervisor.ready: + return nil + default: + return errors.New("process exited before readiness") + } + case <-ctx.Done(): + return ctx.Err() + } +} + +func (supervisor *Supervisor) Wait(ctx context.Context) error { + select { + case <-supervisor.exited: + cleanupErr := supervisor.Stop(ctx) + supervisor.waitMu.Lock() + waitErr := supervisor.waitErr + supervisor.waitMu.Unlock() + return errors.Join(waitErr, cleanupErr) + case <-ctx.Done(): + return errors.Join(ctx.Err(), supervisor.Stop(context.WithoutCancel(ctx))) + } +} + +func (supervisor *Supervisor) Stop(ctx context.Context) error { + supervisor.cleanupOnce.Do(func() { go supervisor.cleanup() }) + select { + case <-supervisor.cleanupDone: + return supervisor.cleanupResult() + case <-ctx.Done(): + return ctx.Err() + } +} + +func (supervisor *Supervisor) cleanup() { + defer close(supervisor.cleanupDone) + if supervisor.pgid < 1 { + return + } + if !processGroupExists(supervisor.pgid) { + supervisor.waitForExit() + return + } + if err := syscall.Kill(-supervisor.pgid, syscall.SIGTERM); err != nil && !errors.Is(err, syscall.ESRCH) { + supervisor.setCleanupError(fmt.Errorf("terminate %s process group: %w", supervisor.spec.Name, err)) + return + } + if waitProcessGroup(supervisor.pgid, supervisor.spec.TerminateTimeout) { + supervisor.waitForExit() + return + } + supervisor.stateMu.Lock() + supervisor.forced = true + supervisor.stateMu.Unlock() + if err := syscall.Kill(-supervisor.pgid, syscall.SIGKILL); err != nil && !errors.Is(err, syscall.ESRCH) { + supervisor.setCleanupError(fmt.Errorf("kill %s process group: %w", supervisor.spec.Name, err)) + return + } + if !waitProcessGroup(supervisor.pgid, supervisor.spec.KillTimeout) { + supervisor.setCleanupError(fmt.Errorf("%s process group %d survived SIGKILL", supervisor.spec.Name, supervisor.pgid)) + return + } + supervisor.waitForExit() +} + +func (supervisor *Supervisor) waitForExit() { + timer := time.NewTimer(supervisor.spec.KillTimeout) + defer timer.Stop() + select { + case <-supervisor.exited: + case <-timer.C: + supervisor.setCleanupError(fmt.Errorf("%s process was not reaped", supervisor.spec.Name)) + } +} + +func (supervisor *Supervisor) setCleanupError(err error) { + supervisor.stateMu.Lock() + supervisor.cleanupErr = errors.Join(supervisor.cleanupErr, err) + supervisor.stateMu.Unlock() +} + +func (supervisor *Supervisor) cleanupResult() error { + supervisor.stateMu.Lock() + defer supervisor.stateMu.Unlock() + return supervisor.cleanupErr +} + +func processGroupExists(pgid int) bool { + err := syscall.Kill(-pgid, 0) + return err == nil || errors.Is(err, syscall.EPERM) +} + +func waitProcessGroup(pgid int, timeout time.Duration) bool { + if !processGroupExists(pgid) { + return true + } + ticker := time.NewTicker(time.Millisecond) + defer ticker.Stop() + timer := time.NewTimer(timeout) + defer timer.Stop() + for { + select { + case <-ticker.C: + if !processGroupExists(pgid) { + return true + } + case <-timer.C: + return !processGroupExists(pgid) + } + } +} + +func (supervisor *Supervisor) PID() int { return supervisor.pid } + +func (supervisor *Supervisor) ProcessGroupID() int { return supervisor.pgid } + +func (supervisor *Supervisor) ForcedCleanup() bool { + supervisor.stateMu.Lock() + defer supervisor.stateMu.Unlock() + return supervisor.forced +} + +func (supervisor *Supervisor) CleanupRecord() CleanupRecord { + supervisor.stateMu.Lock() + defer supervisor.stateMu.Unlock() + return CleanupRecord{Name: supervisor.spec.Name, PID: supervisor.pid, Forced: supervisor.forced, Error: errorString(supervisor.cleanupErr)} +} + +func errorString(err error) string { + if err == nil { + return "" + } + return err.Error() +} diff --git a/integration/agentcompat/internal/process/supervisor_agentcompat.go b/integration/agentcompat/internal/process/supervisor_agentcompat.go new file mode 100644 index 00000000..340cc3fd --- /dev/null +++ b/integration/agentcompat/internal/process/supervisor_agentcompat.go @@ -0,0 +1,7 @@ +//go:build linux && agentcompat + +package process + +func (supervisor *Supervisor) CleanupDoneForTest() <-chan struct{} { + return supervisor.cleanupDone +} diff --git a/integration/agentcompat/internal/process/supervisor_test.go b/integration/agentcompat/internal/process/supervisor_test.go new file mode 100644 index 00000000..fa4b3a2a --- /dev/null +++ b/integration/agentcompat/internal/process/supervisor_test.go @@ -0,0 +1,195 @@ +//go:build linux + +package process + +import ( + "bufio" + "bytes" + "context" + "errors" + "net" + "os" + "os/exec" + "path/filepath" + "strings" + "syscall" + "testing" + "time" +) + +func TestSupervisor_CleanExit(t *testing.T) { + // Given + supervisor := newHelperSupervisor(t.Context(), "clean", nil) + + // When + requireNoError(t, supervisor.Start()) + requireNoError(t, supervisor.WaitReady(t.Context())) + + // Then + requireNoError(t, supervisor.Wait(t.Context())) +} + +func TestSupervisor_RunsChildWithConfiguredCredential(t *testing.T) { + // Given + credentialDirectory, err := os.MkdirTemp("/tmp", "agentcompat-credential-") + requireNoError(t, err) + t.Cleanup(func() { _ = os.RemoveAll(credentialDirectory) }) + requireNoError(t, os.Chmod(credentialDirectory, 0o777)) + marker := filepath.Join(credentialDirectory, "credential.txt") + supervisor := newHelperSupervisor(t.Context(), "credential", []string{helperMarkerEnv + "=" + marker}) + testBinary, err := os.ReadFile(os.Args[0]) + requireNoError(t, err) + executablePath := filepath.Join(credentialDirectory, "process-helper") + requireNoError(t, os.WriteFile(executablePath, testBinary, 0o755)) + supervisor.spec.Path = executablePath + supervisor.spec.Credential = &syscall.Credential{Uid: 65534, Gid: 65534} + + // When + requireNoError(t, supervisor.Start()) + requireNoError(t, supervisor.Wait(t.Context())) + + // Then + content, err := os.ReadFile(marker) + requireNoError(t, err) + if strings.TrimSpace(string(content)) != "65534:65534" { + t.Fatalf("child credential = %q, want 65534:65534", content) + } +} + +func TestSupervisor_KillsProcessTree(t *testing.T) { + // Given + marker := filepath.Join(t.TempDir(), "grandchild.pid") + ctx, cancel := context.WithCancel(t.Context()) + supervisor := newHelperSupervisor(ctx, "tree", []string{helperMarkerEnv + "=" + marker}) + requireNoError(t, supervisor.Start()) + requireNoError(t, supervisor.WaitReady(t.Context())) + grandchildPID := readPID(t, marker) + processGroupID := supervisor.ProcessGroupID() + + // When + cancel() + select { + case <-supervisor.cleanupDone: + case <-time.After(2 * time.Second): + t.Fatal("context cancellation did not complete process-tree cleanup") + } + + // Then + requirePIDGone(t, supervisor.PID()) + requirePIDGone(t, grandchildPID) + if err := syscall.Kill(-processGroupID, 0); !errors.Is(err, syscall.ESRCH) { + t.Fatalf("process group %d remains: %v", processGroupID, err) + } +} + +func TestSupervisor_AdoptsListener(t *testing.T) { + // Given + listener, err := net.Listen("tcp4", "127.0.0.1:0") + requireNoError(t, err) + tcpListener := listener.(*net.TCPListener) + inheritedFile, err := tcpListener.File() + requireNoError(t, err) + requireNoError(t, tcpListener.Close()) + supervisor := newHelperSupervisor(t.Context(), "listener", []string{helperFDEnv + "=3"}) + supervisor.spec.ExtraFiles = []*os.File{inheritedFile} + + // When + requireNoError(t, supervisor.Start()) + requireNoError(t, supervisor.WaitReady(t.Context())) + sample, err := SampleProcess(supervisor.PID()) + requireNoError(t, err) + + // Then + if sample.TCPListenerCount != 1 { + t.Fatalf("TCP listeners = %d, want 1", sample.TCPListenerCount) + } + requireNoError(t, supervisor.Stop(t.Context())) + requirePIDGone(t, supervisor.PID()) +} + +func TestSupervisor_RedactsLogs(t *testing.T) { + // Given + var output bytes.Buffer + supervisor := newHelperSupervisor(t.Context(), "logs", nil) + supervisor.spec.MaxLogBytes = 128 + supervisor.spec.Stdout = &output + supervisor.spec.Stderr = &output + + // When + requireNoError(t, supervisor.Start()) + requireNoError(t, supervisor.WaitReady(t.Context())) + requireNoError(t, supervisor.Wait(t.Context())) + + // Then + logged := output.String() + if strings.Contains(logged, "top-secret") || strings.Contains(logged, "eyJsecret") { + t.Fatalf("secret survived supervisor log redaction: %s", logged) + } + if output.Len() > supervisor.spec.MaxLogBytes*2 { + t.Fatalf("combined log bytes = %d, per-stream limit = %d", output.Len(), supervisor.spec.MaxLogBytes) + } + if !strings.Contains(logged, truncationMarker) { + t.Fatalf("truncation marker missing: %q", logged) + } +} + +func TestSupervisor_RecordsForcedCleanupForSIGTERMIgnoringChild(t *testing.T) { + // Given + resultsDir := t.TempDir() + marker := filepath.Join(t.TempDir(), "forced-grandchild.pid") + supervisor := newHelperSupervisor(t.Context(), "force-tree", []string{helperMarkerEnv + "=" + marker}) + requireNoError(t, supervisor.Start()) + requireNoError(t, supervisor.WaitReady(t.Context())) + grandchildPID := readPID(t, marker) + + // When + requireNoError(t, supervisor.Stop(t.Context())) + receipt := NewCleanupReceipt([]CleanupRecord{supervisor.CleanupRecord()}) + receiptPath := filepath.Join(resultsDir, "cleanup.json") + requireNoError(t, WriteCleanupReceipt(receiptPath, receipt)) + + // Then + if !supervisor.ForcedCleanup() { + t.Fatal("forced cleanup was not recorded") + } + data, err := os.ReadFile(receiptPath) + requireNoError(t, err) + if !strings.Contains(string(data), `"forced": true`) { + t.Fatalf("cleanup receipt = %s", data) + } + requirePIDGone(t, supervisor.PID()) + requirePIDGone(t, grandchildPID) +} + +func TestSupervisor_InterruptSignalCleansProcessTree(t *testing.T) { + // Given + marker := filepath.Join(t.TempDir(), "interrupt-grandchild.pid") + command := exec.Command(os.Args[0], "-test.run=^TestProcessHelper$") + command.Env = append(os.Environ(), helperModeEnv+"=interrupt-probe", helperMarkerEnv+"="+marker) + output, err := command.StdoutPipe() + requireNoError(t, err) + command.Stderr = os.Stderr + requireNoError(t, command.Start()) + scanner := bufio.NewScanner(output) + ready := false + for scanner.Scan() { + if scanner.Text() == "PROBE_READY" { + ready = true + break + } + } + requireNoError(t, scanner.Err()) + if !ready { + t.Fatal("interrupt probe exited before readiness") + } + leaderPID := readPID(t, marker+".leader") + grandchildPID := readPID(t, marker) + + // When + requireNoError(t, command.Process.Signal(syscall.SIGTERM)) + requireNoError(t, command.Wait()) + + // Then + requirePIDGone(t, leaderPID) + requirePIDGone(t, grandchildPID) +} diff --git a/integration/agentcompat/internal/workspace/build.go b/integration/agentcompat/internal/workspace/build.go new file mode 100644 index 00000000..fe0b05d7 --- /dev/null +++ b/integration/agentcompat/internal/workspace/build.go @@ -0,0 +1,56 @@ +//go:build linux + +package workspace + +import ( + "context" + "errors" + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" +) + +type BuildSpec struct { + Name string + SourceDir string + Package string + Tags []string + Ldflags []string + Env []string +} + +func (workspace *Workspace) Build(ctx context.Context, spec BuildSpec) (string, error) { + workspace.closeMu.Lock() + defer workspace.closeMu.Unlock() + if err := workspace.requireOpen(); err != nil { + return "", err + } + if spec.SourceDir == "" || spec.Package == "" { + return "", errors.New("build source and package are required") + } + if err := validateLeafName(spec.Name); err != nil { + return "", err + } + binaryPath := filepath.Join(workspace.binDir, spec.Name) + arguments := []string{"build", "-mod=readonly", "-o", binaryPath} + if len(spec.Tags) > 0 { + arguments = append(arguments, "-tags", strings.Join(spec.Tags, ",")) + } + if len(spec.Ldflags) > 0 { + arguments = append(arguments, "-ldflags", strings.Join(spec.Ldflags, " ")) + } + arguments = append(arguments, spec.Package) + command := exec.CommandContext(ctx, "go", arguments...) + command.Dir = spec.SourceDir + command.Env = spec.Env + if spec.Env == nil { + command.Env = os.Environ() + } + output, err := command.CombinedOutput() + if err != nil { + return "", fmt.Errorf("build %s: %w: %s", spec.Name, err, output) + } + return binaryPath, nil +} diff --git a/integration/agentcompat/internal/workspace/listener.go b/integration/agentcompat/internal/workspace/listener.go new file mode 100644 index 00000000..fd5d755c --- /dev/null +++ b/integration/agentcompat/internal/workspace/listener.go @@ -0,0 +1,90 @@ +//go:build linux + +package workspace + +import ( + "bufio" + "errors" + "fmt" + "os" + "strconv" + "strings" + "sync" + "syscall" +) + +type OwnedListener struct { + file *os.File + address string + inode uint64 + closeOnce sync.Once + closeErr error +} + +type ListenerIdentity struct { + Address string + Inode uint64 +} + +func (listener *OwnedListener) FileDescriptor() int { return int(listener.file.Fd()) } + +func (listener *OwnedListener) Address() string { return listener.address } + +func (listener *OwnedListener) Identity() ListenerIdentity { + return ListenerIdentity{Address: listener.address, Inode: listener.inode} +} + +func (listener *OwnedListener) ExtraFile() (*os.File, error) { + descriptor, err := syscall.Dup(listener.FileDescriptor()) + if err != nil { + return nil, fmt.Errorf("duplicate inherited listener FD: %w", err) + } + return os.NewFile(uintptr(descriptor), "agentcompat-listener"), nil +} + +func (listener *OwnedListener) Close() error { + listener.closeOnce.Do(func() { + if err := listener.file.Close(); err != nil && !errors.Is(err, os.ErrClosed) { + listener.closeErr = fmt.Errorf("close owned listener: %w", err) + } + }) + return listener.closeErr +} + +func socketInode(file *os.File) (uint64, error) { + target, err := os.Readlink("/proc/self/fd/" + strconv.Itoa(int(file.Fd()))) + if err != nil { + return 0, fmt.Errorf("read listener FD link: %w", err) + } + if !strings.HasPrefix(target, "socket:[") || !strings.HasSuffix(target, "]") { + return 0, errors.New("listener FD is not a socket") + } + inode, err := strconv.ParseUint(strings.TrimSuffix(strings.TrimPrefix(target, "socket:["), "]"), 10, 64) + if err != nil { + return 0, fmt.Errorf("parse listener inode: %w", err) + } + return inode, nil +} + +func listenerInodePresent(inode uint64) (bool, error) { + for _, path := range []string{"/proc/self/net/tcp", "/proc/self/net/tcp6"} { + file, err := os.Open(path) + if err != nil { + return false, fmt.Errorf("open listener table %s: %w", path, err) + } + scanner := bufio.NewScanner(file) + for scanner.Scan() { + fields := strings.Fields(scanner.Text()) + if len(fields) >= 10 && fields[3] == "0A" && fields[9] == strconv.FormatUint(inode, 10) { + _ = file.Close() + return true, nil + } + } + scanErr := scanner.Err() + closeErr := file.Close() + if scanErr != nil || closeErr != nil { + return false, errors.Join(scanErr, closeErr) + } + } + return false, nil +} diff --git a/integration/agentcompat/internal/workspace/log.go b/integration/agentcompat/internal/workspace/log.go new file mode 100644 index 00000000..a4fb9684 --- /dev/null +++ b/integration/agentcompat/internal/workspace/log.go @@ -0,0 +1,117 @@ +//go:build linux + +package workspace + +import ( + "bytes" + "errors" + "fmt" + "os" + "sync" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/evidence" +) + +const workspaceTruncationMarker = "[TRUNCATED]\n" + +type LogFile struct { + file *os.File + maxBytes int + written int + pending []byte + dropLine bool + closed bool + closeOnce sync.Once + closeErr error + mu sync.Mutex +} + +func newLogFile(path string, maxBytes int) (*LogFile, error) { + file, err := os.OpenFile(path, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600) + if err != nil { + return nil, fmt.Errorf("create workspace log: %w", err) + } + return &LogFile{file: file, maxBytes: maxBytes}, nil +} + +func (logFile *LogFile) Name() string { return logFile.file.Name() } + +func (logFile *LogFile) WriteString(value string) (int, error) { + return logFile.Write([]byte(value)) +} + +func (logFile *LogFile) Write(data []byte) (int, error) { + logFile.mu.Lock() + defer logFile.mu.Unlock() + if logFile.closed { + return 0, errors.New("write closed workspace log") + } + inputLength := len(data) + for len(data) > 0 { + newline := bytes.IndexByte(data, '\n') + if newline < 0 { + logFile.appendFragment(data) + break + } + logFile.appendFragment(data[:newline+1]) + if err := logFile.flushLine(); err != nil { + return 0, err + } + data = data[newline+1:] + } + return inputLength, nil +} + +func (logFile *LogFile) appendFragment(fragment []byte) { + if logFile.dropLine { + return + } + if len(logFile.pending)+len(fragment) > logFile.maxBytes { + logFile.pending = nil + logFile.dropLine = true + return + } + logFile.pending = append(logFile.pending, fragment...) +} + +func (logFile *LogFile) flushLine() error { + if logFile.dropLine { + logFile.dropLine = false + return logFile.writeBounded(workspaceTruncationMarker) + } + redacted := evidence.Redact(string(logFile.pending)) + logFile.pending = nil + if len(redacted) > logFile.maxBytes-logFile.written { + return logFile.writeBounded(workspaceTruncationMarker) + } + return logFile.writeBounded(redacted) +} + +func (logFile *LogFile) writeBounded(value string) error { + remaining := logFile.maxBytes - logFile.written + if remaining <= 0 || value == "" { + return nil + } + if len(value) > remaining { + value = value[:remaining] + } + written, err := logFile.file.WriteString(value) + logFile.written += written + if err != nil { + return fmt.Errorf("write workspace log: %w", err) + } + return nil +} + +func (logFile *LogFile) Close() error { + logFile.closeOnce.Do(func() { + logFile.mu.Lock() + defer logFile.mu.Unlock() + if len(logFile.pending) > 0 || logFile.dropLine { + logFile.closeErr = logFile.flushLine() + } + logFile.closed = true + logFile.closeErr = errors.Join(logFile.closeErr, logFile.file.Close()) + }) + return logFile.closeErr +} diff --git a/integration/agentcompat/internal/workspace/residue_test.go b/integration/agentcompat/internal/workspace/residue_test.go new file mode 100644 index 00000000..32bbbab7 --- /dev/null +++ b/integration/agentcompat/internal/workspace/residue_test.go @@ -0,0 +1,72 @@ +//go:build linux + +package workspace + +import ( + "context" + "errors" + "io" + "os" + "os/exec" + "strings" + "syscall" + "testing" +) + +func TestWorkspace_PreservesEvidenceWhenProcessGroupRemains(t *testing.T) { + // Given + workspace, err := New(context.Background()) + if err != nil { + t.Fatal(err) + } + root := workspace.Root() + command, input := startWorkspaceHelper(t) + if err := workspace.TrackProcessGroup(command.Process.Pid); err != nil { + t.Fatal(err) + } + + // When + err = workspace.Close() + + // Then + if err == nil || !strings.Contains(err.Error(), "process group") { + t.Fatalf("close error = %v", err) + } + if _, statErr := os.Stat(root); statErr != nil { + t.Fatalf("workspace evidence was removed: %v", statErr) + } + if err := input.Close(); err != nil { + t.Fatal(err) + } + if err := command.Wait(); err != nil { + t.Fatal(err) + } + if err := workspace.Close(); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(root); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("workspace remains after group exit: %v", err) + } +} + +func TestWorkspaceHelper(t *testing.T) { + if os.Getenv("GO_WANT_WORKSPACE_HELPER") != "1" { + return + } + _, _ = io.Copy(io.Discard, os.Stdin) +} + +func startWorkspaceHelper(t *testing.T) (*exec.Cmd, io.WriteCloser) { + t.Helper() + command := exec.Command(os.Args[0], "-test.run=^TestWorkspaceHelper$") + command.Env = append(os.Environ(), "GO_WANT_WORKSPACE_HELPER=1") + command.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} + input, err := command.StdinPipe() + if err != nil { + t.Fatal(err) + } + if err := command.Start(); err != nil { + t.Fatal(err) + } + return command, input +} diff --git a/integration/agentcompat/internal/workspace/supervisor_exit_integration_test.go b/integration/agentcompat/internal/workspace/supervisor_exit_integration_test.go new file mode 100644 index 00000000..6a88c4a4 --- /dev/null +++ b/integration/agentcompat/internal/workspace/supervisor_exit_integration_test.go @@ -0,0 +1,250 @@ +//go:build linux && agentcompat + +package workspace + +import ( + "bufio" + "context" + "errors" + "fmt" + "net" + "os" + "os/exec" + "os/signal" + "path/filepath" + "strconv" + "strings" + "syscall" + "testing" + "time" + + processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process" +) + +const ( + workspaceExitHelperModeEnv = "NEZHA_AGENTCOMPAT_WORKSPACE_EXIT_HELPER" + workspaceExitHelperMarkerEnv = "NEZHA_AGENTCOMPAT_WORKSPACE_EXIT_MARKER" +) + +func TestWorkspace_SupervisorExitedPrecedesProcessGroupAndListenerCleanup(t *testing.T) { + // Given + workspace, err := New(context.Background()) + if err != nil { + t.Fatal(err) + } + root := workspace.Root() + listener, err := net.Listen("tcp4", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + ownedListener, err := workspace.AdoptListener(listener) + if err != nil { + t.Fatal(err) + } + inheritedFile, err := ownedListener.ExtraFile() + if err != nil { + t.Fatal(err) + } + descendantMarker := filepath.Join(t.TempDir(), "descendant.pid") + supervisor := processharness.NewSupervisor(t.Context(), processharness.Spec{ + Name: "workspace-exited-semantics", + Path: os.Args[0], + Args: []string{"-test.run=^TestWorkspaceExitHelper$"}, + Env: append(os.Environ(), workspaceExitHelperModeEnv+"=leader", workspaceExitHelperMarkerEnv+"="+descendantMarker), + ExtraFiles: []*os.File{inheritedFile}, + Stdout: os.Stdout, + Stderr: os.Stderr, + MaxLogBytes: 1024, + TerminateTimeout: time.Second, + KillTimeout: time.Second, + }) + if err := supervisor.Start(); err != nil { + t.Fatal(err) + } + if err := workspace.TrackPID(supervisor.PID()); err != nil { + t.Fatal(err) + } + if err := workspace.TrackProcessGroup(supervisor.ProcessGroupID()); err != nil { + t.Fatal(err) + } + if err := ownedListener.Close(); err != nil { + t.Fatal(err) + } + processGroupID := supervisor.ProcessGroupID() + + // When + select { + case <-supervisor.Exited(): + case <-time.After(2 * time.Second): + t.Fatal("leader did not exit") + } + descendantPID := readWorkspaceExitHelperPID(t, descendantMarker) + + // Then + select { + case <-supervisor.CleanupDoneForTest(): + t.Fatal("cleanup completed before Stop") + default: + } + if _, err := os.Stat(filepath.Join("/proc", strconv.Itoa(descendantPID))); err != nil { + t.Fatalf("descendant exited before cleanup: %v", err) + } + if descendantProcessGroupID := readProcessGroupID(t, descendantPID); descendantProcessGroupID != processGroupID { + t.Fatalf("descendant process group = %d, want %d", descendantProcessGroupID, processGroupID) + } + if err := syscall.Kill(-processGroupID, 0); err != nil { + t.Fatalf("process group exited before cleanup: %v", err) + } + requireProcessHoldsSocket(t, descendantPID, ownedListener.inode) + listenerPresent, err := listenerInodePresent(ownedListener.inode) + if err != nil { + t.Fatal(err) + } + if !listenerPresent { + t.Fatal("descendant listener disappeared before cleanup") + } + if err := workspace.Close(); err == nil { + t.Fatal("workspace closed before descendant cleanup") + } + if _, err := os.Stat(root); err != nil { + t.Fatalf("workspace disappeared before cleanup: %v", err) + } + + stopContext, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + if err := supervisor.Stop(stopContext); err != nil { + t.Fatal(err) + } + select { + case <-supervisor.CleanupDoneForTest(): + case <-time.After(time.Second): + t.Fatal("cleanup completion signal did not close after Stop") + } + if err := workspace.Close(); err != nil { + t.Fatal(err) + } + if _, err := os.Stat(filepath.Join("/proc", strconv.Itoa(descendantPID))); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("descendant remains after Stop: %v", err) + } + if err := syscall.Kill(-processGroupID, 0); !errors.Is(err, syscall.ESRCH) { + t.Fatalf("process group remains after Stop: %v", err) + } + listenerPresent, err = listenerInodePresent(ownedListener.inode) + if err != nil { + t.Fatal(err) + } + if listenerPresent { + t.Fatal("listener remains after Stop") + } + if _, err := os.Stat(root); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("workspace remains after cleanup: %v", err) + } +} + +func TestWorkspaceExitHelper(t *testing.T) { + switch os.Getenv(workspaceExitHelperModeEnv) { + case "": + return + case "leader": + runWorkspaceExitLeader(t) + case "descendant": + runWorkspaceExitDescendant(t) + default: + t.Fatalf("unknown workspace exit helper mode %q", os.Getenv(workspaceExitHelperModeEnv)) + } +} + +func runWorkspaceExitLeader(t *testing.T) { + t.Helper() + listenerFile := os.NewFile(3, "workspace-exit-listener") + child := exec.Command(os.Args[0], "-test.run=^TestWorkspaceExitHelper$") + child.Env = append(os.Environ(), workspaceExitHelperModeEnv+"=descendant", workspaceExitHelperMarkerEnv+"="+os.Getenv(workspaceExitHelperMarkerEnv)) + child.ExtraFiles = []*os.File{listenerFile} + stdout, err := child.StdoutPipe() + if err != nil { + t.Fatal(err) + } + if err := child.Start(); err != nil { + t.Fatal(err) + } + if err := listenerFile.Close(); err != nil { + t.Fatal(err) + } + scanner := bufio.NewScanner(stdout) + if !scanner.Scan() || scanner.Text() != "DESCENDANT_READY" { + t.Fatalf("descendant readiness = %q, err = %v", scanner.Text(), scanner.Err()) + } + fmt.Println("READY") +} + +func runWorkspaceExitDescendant(t *testing.T) { + t.Helper() + listenerFile := os.NewFile(3, "workspace-exit-listener") + listener, err := net.FileListener(listenerFile) + if err != nil { + t.Fatal(err) + } + if err := listenerFile.Close(); err != nil { + t.Fatal(err) + } + defer listener.Close() + if err := os.WriteFile(os.Getenv(workspaceExitHelperMarkerEnv), []byte(strconv.Itoa(os.Getpid())), 0o600); err != nil { + t.Fatal(err) + } + fmt.Println("DESCENDANT_READY") + signals := make(chan os.Signal, 1) + signal.Notify(signals, syscall.SIGTERM) + defer signal.Stop(signals) + <-signals +} + +func readWorkspaceExitHelperPID(t *testing.T, path string) int { + t.Helper() + data, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + pid, err := strconv.Atoi(strings.TrimSpace(string(data))) + if err != nil { + t.Fatal(err) + } + return pid +} + +func readProcessGroupID(t *testing.T, pid int) int { + t.Helper() + data, err := os.ReadFile(filepath.Join("/proc", strconv.Itoa(pid), "stat")) + if err != nil { + t.Fatal(err) + } + commandEnd := strings.LastIndex(string(data), ") ") + if commandEnd < 0 { + t.Fatalf("invalid process stat for PID %d", pid) + } + fields := strings.Fields(string(data)[commandEnd+2:]) + if len(fields) < 3 { + t.Fatalf("process stat for PID %d has %d fields after command", pid, len(fields)) + } + processGroupID, err := strconv.Atoi(fields[2]) + if err != nil { + t.Fatal(err) + } + return processGroupID +} + +func requireProcessHoldsSocket(t *testing.T, pid int, inode uint64) { + t.Helper() + descriptorDirectory := filepath.Join("/proc", strconv.Itoa(pid), "fd") + descriptors, err := os.ReadDir(descriptorDirectory) + if err != nil { + t.Fatal(err) + } + wantedTarget := fmt.Sprintf("socket:[%d]", inode) + for _, descriptor := range descriptors { + target, err := os.Readlink(filepath.Join(descriptorDirectory, descriptor.Name())) + if err == nil && target == wantedTarget { + return + } + } + t.Fatalf("PID %d does not hold %s", pid, wantedTarget) +} diff --git a/integration/agentcompat/internal/workspace/supervisor_integration_test.go b/integration/agentcompat/internal/workspace/supervisor_integration_test.go new file mode 100644 index 00000000..2ba0640e --- /dev/null +++ b/integration/agentcompat/internal/workspace/supervisor_integration_test.go @@ -0,0 +1,104 @@ +//go:build linux + +package workspace + +import ( + "context" + "fmt" + "net" + "os" + "os/signal" + "strconv" + "strings" + "syscall" + "testing" + "time" + + processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process" +) + +const workspaceListenerHelperEnv = "GO_WANT_WORKSPACE_LISTENER_HELPER" + +func TestWorkspace_TransfersListenerToSupervisorAndRemovesResidue(t *testing.T) { + // Given + workspace, err := New(context.Background()) + if err != nil { + t.Fatal(err) + } + root := workspace.Root() + listener, err := net.Listen("tcp4", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + owned, err := workspace.AdoptListener(listener) + if err != nil { + t.Fatal(err) + } + extraFile, err := owned.ExtraFile() + if err != nil { + t.Fatal(err) + } + supervisor := processharness.NewSupervisor(t.Context(), processharness.Spec{ + Name: "workspace-listener", + Path: os.Args[0], + Args: []string{"-test.run=^TestWorkspaceListenerHelper$"}, + Env: append(os.Environ(), workspaceListenerHelperEnv+"=3"), + ExtraFiles: []*os.File{extraFile}, + Stdout: os.Stdout, + Stderr: os.Stderr, + MaxLogBytes: 1024, + TerminateTimeout: 100 * time.Millisecond, + KillTimeout: time.Second, + Readiness: func(_ processharness.Stream, line string) bool { + return strings.Contains(line, "READY") + }, + }) + if err := supervisor.Start(); err != nil { + t.Fatal(err) + } + if err := workspace.TrackPID(supervisor.PID()); err != nil { + t.Fatal(err) + } + if err := workspace.TrackProcessGroup(supervisor.ProcessGroupID()); err != nil { + t.Fatal(err) + } + if err := supervisor.WaitReady(t.Context()); err != nil { + t.Fatal(err) + } + + // When + if err := supervisor.Stop(t.Context()); err != nil { + t.Fatal(err) + } + if err := workspace.Close(); err != nil { + t.Fatal(err) + } + + // Then + if _, err := os.Stat(root); !os.IsNotExist(err) { + t.Fatalf("workspace remains: %v", err) + } +} + +func TestWorkspaceListenerHelper(t *testing.T) { + rawDescriptor := os.Getenv(workspaceListenerHelperEnv) + if rawDescriptor == "" { + return + } + descriptor, err := strconv.Atoi(rawDescriptor) + if err != nil { + t.Fatal(err) + } + file := os.NewFile(uintptr(descriptor), "workspace-listener") + listener, err := net.FileListener(file) + if err != nil { + t.Fatal(err) + } + _ = file.Close() + defer listener.Close() + fmt.Println("READY") + signals := make(chan os.Signal, 1) + signal.Notify(signals, syscall.SIGTERM) + defer signal.Stop(signals) + <-signals +} diff --git a/integration/agentcompat/internal/workspace/workspace.go b/integration/agentcompat/internal/workspace/workspace.go new file mode 100644 index 00000000..7822ac25 --- /dev/null +++ b/integration/agentcompat/internal/workspace/workspace.go @@ -0,0 +1,265 @@ +//go:build linux + +package workspace + +import ( + "context" + "errors" + "fmt" + "net" + "os" + "path/filepath" + "strconv" + "strings" + "sync" + "syscall" +) + +const defaultLogBytes = 1024 * 1024 + +type Workspace struct { + root string + binDir string + logDir string + payloadDir string + done chan struct{} + closeMu sync.Mutex + closed bool + closing bool + mu sync.Mutex + logs []*LogFile + listeners []*OwnedListener + pids map[int]struct{} + groups map[int]struct{} +} + +func New(ctx context.Context) (*Workspace, error) { + root, err := os.MkdirTemp("", "nezha-agentcompat-") + if err != nil { + return nil, fmt.Errorf("create agent compatibility workspace: %w", err) + } + workspace := &Workspace{ + root: root, + binDir: filepath.Join(root, "bin"), + logDir: filepath.Join(root, "logs"), + payloadDir: filepath.Join(root, "payloads"), + done: make(chan struct{}), + pids: make(map[int]struct{}), + groups: make(map[int]struct{}), + } + for _, directory := range []string{workspace.binDir, workspace.logDir, workspace.payloadDir} { + if err := os.Mkdir(directory, 0o700); err != nil { + _ = os.RemoveAll(root) + return nil, fmt.Errorf("create workspace directory %s: %w", filepath.Base(directory), err) + } + } + go workspace.closeOnCancellation(ctx) + return workspace, nil +} + +func (workspace *Workspace) closeOnCancellation(ctx context.Context) { + select { + case <-ctx.Done(): + _ = workspace.Close() + case <-workspace.done: + } +} + +func (workspace *Workspace) Root() string { return workspace.root } + +func (workspace *Workspace) BinaryPath(name string) (string, error) { + workspace.closeMu.Lock() + defer workspace.closeMu.Unlock() + if err := workspace.requireOpen(); err != nil { + return "", err + } + if err := validateLeafName(name); err != nil { + return "", err + } + return filepath.Join(workspace.binDir, name), nil +} + +func (workspace *Workspace) PayloadPath(name string) (string, error) { + workspace.closeMu.Lock() + defer workspace.closeMu.Unlock() + if err := workspace.requireOpen(); err != nil { + return "", err + } + if err := validateLeafName(name); err != nil { + return "", err + } + return filepath.Join(workspace.payloadDir, name), nil +} + +func validateLeafName(name string) error { + if name == "" || name == "." || name == ".." || filepath.Base(name) != name || strings.ContainsRune(name, os.PathSeparator) { + return errors.New("workspace name must be one path component") + } + return nil +} + +func (workspace *Workspace) Log(name string) (*LogFile, error) { + workspace.closeMu.Lock() + defer workspace.closeMu.Unlock() + if err := workspace.requireOpen(); err != nil { + return nil, err + } + if err := validateLeafName(name); err != nil { + return nil, err + } + logFile, err := newLogFile(filepath.Join(workspace.logDir, name+".log"), defaultLogBytes) + if err != nil { + return nil, err + } + workspace.mu.Lock() + workspace.logs = append(workspace.logs, logFile) + workspace.mu.Unlock() + return logFile, nil +} + +func (workspace *Workspace) AdoptListener(listener net.Listener) (*OwnedListener, error) { + workspace.closeMu.Lock() + defer workspace.closeMu.Unlock() + if err := workspace.requireOpen(); err != nil { + return nil, err + } + tcpListener, ok := listener.(*net.TCPListener) + if !ok { + return nil, errors.New("workspace listener must be TCP") + } + address, ok := tcpListener.Addr().(*net.TCPAddr) + if !ok || !address.IP.IsLoopback() { + return nil, errors.New("workspace listener must be loopback TCP") + } + file, err := tcpListener.File() + if err != nil { + return nil, fmt.Errorf("duplicate workspace listener: %w", err) + } + inode, err := socketInode(file) + if err != nil { + _ = file.Close() + return nil, err + } + if err := tcpListener.Close(); err != nil { + _ = file.Close() + return nil, fmt.Errorf("transfer workspace listener ownership: %w", err) + } + owned := &OwnedListener{file: file, address: address.String(), inode: inode} + workspace.mu.Lock() + workspace.listeners = append(workspace.listeners, owned) + workspace.mu.Unlock() + return owned, nil +} + +func (workspace *Workspace) TrackPID(pid int) error { + workspace.closeMu.Lock() + defer workspace.closeMu.Unlock() + if err := workspace.requireOpen(); err != nil { + return err + } + if pid < 1 { + return errors.New("tracked PID must be positive") + } + workspace.mu.Lock() + workspace.pids[pid] = struct{}{} + workspace.mu.Unlock() + return nil +} + +func (workspace *Workspace) TrackProcessGroup(processGroupID int) error { + workspace.closeMu.Lock() + defer workspace.closeMu.Unlock() + if err := workspace.requireOpen(); err != nil { + return err + } + if processGroupID < 1 { + return errors.New("tracked process group ID must be positive") + } + workspace.mu.Lock() + workspace.groups[processGroupID] = struct{}{} + workspace.mu.Unlock() + return nil +} + +func (workspace *Workspace) Close() error { + workspace.closeMu.Lock() + defer workspace.closeMu.Unlock() + if workspace.closed { + return nil + } + workspace.closing = true + if err := workspace.close(); err != nil { + workspace.closing = false + return err + } + workspace.closed = true + close(workspace.done) + return nil +} + +func (workspace *Workspace) requireOpen() error { + if workspace.closed || workspace.closing { + return errors.New("workspace is closing or closed") + } + return nil +} + +func (workspace *Workspace) close() error { + workspace.mu.Lock() + logs := append([]*LogFile(nil), workspace.logs...) + listeners := append([]*OwnedListener(nil), workspace.listeners...) + pids := make([]int, 0, len(workspace.pids)) + for pid := range workspace.pids { + pids = append(pids, pid) + } + groups := make([]int, 0, len(workspace.groups)) + for processGroupID := range workspace.groups { + groups = append(groups, processGroupID) + } + workspace.mu.Unlock() + + var cleanupErrors []error + for _, logFile := range logs { + if err := logFile.Close(); err != nil { + cleanupErrors = append(cleanupErrors, err) + } + } + for _, listener := range listeners { + if err := listener.Close(); err != nil { + cleanupErrors = append(cleanupErrors, err) + } + } + for _, pid := range pids { + if _, err := os.Stat(filepath.Join("/proc", strconv.Itoa(pid))); err == nil { + cleanupErrors = append(cleanupErrors, fmt.Errorf("tracked PID %d remains", pid)) + } else if !errors.Is(err, os.ErrNotExist) { + cleanupErrors = append(cleanupErrors, fmt.Errorf("inspect tracked PID %d: %w", pid, err)) + } + } + for _, processGroupID := range groups { + if err := syscall.Kill(-processGroupID, 0); err == nil || errors.Is(err, syscall.EPERM) { + cleanupErrors = append(cleanupErrors, fmt.Errorf("tracked process group %d remains", processGroupID)) + } else if !errors.Is(err, syscall.ESRCH) { + cleanupErrors = append(cleanupErrors, fmt.Errorf("inspect tracked process group %d: %w", processGroupID, err)) + } + } + for _, listener := range listeners { + present, err := listenerInodePresent(listener.inode) + if err != nil { + cleanupErrors = append(cleanupErrors, err) + } else if present { + cleanupErrors = append(cleanupErrors, fmt.Errorf("listener %s inode %d remains", listener.address, listener.inode)) + } + } + if err := os.RemoveAll(workspace.payloadDir); err != nil { + cleanupErrors = append(cleanupErrors, fmt.Errorf("remove workspace payloads: %w", err)) + } else if _, err := os.Stat(workspace.payloadDir); !errors.Is(err, os.ErrNotExist) { + cleanupErrors = append(cleanupErrors, fmt.Errorf("workspace payload residue remains: %w", err)) + } + if len(cleanupErrors) == 0 { + if err := os.RemoveAll(workspace.root); err != nil { + cleanupErrors = append(cleanupErrors, fmt.Errorf("remove workspace: %w", err)) + } + } + return errors.Join(cleanupErrors...) +} diff --git a/integration/agentcompat/internal/workspace/workspace_test.go b/integration/agentcompat/internal/workspace/workspace_test.go new file mode 100644 index 00000000..d08edbc7 --- /dev/null +++ b/integration/agentcompat/internal/workspace/workspace_test.go @@ -0,0 +1,211 @@ +//go:build linux + +package workspace + +import ( + "context" + "errors" + "net" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestWorkspace_RemovesWorkspace(t *testing.T) { + // Given + workspace, err := New(t.Context()) + if err != nil { + t.Fatal(err) + } + root := workspace.Root() + // When + if err := workspace.Close(); err != nil { + t.Fatal(err) + } + // Then + if _, err := os.Stat(root); !os.IsNotExist(err) { + t.Fatalf("workspace still exists: %v", err) + } +} + +func TestWorkspace_RedactsLogs(t *testing.T) { + // Given + workspace, err := New(t.Context()) + if err != nil { + t.Fatal(err) + } + defer workspace.Close() + // When + logFile, err := workspace.Log("agent") + if err != nil { + t.Fatal(err) + } + if _, err := logFile.WriteString("Authorization: Bearer eyJsecret.secret.secret password=top-secret\n" + strings.Repeat("x", defaultLogBytes*2)); err != nil { + t.Fatal(err) + } + if err := logFile.Close(); err != nil { + t.Fatal(err) + } + // Then + data, err := os.ReadFile(logFile.Name()) + if err != nil { + t.Fatal(err) + } + if strings.Contains(string(data), "top-secret") || strings.Contains(string(data), "eyJsecret") { + t.Fatalf("secret survived log redaction: %s", data) + } + if len(data) > defaultLogBytes { + t.Fatalf("log bytes = %d, limit = %d", len(data), defaultLogBytes) + } +} + +func TestWorkspace_AdoptsListener(t *testing.T) { + // Given + workspace, err := New(t.Context()) + if err != nil { + t.Fatal(err) + } + defer workspace.Close() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + // When + owned, err := workspace.AdoptListener(listener) + if err != nil { + t.Fatal(err) + } + // Then + if _, err := listener.Accept(); err == nil { + t.Fatal("original listener retained ownership") + } + if owned.FileDescriptor() < 3 { + t.Fatalf("listener FD = %d", owned.FileDescriptor()) + } + extraFile, err := owned.ExtraFile() + if err != nil { + t.Fatal(err) + } + if err := extraFile.Close(); err != nil { + t.Fatal(err) + } +} + +func TestWorkspace_BuildsBinaryInRunDirectory(t *testing.T) { + // Given + workspace, err := New(t.Context()) + if err != nil { + t.Fatal(err) + } + defer workspace.Close() + source := t.TempDir() + if err := os.WriteFile(filepath.Join(source, "go.mod"), []byte("module example.com/workspacefixture\n\ngo 1.26.3\n"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(source, "main.go"), []byte("package main\nfunc main() {}\n"), 0o600); err != nil { + t.Fatal(err) + } + + // When + binaryPath, err := workspace.Build(t.Context(), BuildSpec{Name: "fixture", SourceDir: source, Package: "."}) + + // Then + if err != nil { + t.Fatal(err) + } + if filepath.Dir(binaryPath) != filepath.Join(workspace.Root(), "bin") { + t.Fatalf("binary path = %s", binaryPath) + } + if info, err := os.Stat(binaryPath); err != nil || info.Mode()&0o111 == 0 { + t.Fatalf("binary is not executable: info=%v err=%v", info, err) + } +} + +func TestWorkspace_PreservesEvidenceWhenTrackedPIDRemains(t *testing.T) { + // Given + workspace, err := New(context.Background()) + if err != nil { + t.Fatal(err) + } + root := workspace.Root() + if err := workspace.TrackPID(os.Getpid()); err != nil { + t.Fatal(err) + } + + // When + err = workspace.Close() + + // Then + if err == nil || !strings.Contains(err.Error(), "tracked PID") { + t.Fatalf("close error = %v", err) + } + if _, statErr := os.Stat(root); statErr != nil { + t.Fatalf("workspace evidence was removed: %v", statErr) + } + if removeErr := os.RemoveAll(root); removeErr != nil && !errors.Is(removeErr, os.ErrNotExist) { + t.Fatal(removeErr) + } +} + +func TestWorkspace_RetriesCleanupAfterTrackedPIDExits(t *testing.T) { + // Given + workspace, err := New(context.Background()) + if err != nil { + t.Fatal(err) + } + root := workspace.Root() + command, input := startWorkspaceHelper(t) + if err := workspace.TrackPID(command.Process.Pid); err != nil { + t.Fatal(err) + } + if err := workspace.Close(); err == nil { + t.Fatal("workspace closed while tracked PID was live") + } + + // When + if err := input.Close(); err != nil { + t.Fatal(err) + } + if err := command.Wait(); err != nil { + t.Fatal(err) + } + if err := workspace.Close(); err != nil { + t.Fatal(err) + } + + // Then + if _, err := os.Stat(root); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("workspace remains after retry: %v", err) + } +} + +func TestWorkspace_RejectsResourcesAfterClose(t *testing.T) { + // Given + workspace, err := New(context.Background()) + if err != nil { + t.Fatal(err) + } + if err := workspace.Close(); err != nil { + t.Fatal(err) + } + listener, err := net.Listen("tcp4", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer listener.Close() + + // When / Then + if _, err := workspace.Log("late"); err == nil { + t.Fatal("closed workspace accepted a log") + } + if _, err := workspace.AdoptListener(listener); err == nil { + t.Fatal("closed workspace accepted a listener") + } + if err := workspace.TrackPID(os.Getpid()); err == nil { + t.Fatal("closed workspace accepted a PID") + } + if err := workspace.TrackProcessGroup(os.Getpid()); err == nil { + t.Fatal("closed workspace accepted a process group") + } +}