test(agentcompat): add process supervision harness

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-20 04:43:28 +00:00
co-authored by naiba/CloudCode
parent 0543995ec3
commit 3a9e0c7887
22 changed files with 3050 additions and 0 deletions
@@ -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
}
@@ -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
}
@@ -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
}
@@ -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
}
@@ -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)
}
@@ -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
}
@@ -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...)
}
@@ -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")
}
}