mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 09:40:12 +00:00
test(agentcompat): add process supervision harness
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user