aboutsummaryrefslogtreecommitdiffhomepage
path: root/test/sandbox/tester/main.go
diff options
context:
space:
mode:
Diffstat (limited to 'test/sandbox/tester/main.go')
-rw-r--r--test/sandbox/tester/main.go224
1 files changed, 224 insertions, 0 deletions
diff --git a/test/sandbox/tester/main.go b/test/sandbox/tester/main.go
new file mode 100644
index 00000000..0782a5e0
--- /dev/null
+++ b/test/sandbox/tester/main.go
@@ -0,0 +1,224 @@
+//go:build tester
+
+// The sandbox tester runs within a cmd/hakurei container and validates its
+// state. Since the test environment is relatively predictable, the tester can
+// make various assumptions about the host.
+package main
+
+import (
+ "errors"
+ "log"
+ "net"
+ "os"
+ "os/signal"
+ "path/filepath"
+ "syscall"
+
+ "hakurei.app/test/internal/mountinfo"
+ "hakurei.app/test/sandbox/testdata"
+)
+
+//#include <sys/quota.h>
+import "C"
+
+// mustAbs returns s, or terminates the program if s is not absolute.
+func mustAbs(s string) string {
+ if !filepath.IsAbs(s) {
+ log.Fatalf("%q is not absolute", s)
+ }
+ return s
+}
+
+func main() {
+ log.SetFlags(0)
+ log.SetPrefix("tester: ")
+
+ if len(os.Args) != 2 {
+ log.Fatal("tester requires 1 argument")
+ }
+ want := testdata.Get(os.Args[1])
+ log.SetPrefix("tester: " + os.Args[1] + " ")
+
+ 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 {
+ log.Fatalf("[FAIL] %s", err)
+ } else if err = os.Remove(pathname); err != nil {
+ log.Fatalf("[FAIL] %s", err)
+ } else {
+ log.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) {
+ log.Fatalf("got more than %d environment variables", len(want.Env))
+ }
+ if got != want.Env[i] {
+ fail = true
+ log.Printf("[FAIL] %s", got)
+ } else {
+ log.Printf("[ OK ] %s", got)
+ }
+ }
+
+ i++
+ if i != len(want.Env) {
+ log.Fatalf("got %d environment variables, want %d", i, len(want.Env))
+ }
+
+ if fail {
+ log.Fatalf("[FAIL] some environment variables did not match")
+ }
+ } else {
+ log.Printf("[SKIP] skipping environ check")
+ }
+
+ if want.FS != nil {
+ if err := want.FS.Compare(log.Printf, ".", os.DirFS("/")); err != nil {
+ log.Fatalf("%v", err)
+ }
+ } else {
+ log.Printf("[SKIP] skipping fs check")
+ }
+
+ if want.Mount != nil {
+ var fail bool
+
+ m, err := mountinfo.Open("")
+ if err != nil {
+ log.Fatal(err)
+ }
+
+ i := 0
+ var ent mountinfo.Entry
+ for m.Next() {
+ m.Copy(&ent)
+
+ if i == len(want.Mount) {
+ log.Fatalf("got more than %d entries", i)
+ }
+ if !ent.EqualWithIgnore(want.Mount[i], "//ignore") {
+ fail = true
+ log.Printf("[FAIL] %s", &ent)
+ } else {
+ log.Printf("[ OK ] %s", &ent)
+ }
+
+ i++
+ }
+ if err = m.Err(); err != nil {
+ log.Fatalf("%v", err)
+ }
+
+ if i != len(want.Mount) {
+ log.Fatalf("got %d entries, want %d", i, len(want.Mount))
+ }
+
+ if fail {
+ log.Fatalf("[FAIL] some mount points did not match")
+ }
+ } else {
+ log.Printf("[SKIP] skipping mounts check")
+ }
+
+ if want.Seccomp {
+ const NULL = 0
+
+ for _, tc := range []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},
+ } {
+ if _, _, errno := syscall.Syscall6(tc.trap, tc.a1, tc.a2, tc.a3, tc.a4, tc.a5, tc.a6); errno != tc.errno {
+ log.Fatalf("[FAIL] %s: %v, want %v", tc.name, errno, tc.errno)
+ }
+ log.Printf("[ OK ] %s: %v", tc.name, tc.errno)
+ }
+ } else {
+ log.Printf("[SKIP] skipping seccomp check")
+ }
+
+ if want.TrySocket != "" {
+ retry:
+ 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)
+ }
+ }
+
+ if errors.Is(
+ abstractErr,
+ syscall.EAGAIN,
+ ) || errors.Is(
+ pathnameErr,
+ syscall.EAGAIN,
+ ) {
+ goto retry
+ }
+
+ abstractWantErr := error(want.ErrnoAbstract)
+ pathnameWantErr := error(want.ErrnoPathname)
+ if want.ErrnoAbstract == 0 {
+ abstractWantErr = nil
+ }
+ if want.ErrnoPathname == 0 {
+ 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)
+ }
+ }
+
+ s := make(chan os.Signal, 1)
+ signal.Notify(s, syscall.SIGTERM)
+ if _, err := os.Stdout.Write(make([]byte, 8)); err != nil {
+ log.Fatalf("cannot notify testsuite: %v", err)
+ }
+ <-s
+}