From 5647daa6204213d0f78c9853de572d50bd9af030 Mon Sep 17 00:00:00 2001 From: Ophestra Date: Thu, 1 Oct 2026 22:34:51 +0900 Subject: test/sandbox: migrate tests This benefits even more than the cmd/sharefs test suite, the slow python-based test script was a major bottleneck. Replacing the nix-represented test cases with compound literals also significantly increases readability. Signed-off-by: Ophestra --- test/sandbox/main.go | 181 +++++++++++++++++++++++++++++++++++++++++++++ test/sandbox/seccomp.patch | 18 +++++ test/sandbox/sum_amd64.go | 7 ++ test/sandbox/tool/main.go | 60 +-------------- 4 files changed, 207 insertions(+), 59 deletions(-) create mode 100644 test/sandbox/main.go create mode 100644 test/sandbox/seccomp.patch create mode 100644 test/sandbox/sum_amd64.go (limited to 'test/sandbox') diff --git a/test/sandbox/main.go b/test/sandbox/main.go new file mode 100644 index 00000000..a407d75d --- /dev/null +++ b/test/sandbox/main.go @@ -0,0 +1,181 @@ +//go:build testsuite + +// The sandbox test program runs cmd/hakurei with configurations simulating +// several common workloads and inspects the resulting container states. +package main + +import ( + "context" + "log" + "os" + "path/filepath" + "slices" + "strconv" + "sync" + "sync/atomic" + "syscall" + + "hakurei.app/fhs" + "hakurei.app/hst" + "hakurei.app/test/internal/testsuite" +) + +// mustStart starts a hakurei container and returns the pid of a process within +// the container. This process must be terminated by the caller. +func mustStart( + ctx context.Context, + serial uint64, + username string, + files ...*os.File, +) (pid int, done <-chan error) { + _serial := strconv.FormatUint(serial, 10) + _, done = testsuite.MustStartAs( + ctx, username, files, + "hakurei", "exec", + "sleep", "infinity", _serial, + ) + + var ( + s testsuite.StatScanner + + stat syscall.Stat_t + ) + for s.Scan() { + select { + case err := <-done: + if err == nil { + log.Fatal("test process terminated unexpectedly") + } + log.Fatal(err) + default: + break + } + + if s.Stat().Comm != "sleep" { + continue + } + + if args, err := s.Stat().Args(); err != nil { + log.Fatal(err) + } else if !slices.Equal(args, []string{ + "sleep", + "infinity", + _serial, + }) { + continue + } + + if err := s.Stat().Stat(&stat); err != nil { + log.Fatal(err) + } + + id := hst.ToUser[uint32](0, 0) + if stat.Uid != id || stat.Gid != id { + continue + } + + break + } + if err := s.Err(); err != nil { + log.Fatal(err) + } + pid = s.Stat().PID + return +} + +func main() { + go testsuite.ReceiveSignals() + username := testsuite.GetUsername() + + // the signal handler does not wait for termination + ctx := context.Background() + + var wg sync.WaitGroup + defer wg.Wait() + + var serial atomic.Uint64 + newSerial := func() uint64 { serial.Add(1); return serial.Load() } + + testsuite.MustRunAs( + username, "-i", + "hakurei", "exec", "capsh", "--print", + ) + wg.Go(func() { + defer log.Println("validated capabilities/securebits in user namespace") + + testsuite.MustRunAs( + username, "-i", + "hakurei", "exec", "capsh", "--has-no-new-privs", + ) + + for _, p := range []byte{'a', 'b', 'i', 'p'} { + testsuite.MustFailAs( + username, "-i", + "hakurei", "exec", "capsh", "--has-"+string(p)+"=CAP_SYS_ADMIN", + ) + } + testsuite.MustFailAs( + username, "-i", + "hakurei", "exec", "umount", "-R", "/dev", + ) + }) + + wg.Go(func() { + defer log.Println("validated pd seccomp outcome") + + c, cancel := context.WithCancel(ctx) + defer cancel() + + pid, done := mustStart(c, newSerial(), username) + testsuite.MustCheckFilter(pid, pdSum) + if err := testsuite.FilterTerminated(<-done); err != nil { + log.Fatal(err) + } + }) + + wg.Go(func() { + defer log.Println("validated fd leak") + + c, cancel := context.WithCancel(ctx) + defer cancel() + + pid, done := mustStart(c, newSerial(), username, os.Stdin, os.Stdout, os.Stderr) + prefix := filepath.Join(fhs.Proc, strconv.Itoa(pid), "fd") + + var fail bool + if entries, err := os.ReadDir(prefix); err != nil { + log.Fatal(err.Error()) + } else { + for _, ent := range entries { + var fd int + if fd, err = strconv.Atoi(ent.Name()); err != nil { + log.Fatal(err.Error()) + } + + // skip standard streams + if fd <= 2 { + continue + } + fail = true + + var d string + if d, err = os.Readlink(filepath.Join( + prefix, + ent.Name(), + )); err != nil { + log.Fatal(err.Error()) + } + log.Printf("extra fd %d -> %s", fd, d) + } + } + if fail { + log.Fatal("file descriptors leaked") + } + + if err := syscall.Kill(pid, syscall.SIGTERM); err != nil { + log.Fatalf("cannot terminate anchor: %v", err) + } else if err = testsuite.FilterTerminated(<-done); err != nil { + log.Fatal(err) + } + }) +} diff --git a/test/sandbox/seccomp.patch b/test/sandbox/seccomp.patch new file mode 100644 index 00000000..ddabc71e --- /dev/null +++ b/test/sandbox/seccomp.patch @@ -0,0 +1,18 @@ +diff --git a/kernel/seccomp.c b/kernel/seccomp.c +index 25f62867a16d..7b63ccc8daf4 100644 +--- a/kernel/seccomp.c ++++ b/kernel/seccomp.c +@@ -2216,8 +2216,12 @@ long seccomp_get_filter(struct task_struct *task, unsigned long filter_off, + struct seccomp_filter *filter; + struct sock_fprog_kern *fprog; + long ret; ++ struct user_namespace *user_ns = current_user_ns(); + +- if (!capable(CAP_SYS_ADMIN) || ++ if (in_userns(user_ns, task_cred_xxx(task, user_ns))) { ++ if (!ns_capable(user_ns, CAP_SYS_ADMIN)) ++ return -EACCES; ++ } else if (!capable(CAP_SYS_ADMIN) || + current->seccomp.mode != SECCOMP_MODE_DISABLED) { + return -EACCES; + } diff --git a/test/sandbox/sum_amd64.go b/test/sandbox/sum_amd64.go new file mode 100644 index 00000000..da03f12b --- /dev/null +++ b/test/sandbox/sum_amd64.go @@ -0,0 +1,7 @@ +//go:build testsuite + +package main + +const ( + pdSum = "c698b081ff957afe17a6d94374537d37f2a63f6f9dd75da7546542407a9e32476ebda3312ba7785d7f618542bcfaf27ca27dcc2dddba852069d28bcfe8cad39a" +) diff --git a/test/sandbox/tool/main.go b/test/sandbox/tool/main.go index 889142d4..490c3c5e 100644 --- a/test/sandbox/tool/main.go +++ b/test/sandbox/tool/main.go @@ -4,12 +4,9 @@ package main import ( "flag" - "fmt" "log" "os" "os/signal" - "strconv" - "strings" "syscall" "hakurei.app/test/internal/sandbox" @@ -47,60 +44,5 @@ func main() { return } - switch args[0] { - case "filter": - if len(args) != 2 { - log.Fatal("invalid argument") - } - - if pid, err := strconv.Atoi(strings.TrimSpace(args[1])); err != nil { - log.Fatalf("%s", err) - } else if pid < 1 { - log.Fatalf("%d out of range", pid) - } else { - sandbox.MustCheckFilter(pid, flagBpfHash) - if err = syscall.Kill(pid, syscall.SIGINT); err != nil { - log.Fatalf("cannot signal check process: %v", err) - } - } - - case "hash": // this eases the pain of passing the hash to python - fmt.Print(flagBpfHash) - - case "fd": - if len(args) != 2 { - log.Fatal("invalid argument") - } - prefix := fmt.Sprintf("/proc/%s/fd/", args[1]) - - var fail bool - if entries, err := os.ReadDir(prefix); err != nil { - log.Fatal(err.Error()) - } else { - for _, ent := range entries { - var fd int - if fd, err = strconv.Atoi(ent.Name()); err != nil { - log.Fatal(err.Error()) - } - - // skip standard streams - if fd <= 2 { - continue - } - fail = true - - var d string - if d, err = os.Readlink(prefix + ent.Name()); err != nil { - log.Fatal(err.Error()) - } - log.Printf("[FAIL] extra fd %d -> %s", fd, d) - } - } - if fail { - log.Fatal("[FAIL] file descriptors leaked") - } - - default: - log.Fatal("invalid argument") - } + log.Fatal("invalid argument") } -- cgit v1.3.1