From a7383510fb05abc98b992240cb76ad4b6c598956 Mon Sep 17 00:00:00 2001 From: Ophestra Date: Sun, 4 Oct 2026 01:05:10 +0900 Subject: test/sandbox: migrate tests This significantly improves performance, removing overhead of nix, python, and virtualisation. Running this in an unprivileged container required patching the kernel, but since special runner setup was already needed, that was an acceptable tradeoff. Signed-off-by: Ophestra --- test/internal/sandbox/assert.go | 247 ----------------------------------- test/internal/sandbox/assert_test.go | 34 ----- test/internal/sandbox/seccomp.go | 46 ------- test/internal/testsuite/proc.go | 23 ++-- test/internal/testsuite/ptrace.go | 44 ++++--- test/internal/testsuite/testsuite.go | 231 ++++++++++++++++++++++++++++++++ 6 files changed, 270 insertions(+), 355 deletions(-) delete mode 100644 test/internal/sandbox/assert.go delete mode 100644 test/internal/sandbox/assert_test.go delete mode 100644 test/internal/sandbox/seccomp.go (limited to 'test/internal') diff --git a/test/internal/sandbox/assert.go b/test/internal/sandbox/assert.go deleted file mode 100644 index 1194befb..00000000 --- a/test/internal/sandbox/assert.go +++ /dev/null @@ -1,247 +0,0 @@ -//go:build testtool - -// Package sandbox provides utilities for checking sandbox outcome. -// -// This package must never be used outside integration tests, there is a much -// better native implementation of mountinfo in the public sandbox/vfs package. -// Files in this package are excluded by the build system to prevent accidental -// misuse. -package sandbox - -import ( - "encoding/json" - "errors" - "io/fs" - "log" - "net" - "os" - "path/filepath" - "syscall" - - "hakurei.app/test/internal/mountinfo" - "hakurei.app/test/internal/testsuite" -) - -var ( - assert = log.New(os.Stderr, "sandbox: ", 0) - printfFunc = assert.Printf - fatalfFunc = assert.Fatalf -) - -func printf(format string, v ...any) { printfFunc(format, v...) } -func fatalf(format string, v ...any) { fatalfFunc(format, v...) } - -type TestCase struct { - Env []string `json:"env"` - FS *testsuite.FS `json:"fs"` - Mount []*mountinfo.Entry `json:"mount"` - Seccomp bool `json:"seccomp"` - - TrySocket string `json:"try_socket,omitempty"` - SocketAbstract bool `json:"socket_abstract,omitempty"` - SocketPathname bool `json:"socket_pathname,omitempty"` -} - -type T struct { - FS fs.FS - - MountsPath string -} - -func (t *T) MustCheckFile(wantFilePath string) { - var want *TestCase - mustDecode(wantFilePath, &want) - t.MustCheck(want) -} - -func mustAbs(s string) string { - if !filepath.IsAbs(s) { - fatalf("[FAIL] %q is not absolute", s) - panic("unreachable") - } - return s -} - -func (t *T) MustCheck(want *TestCase) { - checkWritableDirPaths := []string{ - "/dev/shm", - "/tmp", - os.Getenv("XDG_RUNTIME_DIR"), - } - for _, a := range checkWritableDirPaths { - pathname := filepath.Join(mustAbs(a), ".hakurei-check") - if err := os.WriteFile(pathname, make([]byte, 1<<8), 0600); err != nil { - fatalf("[FAIL] %s", err) - } else if err = os.Remove(pathname); err != nil { - fatalf("[FAIL] %s", err) - } else { - printf("[ OK ] %s is writable", a) - } - } - - if want.Env != nil { - var ( - fail bool - i int - got string - ) - for i, got = range os.Environ() { - if i == len(want.Env) { - fatalf("got more than %d environment variables", len(want.Env)) - } - if got != want.Env[i] { - fail = true - printf("[FAIL] %s", got) - } else { - printf("[ OK ] %s", got) - } - } - - i++ - if i != len(want.Env) { - fatalf("got %d environment variables, want %d", i, len(want.Env)) - } - - if fail { - fatalf("[FAIL] some environment variables did not match") - } - } else { - printf("[SKIP] skipping environ check") - } - - if want.FS != nil && t.FS != nil { - if err := want.FS.Compare(printfFunc, ".", t.FS); err != nil { - fatalf("%v", err) - } - } else { - printf("[SKIP] skipping fs check") - } - - if want.Mount != nil { - var fail bool - m := mustParseMountinfo(t.MountsPath) - i := 0 - var ent mountinfo.Entry - for m.Next() { - m.Copy(&ent) - - if i == len(want.Mount) { - fatalf("got more than %d entries", i) - } - if !ent.EqualWithIgnore(want.Mount[i], "//ignore") { - fail = true - printf("[FAIL] %s", &ent) - } else { - printf("[ OK ] %s", &ent) - } - - i++ - } - if err := m.Err(); err != nil { - fatalf("%v", err) - } - - if i != len(want.Mount) { - fatalf("got %d entries, want %d", i, len(want.Mount)) - } - - if fail { - fatalf("[FAIL] some mount points did not match") - } - } else { - printf("[SKIP] skipping mounts check") - } - - if want.Seccomp { - if trySyscalls() != nil { - os.Exit(1) - } - } else { - printf("[SKIP] skipping seccomp check") - } - - if want.TrySocket != "" { - abstractConn, abstractErr := net.Dial("unix", "@"+want.TrySocket) - pathnameConn, pathnameErr := net.Dial("unix", want.TrySocket) - ok := true - - if abstractErr == nil { - if err := abstractConn.Close(); err != nil { - ok = false - log.Printf("Close: %v", err) - } - } - if pathnameErr == nil { - if err := pathnameConn.Close(); err != nil { - ok = false - log.Printf("Close: %v", err) - } - } - - abstractWantErr := error(syscall.EPERM) - pathnameWantErr := error(syscall.ENOENT) - if want.SocketAbstract { - abstractWantErr = nil - } - if want.SocketPathname { - pathnameWantErr = nil - } - - if !errors.Is(abstractErr, abstractWantErr) { - ok = false - log.Printf("abstractErr: %v, want %v", abstractErr, abstractWantErr) - } - if !errors.Is(pathnameErr, pathnameWantErr) { - ok = false - log.Printf("pathnameErr: %v, want %v", pathnameErr, pathnameWantErr) - } - - if !ok { - os.Exit(1) - } - } -} - -func MustCheckFilter(pid int, want string) { - err := testsuite.CheckFilter(pid, 0, want) - if err == nil { - return - } - - e, ok := errors.AsType[*os.SyscallError](err) - if !ok { - fatalf("%s", err) - } - switch e.Syscall { - case "PTRACE_ATTACH": - fatalf("cannot attach to process %d: %v", pid, err) - case "PTRACE_SECCOMP_GET_FILTER": - if errors.Is(e.Err, syscall.ENOENT) { - fatalf("seccomp filter not installed for process %d", pid) - } - fatalf("cannot get filter: %v", err) - default: - fatalf("cannot check filter: %v", err) - } - - *(*int)(nil) = 0 // not reached -} - -func mustDecode(wantFilePath string, v any) { - if f, err := os.Open(wantFilePath); err != nil { - fatalf("cannot open %q: %v", wantFilePath, err) - } else if err = json.NewDecoder(f).Decode(v); err != nil { - fatalf("cannot decode %q: %v", wantFilePath, err) - } else if err = f.Close(); err != nil { - fatalf("cannot close %q: %v", wantFilePath, err) - } -} - -func mustParseMountinfo(name string) *mountinfo.Iter { - m, err := mountinfo.Open(name) - if err != nil { - fatalf("%v", err) - panic("unreachable") - } - return m -} diff --git a/test/internal/sandbox/assert_test.go b/test/internal/sandbox/assert_test.go deleted file mode 100644 index 012ae23d..00000000 --- a/test/internal/sandbox/assert_test.go +++ /dev/null @@ -1,34 +0,0 @@ -//go:build testtool - -package sandbox - -import ( - "encoding/json" - "os" - "path/filepath" - "testing" -) - -type F func(format string, v ...any) - -func SwapPrint(f F) (old F) { old = printfFunc; printfFunc = f; return } -func SwapFatal(f F) (old F) { old = fatalfFunc; fatalfFunc = f; return } - -func MustWantFile(t *testing.T, v any) (wantFile string) { - wantFile = filepath.Join(t.TempDir(), "want.json") - if f, err := os.OpenFile(wantFile, os.O_CREATE|os.O_WRONLY, 0400); err != nil { - t.Fatalf("cannot create %q: %v", wantFile, err) - } else if err = json.NewEncoder(f).Encode(v); err != nil { - t.Fatalf("cannot encode to %q: %v", wantFile, err) - } else if err = f.Close(); err != nil { - t.Fatalf("cannot close %q: %v", wantFile, err) - } - - t.Cleanup(func() { - if err := os.Remove(wantFile); err != nil { - t.Fatalf("cannot remove %q: %v", wantFile, err) - } - }) - - return -} diff --git a/test/internal/sandbox/seccomp.go b/test/internal/sandbox/seccomp.go deleted file mode 100644 index 1d8cd457..00000000 --- a/test/internal/sandbox/seccomp.go +++ /dev/null @@ -1,46 +0,0 @@ -//go:build testtool - -package sandbox - -import ( - "os" - "syscall" -) - -/* -#include -*/ -import "C" - -const NULL = 0 - -func trySyscalls() error { - testCases := []struct { - name string - errno syscall.Errno - - trap, a1, a2, a3, a4, a5, a6 uintptr - }{ - {"syslog", syscall.EPERM, syscall.SYS_SYSLOG, 0, NULL, NULL, NULL, NULL, NULL}, - {"acct", syscall.EPERM, syscall.SYS_ACCT, 0, NULL, NULL, NULL, NULL, NULL}, - {"quotactl", syscall.EPERM, syscall.SYS_QUOTACTL, C.Q_GETQUOTA, NULL, uintptr(os.Getuid()), NULL, NULL, NULL}, - {"add_key", syscall.EPERM, syscall.SYS_ADD_KEY, NULL, NULL, NULL, NULL, NULL, NULL}, - {"keyctl", syscall.EPERM, syscall.SYS_KEYCTL, NULL, NULL, NULL, NULL, NULL, NULL}, - {"request_key", syscall.EPERM, syscall.SYS_REQUEST_KEY, NULL, NULL, NULL, NULL, NULL, NULL}, - {"move_pages", syscall.EPERM, syscall.SYS_MOVE_PAGES, uintptr(os.Getpid()), NULL, NULL, NULL, NULL, NULL}, - {"mbind", syscall.EPERM, syscall.SYS_MBIND, NULL, NULL, NULL, NULL, NULL, NULL}, - {"get_mempolicy", syscall.EPERM, syscall.SYS_GET_MEMPOLICY, NULL, NULL, NULL, NULL, NULL, NULL}, - {"set_mempolicy", syscall.EPERM, syscall.SYS_SET_MEMPOLICY, NULL, NULL, NULL, NULL, NULL, NULL}, - {"migrate_pages", syscall.EPERM, syscall.SYS_MIGRATE_PAGES, NULL, NULL, NULL, NULL, NULL, NULL}, - } - - for _, tc := range testCases { - if _, _, errno := syscall.Syscall6(tc.trap, tc.a1, tc.a2, tc.a3, tc.a4, tc.a5, tc.a6); errno != tc.errno { - printf("[FAIL] %s: %v, want %v", tc.name, errno, tc.errno) - return errno - } - printf("[ OK ] %s: %v", tc.name, tc.errno) - } - - return nil -} diff --git a/test/internal/testsuite/proc.go b/test/internal/testsuite/proc.go index f2ad8857..e7ef1aed 100644 --- a/test/internal/testsuite/proc.go +++ b/test/internal/testsuite/proc.go @@ -10,6 +10,8 @@ import ( "strings" "syscall" "unsafe" + + "hakurei.app/fhs" ) // Stat represents status information read from /proc/pid/stat. @@ -144,13 +146,9 @@ type Stat struct { CGuestTime int } -// fhsProc points to a virtual kernel file system exposing the process list and -// other functionality. -const fhsProc = "/proc/" - // Executable is like [os.Executable], but for the process referred to by s. func (s *Stat) Executable() (string, error) { - path, err := os.Readlink(filepath.Join(fhsProc, strconv.Itoa(s.PID), "exe")) + path, err := os.Readlink(filepath.Join(fhs.Proc, strconv.Itoa(s.PID), "exe")) // When the executable has been deleted then Readlink returns a // path appended with " (deleted)". @@ -159,7 +157,7 @@ func (s *Stat) Executable() (string, error) { // Stat populates stat with the proc filesystem entry referred to by s. func (s *Stat) Stat(stat *syscall.Stat_t) (err error) { - err = syscall.Stat(filepath.Join(fhsProc, strconv.Itoa(s.PID)), stat) + err = syscall.Stat(filepath.Join(fhs.Proc, strconv.Itoa(s.PID)), stat) if err != nil { err = os.NewSyscallError("stat", err) } @@ -168,7 +166,7 @@ func (s *Stat) Stat(stat *syscall.Stat_t) (err error) { // Args reads arguments of the process referred to by s. func (s *Stat) Args() ([]string, error) { - p, err := os.ReadFile(filepath.Join(fhsProc, strconv.Itoa(s.PID), "cmdline")) + p, err := os.ReadFile(filepath.Join(fhs.Proc, strconv.Itoa(s.PID), "cmdline")) if err != nil { return nil, err } @@ -286,6 +284,11 @@ type StatScanner struct { err error } +// IsNotExist returns whether an error is [os.ErrNotExist] or ESRCH. +func IsNotExist(err error) bool { + return errors.Is(err, os.ErrNotExist) || errors.Is(err, syscall.ESRCH) +} + // Scan reads a process status information entry. It returns false if an // unrecoverable error is encountered, after which Scan no longer scans new // entries. @@ -295,7 +298,7 @@ func (s *StatScanner) Scan() bool { } if s.wrapped = s.i == len(s.dents); s.wrapped { - if s.dents, s.err = os.ReadDir(fhsProc); s.err != nil { + if s.dents, s.err = os.ReadDir(fhs.Proc); s.err != nil { return false } s.i = 0 @@ -318,9 +321,9 @@ func (s *StatScanner) Scan() bool { } var p []byte - p, err = os.ReadFile(filepath.Join(fhsProc, dent.Name(), "stat")) + p, err = os.ReadFile(filepath.Join(fhs.Proc, dent.Name(), "stat")) if err != nil { - if errors.Is(err, os.ErrNotExist) || errors.Is(err, syscall.ESRCH) { + if IsNotExist(err) { continue } s.err = err diff --git a/test/internal/testsuite/ptrace.go b/test/internal/testsuite/ptrace.go index 4fcf1508..ccf0900c 100644 --- a/test/internal/testsuite/ptrace.go +++ b/test/internal/testsuite/ptrace.go @@ -2,7 +2,7 @@ package testsuite import ( "crypto/sha512" - "encoding/hex" + "encoding/base64" "errors" "fmt" "os" @@ -58,10 +58,26 @@ func ptraceAttach(pid int) error { } return os.NewSyscallError("wait4", err) } - break - } + switch { + case status.Stopped(): + return nil - return nil + case status.Continued(): + continue + + case status.Signaled(): + return fmt.Errorf( + "tracee terminated by signal %s", + status.Signal(), + ) + + case status.Exited(): + return fmt.Errorf( + "tracee terminated unexpectedly with code %d", + status.ExitStatus(), + ) + } + } } // ptraceDetach detaches from the attached process referred to by pid. @@ -95,8 +111,8 @@ func getFilter(pid, index int) ([]syscall.SockFilter, error) { } // CheckFilter checks the process at pid to have its first filter's contents -// match the sha512 checksum specified in hexadecimal string representation. -func CheckFilter(pid, index int, sum string) (err error) { +// match the specified sha512 checksum. +func CheckFilter(pid, index int, sum [sha512.Size]byte) (err error) { if err = ptraceAttach(pid); err != nil { return } @@ -106,15 +122,7 @@ func CheckFilter(pid, index int, sum string) (err error) { } }() - var ( - buf []syscall.SockFilter - want []byte - ) - - if want, err = hex.DecodeString(sum); err != nil { - return - } - + var buf []syscall.SockFilter h := sha512.New() if buf, err = getFilter(pid, index); err != nil { return @@ -125,11 +133,11 @@ func CheckFilter(pid, index int, sum string) (err error) { )) } - if got := h.Sum(nil); string(got) != string(want) { + if got := h.Sum(nil); string(got) != string(sum[:]) { return fmt.Errorf( "bad filter\n\t got: %s\n\twant: %s", - hex.EncodeToString(got), - sum, + base64.StdEncoding.EncodeToString(got), + base64.StdEncoding.EncodeToString(sum[:]), ) } return diff --git a/test/internal/testsuite/testsuite.go b/test/internal/testsuite/testsuite.go index 6b4cd717..00eb2f92 100644 --- a/test/internal/testsuite/testsuite.go +++ b/test/internal/testsuite/testsuite.go @@ -5,12 +5,19 @@ package testsuite import ( + "bufio" + "context" + "crypto/sha512" + "errors" "log" "os" "os/exec" "os/signal" "os/user" + "strconv" + "sync" "syscall" + "time" ) // ReceiveSignals blocks until a termination signal arrives, and terminates. @@ -39,7 +46,231 @@ func MustRun(command ...string) { } } +// ErrUnexpectedSuccess is returned for processes expected to exit with a +// non-zero code, but failed to do so. +var ErrUnexpectedSuccess = errors.New("process unexpectedly exited with code 0") + +// MustFail runs command and terminates the testsuite if the program fails to +// start or exits with code 0. +func MustFail(command ...string) { + cmd := exec.Command(command[0], command[1:]...) + cmd.Stdout, cmd.Stderr = os.Stdout, os.Stderr + if err := cmd.Run(); err == nil { + log.Fatal(ErrUnexpectedSuccess) + } else if e, ok := errors.AsType[*exec.ExitError](err); !ok { + log.Fatal(err) + } else if !e.Exited() { + log.Fatal(e) + } +} + // MustRunAs wraps [MustRun] for sudo. func MustRunAs(username string, command ...string) { MustRun(append([]string{"sudo", "-u", username}, command...)...) } + +// MustFailAs wraps [MustFail] for sudo. +func MustFailAs(username string, command ...string) { + MustFail(append([]string{"sudo", "-u", username}, command...)...) +} + +// MustStart starts cmd and returns a channel delivering its wait error. +func MustStart(cmd *exec.Cmd) (done <-chan error) { + if err := cmd.Start(); err != nil { + log.Fatal(err) + } + d := make(chan error) + go func() { d <- cmd.Wait() }() + return d +} + +// MustStartAs wraps [MustStart] for sudo. +func MustStartAs( + ctx context.Context, + username string, + files []*os.File, + command ...string, +) (proc *os.Process, done <-chan error) { + sudoArgs := []string{ + "-u", username, + } + if len(files) != 0 { + sudoArgs = append(sudoArgs, "-C", strconv.Itoa(len(files)+4)) + } + sudoArgs = append(sudoArgs, "--") + cmd := exec.CommandContext(ctx, "sudo", append(sudoArgs, command...)...) + cmd.Stdout, cmd.Stderr = os.Stdout, os.Stderr + cmd.ExtraFiles = files + cmd.SysProcAttr = &syscall.SysProcAttr{Pdeathsig: syscall.SIGTERM} + return cmd.Process, MustStart(cmd) +} + +// MustCheckFilter is like [CheckFilter], but terminates the test suite if a +// non-nil error is returned. Otherwise, the tracee is terminated after it +// resumes. +func MustCheckFilter(pid int, sum [sha512.Size]byte) { + // podman installs its own filter + if err := CheckFilter(pid, 1, sum); err != nil { + log.Fatal(err) + } else if err = syscall.Kill(pid, syscall.SIGTERM); err != nil { + log.Fatalf("cannot terminate tracee: %v", err) + } +} + +// FilterTerminated returns a non-nil error if err is not an [exec.ExitError] +// describing a process terminated by a syscall.SIGTERM signal. +func FilterTerminated(err error) error { + if err == nil { + return ErrUnexpectedSuccess + } + + e, ok := errors.AsType[*exec.ExitError](err) + if !ok { + return err + } + + if e.ExitCode() == 0x80+int(syscall.SIGTERM) { + return nil + } + return e +} + +// Poll repeatedly runs command until it succeeds. +func Poll(d time.Duration, command ...string) { + for range time.NewTicker(d).C { + cmd := exec.Command(command[0], command[1:]...) + if err := cmd.Run(); err != nil { + if e, ok := errors.AsType[*exec.ExitError](err); ok && e.Exited() { + continue + } + log.Fatal(err) + } + break + } +} + +const ( + // XDGRuntimeDir is the hardcoded XDG runtime directory for the user + // described by [GetUser]. + XDGRuntimeDir = "/var/run/user/1000" + + // XDGRuntimeEnv is the environment variable string for XDG_RUNTIME_DIR. + XDGRuntimeEnv = "XDG_RUNTIME_DIR=" + XDGRuntimeDir +) + +// MustStartSessionBus starts a session bus that is never explicitly terminated. +// The test suite is terminated if the session bus daemon terminates. +func MustStartSessionBus(username string) (dbusEnv string) { + r, w, err := os.Pipe() + if err != nil { + log.Fatal(err) + } + + // this is never explicitly terminated + _, done := MustStartAs( + context.Background(), username, []*os.File{w}, + "dbus-daemon", + "--print-address=3", + "--address=unix:path="+XDGRuntimeDir+"/dbus", + "--session", + "--nofork", + "--nopidfile", + ) + + go func() { + if _err := <-done; _err != nil { + log.Fatal(_err) + } + log.Fatal("session bus terminated unexpectedly") + }() + + dbusEnv, err = bufio.NewReader(r).ReadString('\n') + if err != nil { + log.Fatal(err) + } + dbusEnv = dbusEnv[:len(dbusEnv)-1] + log.Printf("dbus listening on %s", dbusEnv) + dbusEnv = "DBUS_SESSION_BUS_ADDRESS=" + dbusEnv + + if err = r.Close(); err != nil { + log.Fatal(err) + } + return +} + +const ( + // SwayEnv is the environment variable string for the sway IPC socket. + SwayEnv = "SWAYSOCK=" + XDGRuntimeDir + "/sway" + // WaylandEnv is the environment variable string for the wayland display. + WaylandEnv = "WAYLAND_DISPLAY=wayland-1" +) + +// MustStartSway starts the sway wayland display server which must be terminated +// by calling [TerminateSway]. +func MustStartSway( + wg *sync.WaitGroup, + username, dbusEnv string, +) { + wg.Go(func() { + // this is terminated via swaymsg + _, done := MustStartAs( + context.Background(), username, nil, "env", + "WLR_BACKENDS=headless", + XDGRuntimeEnv, + SwayEnv, + dbusEnv, + "sway", + ) + if err := <-done; err != nil { + log.Fatal(err) + } + }) + + Poll(50*time.Millisecond, "sudo", "-u", username, SwayEnv, "swaymsg") + log.Printf("sway available via %s", SwayEnv) +} + +// TerminateSway requests for the sway server to terminate via sway IPC. +func TerminateSway(username string) { + MustFailAs(username, SwayEnv, "swaymsg", "exit") +} + +// MustStartPipeWire starts a PipeWire server that is never explicitly +// terminated. The test suite is terminated if the PipeWire server terminates. +func MustStartPipeWire(username, dbusEnv string) { + // this is never explicitly terminated + _, done := MustStartAs( + context.Background(), username, nil, "env", + XDGRuntimeEnv, + dbusEnv, + "pipewire", + ) + + go func() { + if _err := <-done; _err != nil { + log.Fatal(_err) + } + log.Fatal("pipewire terminated unexpectedly") + }() + + Poll(50*time.Millisecond, "sudo", "-u", username, + XDGRuntimeEnv, + dbusEnv, + "wpctl", + "status", + ) + + _, _done := MustStartAs( + context.Background(), username, nil, "env", + XDGRuntimeEnv, + dbusEnv, + "wireplumber", + ) + + go func() { + if _err := <-_done; _err != nil { + log.Fatal(_err) + } + log.Fatal("wireplumber terminated unexpectedly") + }() +} -- cgit v1.3.1