diff options
Diffstat (limited to 'test/internal')
| -rw-r--r-- | test/internal/sandbox/assert.go | 247 | ||||
| -rw-r--r-- | test/internal/sandbox/assert_test.go | 34 | ||||
| -rw-r--r-- | test/internal/sandbox/seccomp.go | 46 | ||||
| -rw-r--r-- | test/internal/testsuite/proc.go | 23 | ||||
| -rw-r--r-- | test/internal/testsuite/ptrace.go | 44 | ||||
| -rw-r--r-- | test/internal/testsuite/testsuite.go | 231 |
6 files changed, 270 insertions, 355 deletions
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 <sys/quota.h> -*/ -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") + }() +} |
