diff options
Diffstat (limited to 'internal/testsuite')
| -rw-r--r-- | internal/testsuite/fs.go | 129 | ||||
| -rw-r--r-- | internal/testsuite/fs_test.go | 85 | ||||
| -rw-r--r-- | internal/testsuite/mountinfo/mountinfo.go | 177 | ||||
| -rw-r--r-- | internal/testsuite/mountinfo/mountinfo_guard.go | 15 | ||||
| -rw-r--r-- | internal/testsuite/mountinfo/mountinfo_test.go | 146 | ||||
| -rw-r--r-- | internal/testsuite/proc.go | 353 | ||||
| -rw-r--r-- | internal/testsuite/proc_test.go | 33 | ||||
| -rw-r--r-- | internal/testsuite/ptrace.go | 144 | ||||
| -rw-r--r-- | internal/testsuite/ptrace_test.go | 13 | ||||
| -rw-r--r-- | internal/testsuite/testsuite.go | 290 | ||||
| -rw-r--r-- | internal/testsuite/testsuite_guard.go | 15 | ||||
| -rw-r--r-- | internal/testsuite/testsuite_root.go | 17 |
12 files changed, 1417 insertions, 0 deletions
diff --git a/internal/testsuite/fs.go b/internal/testsuite/fs.go new file mode 100644 index 00000000..9acd24b6 --- /dev/null +++ b/internal/testsuite/fs.go @@ -0,0 +1,129 @@ +package testsuite + +import ( + "errors" + "fmt" + "io/fs" + "path/filepath" + "strings" +) + +var ( + // ErrFSBadLength is returned by [FS.Compare] for a directory with an + // unexpected amount of dents. + ErrFSBadLength = errors.New("bad dir length") + // ErrFSBadData is returned by [FS.Compare] for a file with unexpected + // contents. + ErrFSBadData = errors.New("data differs") + // ErrFSBadMode is returned by [FS.Compare] for an entry with unexpected + // mode. + ErrFSBadMode = errors.New("mode differs") + // ErrFSInvalidEnt is returned by [FS.Compare] if an invalid [FS] is visited. + ErrFSInvalidEnt = errors.New("invalid entry condition") +) + +// FS represents part of a filesystem hierarchy. +type FS struct { + // Expected mode of corresponding entry. + Mode fs.FileMode `json:"mode"` + // Expected directory contents. The directory is not descended if Dir is nil. + Dir map[string]*FS `json:"dir"` + // Expected file contents. The file is not read if Data is nil. + Data *string `json:"data"` +} + +// dprintf calls printf if it is non-nil. +func dprintf(printf func(format string, a ...any), format string, a ...any) { + if printf == nil { + return + } + printf(format, a...) +} + +// printDir prints a failed [FS.Compare] directory. +func printDir( + printf func(format string, a ...any), + prefix string, + dir []fs.DirEntry, +) { + names := make([]string, len(dir)) + for i, ent := range dir { + name := ent.Name() + if ent.IsDir() { + name += "/" + } + names[i] = fmt.Sprintf("%q", name) + } + dprintf(printf, "[FAIL] d %s: %s", prefix, strings.Join(names, " ")) +} + +// Compare compares the contents of prefix against the hierarchy described by s. +func (s *FS) Compare( + printf func(format string, a ...any), + prefix string, + e fs.FS, +) error { + if s.Data != nil { + if s.Dir != nil { + panic("invalid state") + } + panic("invalid compare call") + } + + if s.Dir == nil { + dprintf(printf, "[ OK ] s %s", prefix) + return nil + } + + var dir []fs.DirEntry + if d, err := fs.ReadDir(e, prefix); err != nil { + return err + } else if len(d) != len(s.Dir) { + printDir(printf, prefix, d) + return ErrFSBadLength + } else { + dir = d + } + + for _, got := range dir { + name := got.Name() + + if want, ok := s.Dir[name]; !ok { + printDir(printf, prefix, dir) + return fs.ErrNotExist + } else if want.Dir != nil && !got.IsDir() { + printDir(printf, prefix, dir) + return ErrFSInvalidEnt + } else { + name = filepath.Join(prefix, name) + + if fi, err := got.Info(); err != nil { + return err + } else if fi.Mode() != want.Mode { + dprintf(printf, "[FAIL] m %s: %#o, want %#o", + name, uint32(fi.Mode()), uint32(want.Mode)) + return ErrFSBadMode + } + + if want.Data != nil { + if want.Dir != nil { + panic("invalid state") + } + if v, err := fs.ReadFile(e, name); err != nil { + return err + } else if string(v) != *want.Data { + dprintf(printf, + "[FAIL] f %s\n\t got: %s\n\twant: %s", + name, v, *want.Data, + ) + return ErrFSBadData + } + dprintf(printf, "[ OK ] f %s", name) + } else if err := want.Compare(printf, name, e); err != nil { + return err + } + } + } + dprintf(printf, "[ OK ] d %s", prefix) + return nil +} diff --git a/internal/testsuite/fs_test.go b/internal/testsuite/fs_test.go new file mode 100644 index 00000000..4eea866d --- /dev/null +++ b/internal/testsuite/fs_test.go @@ -0,0 +1,85 @@ +package testsuite_test + +import ( + "bytes" + "errors" + "fmt" + "io/fs" + "testing" + "testing/fstest" + + "hakurei.app/internal/testsuite" +) + +func TestCompare(t *testing.T) { + var ( + fsPasswdSample = "u0_a20:x:65534:65534:Hakurei:/var/lib/persist/module/hakurei/u0/a20:/run/current-system/sw/bin/zsh" + fsGroupSample = "hakurei:x:65534:" + ) + + testCases := []struct { + name string + + sample fstest.MapFS + want *testsuite.FS + wantOut string + wantErr error + }{ + {"skip", fstest.MapFS{}, &testsuite.FS{}, "[ OK ] s .\x00", nil}, + {"simple pass", fstest.MapFS{".hakurei": {Mode: 0x800001ed}}, + &testsuite.FS{Dir: map[string]*testsuite.FS{".hakurei": {Mode: 0x800001ed}}}, + "[ OK ] s .hakurei\x00[ OK ] d .\x00", nil}, + {"bad length", fstest.MapFS{".hakurei": {Mode: 0x800001ed}}, + &testsuite.FS{Dir: make(map[string]*testsuite.FS)}, + "[FAIL] d .: \".hakurei/\"\x00", testsuite.ErrFSBadLength}, + {"top level bad mode", fstest.MapFS{".hakurei": {Mode: 0x800001ed}}, + &testsuite.FS{Dir: map[string]*testsuite.FS{".hakurei": {Mode: 0xdeadbeef}}}, + "[FAIL] m .hakurei: 020000000755, want 033653337357\x00", testsuite.ErrFSBadMode}, + {"invalid entry condition", fstest.MapFS{"test": {Data: []byte{'0'}, Mode: 0644}}, + &testsuite.FS{Dir: map[string]*testsuite.FS{"test": {Dir: make(map[string]*testsuite.FS)}}}, + "[FAIL] d .: \"test\"\x00", testsuite.ErrFSInvalidEnt}, + {"nonexistent", fstest.MapFS{"test": {Data: []byte{'0'}, Mode: 0644}}, + &testsuite.FS{Dir: map[string]*testsuite.FS{".test": {}}}, + "[FAIL] d .: \"test\"\x00", fs.ErrNotExist}, + {"file", fstest.MapFS{"etc": {Mode: 0x800001c0}, + "etc/passwd": {Data: []byte(fsPasswdSample), Mode: 0644}, + "etc/group": {Data: []byte(fsGroupSample), Mode: 0644}, + }, &testsuite.FS{Dir: map[string]*testsuite.FS{"etc": {Mode: 0x800001c0, Dir: map[string]*testsuite.FS{ + "passwd": {Mode: 0x1a4, Data: &fsPasswdSample}, + "group": {Mode: 0x1a4, Data: &fsGroupSample}, + }}}}, "[ OK ] f etc/group\x00[ OK ] f etc/passwd\x00[ OK ] d etc\x00[ OK ] d .\x00", nil}, + {"file differ", fstest.MapFS{"etc": {Mode: 0x800001c0}, + "etc/passwd": {Data: []byte(fsPasswdSample), Mode: 0644}, + "etc/group": {Data: []byte(fsGroupSample), Mode: 0644}, + }, &testsuite.FS{Dir: map[string]*testsuite.FS{"etc": {Mode: 0x800001c0, Dir: map[string]*testsuite.FS{ + "passwd": {Mode: 0x1a4, Data: &fsGroupSample}, + "group": {Mode: 0x1a4, Data: &fsGroupSample}, + }}}}, "[ OK ] f etc/group\x00[FAIL] f etc/passwd\n\t got: u0_a20:x:65534:65534:Hakurei:/var/lib/persist/module/hakurei/u0/a20:/run/current-system/sw/bin/zsh\n\twant: hakurei:x:65534:\x00", testsuite.ErrFSBadData}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + var buf bytes.Buffer + + err := tc.want.Compare( + func(format string, a ...any) { + _, _ = fmt.Fprintf(&buf, format+"\x00", a...) + }, + ".", tc.sample, + ) + if !errors.Is(err, tc.wantErr) { + t.Errorf( + "Compare: error = %v; wantErr %v", + err, tc.wantErr, + ) + } + + if buf.String() != tc.wantOut { + t.Errorf( + "Compare: output %q; want %q", + &buf, tc.wantOut, + ) + } + }) + } +} diff --git a/internal/testsuite/mountinfo/mountinfo.go b/internal/testsuite/mountinfo/mountinfo.go new file mode 100644 index 00000000..19892cd2 --- /dev/null +++ b/internal/testsuite/mountinfo/mountinfo.go @@ -0,0 +1,177 @@ +// Package mountinfo provides util-linux bindings for parsing +// proc_pid_mountinfo(5). +// +// This package must never be used outside integration tests, a much better +// implementation can be found in package vfs. +// +// Attempting to import this package outside testing causes the resulting +// program to panic. +package mountinfo + +/* +#cgo linux pkg-config: --static mount + +#include <stdlib.h> +#include <stdio.h> +#include <libmount.h> + +const char *HAKUREI_MOUNTINFO_PATH = "/proc/self/mountinfo"; +*/ +import "C" + +import ( + "errors" + "fmt" + "runtime" + "unsafe" +) + +var ( + // ErrParse is returned by [Open] when encountering a bad record. + ErrParse = errors.New("invalid mountinfo record") + // ErrIter is returned by [Open] if an iterator cannot be allocated. + ErrIter = errors.New("cannot allocate iterator") + // ErrIterAdvance is stored when the iterator is unable to advance. + ErrIterAdvance = errors.New("unable to advance iterator") +) + +type ( + // Iter refers to libmnt iterator state. + Iter struct { + // Last stored error. + err error + // Whether iteration has concluded. + ok bool + // Whether Close had already been called. + closed bool + + tb *C.struct_libmnt_table + itr *C.struct_libmnt_iter + + fs *C.struct_libmnt_fs + } + + // Entry represents deterministic mountinfo parts of a libmnt_fs entry. + Entry struct { + // mount ID: a unique ID for the mount (may be reused after umount(2)). + ID int `json:"id"` + // parent ID: the ID of the parent mount (or of self for the root of + // this mount namespace's mount tree). + Parent int `json:"parent"` + // root: the pathname of the directory in the filesystem which forms the + // root of this mount. + Root string `json:"root"` + // mount point: the pathname of the mount point relative to the + // process's root directory. + Target string `json:"target"` + // mount options: per-mount options (see mount(2)). + VfsOptstr string `json:"vfs_optstr"` + // filesystem type: the filesystem type in the form "type[.subtype]". + FsType string `json:"fstype"` + // mount source: filesystem-specific information or "none". + Source string `json:"source"` + // super options: per-superblock options (see mount(2)). + FsOptstr string `json:"fs_optstr"` + } +) + +// Copy populates v with the current record. +func (m *Iter) Copy(v *Entry) { + if m.fs == nil { + panic("invalid entry") + } + v.ID = int(C.mnt_fs_get_id(m.fs)) + v.Parent = int(C.mnt_fs_get_parent_id(m.fs)) + v.Root = C.GoString(C.mnt_fs_get_root(m.fs)) + v.Target = C.GoString(C.mnt_fs_get_target(m.fs)) + v.VfsOptstr = C.GoString(C.mnt_fs_get_vfs_options(m.fs)) + v.FsType = C.GoString(C.mnt_fs_get_fstype(m.fs)) + v.Source = C.GoString(C.mnt_fs_get_source(m.fs)) + v.FsOptstr = C.GoString(C.mnt_fs_get_fs_options(m.fs)) +} + +// Err returns the saved iterator error. +func (m *Iter) Err() error { return m.err } + +// Open opens a mountinfo document. If name is an empty string, the mountinfo +// document of the current process is opened instead. +func Open(name string) (*Iter, error) { + var m Iter + if name == "" { + m.tb = C.mnt_new_table_from_file(C.HAKUREI_MOUNTINFO_PATH) + } else { + _name := C.CString(name) + m.tb = C.mnt_new_table_from_file(_name) + C.free(unsafe.Pointer(_name)) + } + if m.tb == nil { + return nil, ErrParse + } + m.itr = C.mnt_new_iter(C.MNT_ITER_FORWARD) + if m.itr == nil { + C.mnt_unref_table(m.tb) + return nil, ErrIter + } + m.ok = true + + runtime.SetFinalizer(&m, (*Iter).Close) + return &m, nil +} + +// Close frees the iterator. +func (m *Iter) Close() { + if m.closed { + return + } + if m.tb == nil { + panic("unref called before open") + } + + C.mnt_unref_table(m.tb) + C.mnt_free_iter(m.itr) + m.closed = true + runtime.SetFinalizer(m, nil) +} + +// Reset resets the iterator to the first record for reuse. +func (m *Iter) Reset() { + if m.err != nil { + panic("attempting to reset a faulted iterator") + } + m.ok = true + C.mnt_reset_iter(m.itr, -1) +} + +// Next advances the iterator to the next record. The record may be copied if +// Next returns true. +func (m *Iter) Next() bool { + if !m.ok || m.err != nil { + return false + } + + r := C.mnt_table_next_fs(m.tb, m.itr, &m.fs) + if r < 0 { + m.err = ErrIterAdvance + } + m.ok = r == 0 + return m.ok +} + +// EqualWithIgnore compares e with want, ignoring fields with the specified +// ignore value. +func (e *Entry) EqualWithIgnore(want *Entry, ignore string) bool { + return (e.ID == want.ID || want.ID == -1) && + (e.Parent == want.Parent || want.Parent == -1) && + (e.Root == want.Root || want.Root == ignore) && + (e.Target == want.Target || want.Target == ignore) && + (e.VfsOptstr == want.VfsOptstr || want.VfsOptstr == ignore) && + (e.FsType == want.FsType || want.FsType == ignore) && + (e.Source == want.Source || want.Source == ignore) && + (e.FsOptstr == want.FsOptstr || want.FsOptstr == ignore) +} + +// String returns a text representation of e loosely following the kernel format. +func (e *Entry) String() string { + return fmt.Sprintf("%d %d %s %s %s %s %s %s", + e.ID, e.Parent, e.Root, e.Target, e.VfsOptstr, e.FsType, e.Source, e.FsOptstr) +} diff --git a/internal/testsuite/mountinfo/mountinfo_guard.go b/internal/testsuite/mountinfo/mountinfo_guard.go new file mode 100644 index 00000000..aaf63ac4 --- /dev/null +++ b/internal/testsuite/mountinfo/mountinfo_guard.go @@ -0,0 +1,15 @@ +//go:build !testsuite && !tester + +package mountinfo + +import ( + "os" + "testing" +) + +func init() { + if !testing.Testing() { + println("package mountinfo imported in non-testsuite program") + os.Exit(1) + } +} diff --git a/internal/testsuite/mountinfo/mountinfo_test.go b/internal/testsuite/mountinfo/mountinfo_test.go new file mode 100644 index 00000000..93be3880 --- /dev/null +++ b/internal/testsuite/mountinfo/mountinfo_test.go @@ -0,0 +1,146 @@ +package mountinfo_test + +import ( + "os" + "path/filepath" + "testing" + + "hakurei.app/internal/testsuite/mountinfo" +) + +func TestMountinfo(t *testing.T) { + testCases := []struct { + name string + + sample string + want []*mountinfo.Entry + }{ + {"util-linux", `15 20 0:3 / /proc rw,relatime - proc /proc rw +16 20 0:15 / /sys rw,relatime - sysfs /sys rw +17 20 0:5 / /dev rw,relatime - devtmpfs udev rw,size=1983516k,nr_inodes=495879,mode=755 +18 17 0:10 / /dev/pts rw,relatime - devpts devpts rw,gid=5,mode=620,ptmxmode=000 +19 17 0:16 / /dev/shm rw,relatime - tmpfs tmpfs rw +20 1 8:4 / / rw,noatime - ext3 /dev/sda4 rw,errors=continue,user_xattr,acl,barrier=0,data=ordered +21 16 0:17 / /sys/fs/cgroup rw,nosuid,nodev,noexec,relatime - tmpfs tmpfs rw,mode=755 +22 21 0:18 / /sys/fs/cgroup/systemd rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,release_agent=/lib/systemd/systemd-cgroups-agent,name=systemd +23 21 0:19 / /sys/fs/cgroup/cpuset rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,cpuset +24 21 0:20 / /sys/fs/cgroup/ns rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,ns +25 21 0:21 / /sys/fs/cgroup/cpu rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,cpu +26 21 0:22 / /sys/fs/cgroup/cpuacct rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,cpuacct +27 21 0:23 / /sys/fs/cgroup/memory rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,memory +28 21 0:24 / /sys/fs/cgroup/devices rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,devices +29 21 0:25 / /sys/fs/cgroup/freezer rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,freezer +30 21 0:26 / /sys/fs/cgroup/net_cls rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,net_cls +31 21 0:27 / /sys/fs/cgroup/blkio rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,blkio +32 16 0:28 / /sys/kernel/security rw,relatime - autofs systemd-1 rw,fd=22,pgrp=1,timeout=300,minproto=5,maxproto=5,direct +33 17 0:29 / /dev/hugepages rw,relatime - autofs systemd-1 rw,fd=23,pgrp=1,timeout=300,minproto=5,maxproto=5,direct +34 16 0:30 / /sys/kernel/debug rw,relatime - autofs systemd-1 rw,fd=24,pgrp=1,timeout=300,minproto=5,maxproto=5,direct +35 15 0:31 / /proc/sys/fs/binfmt_misc rw,relatime - autofs systemd-1 rw,fd=25,pgrp=1,timeout=300,minproto=5,maxproto=5,direct +36 17 0:32 / /dev/mqueue rw,relatime - autofs systemd-1 rw,fd=26,pgrp=1,timeout=300,minproto=5,maxproto=5,direct +37 15 0:14 / /proc/bus/usb rw,relatime - usbfs /proc/bus/usb rw +38 33 0:33 / /dev/hugepages rw,relatime - hugetlbfs hugetlbfs rw +39 36 0:12 / /dev/mqueue rw,relatime - mqueue mqueue rw +40 20 8:6 / /boot rw,noatime - ext3 /dev/sda6 rw,errors=continue,barrier=0,data=ordered +41 20 253:0 / /home/kzak rw,noatime - ext4 /dev/mapper/kzak-home rw,barrier=1,data=ordered +42 35 0:34 / /proc/sys/fs/binfmt_misc rw,relatime - binfmt_misc none rw +43 16 0:35 / /sys/fs/fuse/connections rw,relatime - fusectl fusectl rw +44 41 0:36 / /home/kzak/.gvfs rw,nosuid,nodev,relatime - fuse.gvfs-fuse-daemon gvfs-fuse-daemon rw,user_id=500,group_id=500 +45 20 0:37 / /var/lib/nfs/rpc_pipefs rw,relatime - rpc_pipefs sunrpc rw +47 20 0:38 / /mnt/sounds rw,relatime - cifs //foo.home/bar/ rw,unc=\\foo.home\bar,username=kzak,domain=SRGROUP,uid=0,noforceuid,gid=0,noforcegid,addr=192.168.111.1,posixpaths,serverino,acl,rsize=16384,wsize=57344 +49 20 0:56 / /mnt/test/foobar rw,relatime,nosymfollow shared:323 - tmpfs tmpfs rw`, []*mountinfo.Entry{ + e(15, 20, "/", "/proc", "rw,relatime", "proc", "/proc", "rw"), + e(16, 20, "/", "/sys", "rw,relatime", "sysfs", "/sys", "rw"), + e(17, 20, "/", "/dev", "rw,relatime", "devtmpfs", "udev", "rw,size=1983516k,nr_inodes=495879,mode=755"), + e(18, 17, "/", "/dev/pts", "rw,relatime", "devpts", "devpts", "rw,gid=5,mode=620,ptmxmode=000"), + e(19, 17, "/", "/dev/shm", "rw,relatime", "tmpfs", "tmpfs", "rw"), + e(20, 1, "/", "/", "rw,noatime", "ext3", "/dev/sda4", "rw,errors=continue,user_xattr,acl,barrier=0,data=ordered"), + e(21, 16, "/", "/sys/fs/cgroup", "rw,nosuid,nodev,noexec,relatime", "tmpfs", "tmpfs", "rw,mode=755"), + e(22, 21, "/", "/sys/fs/cgroup/systemd", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,release_agent=/lib/systemd/systemd-cgroups-agent,name=systemd"), + e(23, 21, "/", "/sys/fs/cgroup/cpuset", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,cpuset"), + e(24, 21, "/", "/sys/fs/cgroup/ns", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,ns"), + e(25, 21, "/", "/sys/fs/cgroup/cpu", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,cpu"), + e(26, 21, "/", "/sys/fs/cgroup/cpuacct", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,cpuacct"), + e(27, 21, "/", "/sys/fs/cgroup/memory", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,memory"), + e(28, 21, "/", "/sys/fs/cgroup/devices", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,devices"), + e(29, 21, "/", "/sys/fs/cgroup/freezer", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,freezer"), + e(30, 21, "/", "/sys/fs/cgroup/net_cls", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,net_cls"), + e(31, 21, "/", "/sys/fs/cgroup/blkio", "rw,nosuid,nodev,noexec,relatime", "cgroup", "cgroup", "rw,blkio"), + e(32, 16, "/", "/sys/kernel/security", "rw,relatime", "autofs", "systemd-1", "rw,fd=22,pgrp=1,timeout=300,minproto=5,maxproto=5,direct"), + e(33, 17, "/", "/dev/hugepages", "rw,relatime", "autofs", "systemd-1", "rw,fd=23,pgrp=1,timeout=300,minproto=5,maxproto=5,direct"), + e(34, 16, "/", "/sys/kernel/debug", "rw,relatime", "autofs", "systemd-1", "rw,fd=24,pgrp=1,timeout=300,minproto=5,maxproto=5,direct"), + e(35, 15, "/", "/proc/sys/fs/binfmt_misc", "rw,relatime", "autofs", "systemd-1", "rw,fd=25,pgrp=1,timeout=300,minproto=5,maxproto=5,direct"), + e(36, 17, "/", "/dev/mqueue", "rw,relatime", "autofs", "systemd-1", "rw,fd=26,pgrp=1,timeout=300,minproto=5,maxproto=5,direct"), + e(37, 15, "/", "/proc/bus/usb", "rw,relatime", "usbfs", "/proc/bus/usb", "rw"), + e(38, 33, "/", "/dev/hugepages", "rw,relatime", "hugetlbfs", "hugetlbfs", "rw"), + e(39, 36, "/", "/dev/mqueue", "rw,relatime", "mqueue", "mqueue", "rw"), + e(40, 20, "/", "/boot", "rw,noatime", "ext3", "/dev/sda6", "rw,errors=continue,barrier=0,data=ordered"), + e(41, 20, "/", "/home/kzak", "rw,noatime", "ext4", "/dev/mapper/kzak-home", "rw,barrier=1,data=ordered"), + e(42, 35, "/", "/proc/sys/fs/binfmt_misc", "rw,relatime", "binfmt_misc", "none", "rw"), + e(43, 16, "/", "/sys/fs/fuse/connections", "rw,relatime", "fusectl", "fusectl", "rw"), + e(44, 41, "/", "/home/kzak/.gvfs", "rw,nosuid,nodev,relatime", "fuse.gvfs-fuse-daemon", "gvfs-fuse-daemon", "rw,user_id=500,group_id=500"), + e(45, 20, "/", "/var/lib/nfs/rpc_pipefs", "rw,relatime", "rpc_pipefs", "sunrpc", "rw"), + e(47, 20, "/", "/mnt/sounds", "rw,relatime", "cifs", "//foo.home/bar/", "rw,unc=\\\\foo.home\\bar,username=kzak,domain=SRGROUP,uid=0,noforceuid,gid=0,noforcegid,addr=192.168.111.1,posixpaths,serverino,acl,rsize=16384,wsize=57344"), + e(49, 20, "/", "/mnt/test/foobar", "rw,relatime,nosymfollow", "tmpfs", "tmpfs", "rw"), + }}, + } + + for _, tc := range testCases { + name := filepath.Join(t.TempDir(), "sample") + if err := os.WriteFile(name, []byte(tc.sample), 0400); err != nil { + t.Fatalf("cannot write sample: %v", err) + } + + t.Run(tc.name, func(t *testing.T) { + m, err := mountinfo.Open(name) + if err != nil { + t.Fatalf("Open: error = %v", err) + } + t.Cleanup(m.Close) + + i := 0 + var ent mountinfo.Entry + for m.Next() { + m.Copy(&ent) + + if i == len(tc.want) { + t.Errorf("Next: got more than %d entries", i) + t.FailNow() + } + if !ent.EqualWithIgnore(tc.want[i], "\x00") { + t.Errorf("Next: entry %d\n got: %#v\nwant: %#v", i, + ent, &tc.want[i]) + t.FailNow() + } else { + t.Logf("%s", &ent) + } + + i++ + } + + if err = m.Err(); err != nil { + t.Fatalf("Err: %v", err) + } + }) + + if err := os.Remove(name); err != nil { + t.Fatalf("cannot remove %q: %v", name, err) + } + } +} + +func e( + id, parent int, + root, target, vfsOptstr string, + fsType, source, fsOptstr string, +) *mountinfo.Entry { + return &mountinfo.Entry{ + ID: id, + Parent: parent, + Root: root, + Target: target, + VfsOptstr: vfsOptstr, + FsType: fsType, + Source: source, + FsOptstr: fsOptstr, + } +} diff --git a/internal/testsuite/proc.go b/internal/testsuite/proc.go new file mode 100644 index 00000000..e7ef1aed --- /dev/null +++ b/internal/testsuite/proc.go @@ -0,0 +1,353 @@ +package testsuite + +import ( + "bytes" + "errors" + "fmt" + "os" + "path/filepath" + "strconv" + "strings" + "syscall" + "unsafe" + + "hakurei.app/fhs" +) + +// Stat represents status information read from /proc/pid/stat. +type Stat struct { + // The process ID. + PID int + // The filename of the executable, with parenthesis stripped. + Comm string + // One of the following characters, indicating process state: + // + // R Running + // + // S Sleeping in an interruptible wait + // + // D Waiting in uninterruptible disk sleep + // + // Z Zombie + // + // T Stopped (on a signal) or (before Linux + // 2.6.33) trace stopped + // + // t Tracing stop (Linux 2.6.33 onward) + // + // W Paging (only before Linux 2.6.0) + // + // X Dead (from Linux 2.6.0 onward) + // + // x Dead (Linux 2.6.33 to 3.13 only) + // + // K Wakekill (Linux 2.6.33 to 3.13 only) + // + // W Waking (Linux 2.6.33 to 3.13 only) + // + // P Parked (Linux 3.9 to 3.13 only) + // + // I Idle (Linux 4.14 onward) + State byte + // The process ID of the parent of this process. + PPID int + // The process group ID of the process. + PGRP int + // The session ID of the process. + Session int + // The controlling terminal of the process. + TTYNR int + // The ID of the foreground process group of the controlling terminal of the + // process. + TPGID int + // The kernel flags word of the process. For bit meanings, see the PF_* + // defines in the Linux kernel source file include/linux/sched.h. + Flags uint + // The number of minor faults the process has made which have not required + // loading a memory page from disk. + MinFlt uint + // The number of minor faults that the process's waited-for children have + // made. + CMinFlt uint + // The number of major faults the process has made which have required + // loading a memory page from disk. + MajFlt uint + // The number of major faults that the process's waited-for children have + // made. + CMajFlt uint + // Amount of time that this process has been scheduled in user mode, + // measured in clock ticks. + UTime uint + // Amount of time that this process has been scheduled in kernel mode, + // measured in clock ticks. + STime uint + // Amount of time that this process's waited-for children have been + // scheduled in user mode, measured in clock ticks. + CUTime int + // Amount of time that this process's waited-for children have been + // scheduled in kernel mode, measured in clock ticks. + CSTime int + // For processes running a real-time scheduling policy, this is the negated + // scheduling priority, minus one. + Priority int + // The nice value, a value in the range 19 (low priority) to -20 (high + // priority). + Nice int + // Number of threads in this process. + NumThreads int + + // unmaintained field: itrealvalue + + // The time the process started after system boot. Since Linux 2.6, the + // value is expressed in clock ticks. + StartTime uint64 + // Virtual memory size in bytes. + VSize uint + // Resident set size in pages. + RSS int + // Soft limit in bytes on the rss of the process. + RSSLim uint64 + // The address above which program text can run. + StartCode uint64 + // The address below which program text can run. + EndCode uint64 + // The address of the start (i.e., bottom) of the stack. + StartStack uint64 + // The current value of ESP (stack pointer), as found in the kernel stack + // page for the process. + KSTKESP uint64 + // The current EIP (instruction pointer). + KSTKEIP uint64 + + // obsolete fields: signal, blocked, sigignore, sigcatch + + // This is the "channel" in which the process is waiting. It is the address + // of a location in the kernel where the process is sleeping. + WChan uint64 + + // unmaintained fields: nswap, cnswap + + // Signal to be sent to parent when we die. + ExitSignal int + // CPU number last executed on. + Processor int + // Real-time scheduling priority, a number in the range 1 to 99 for processes + // scheduled under a real-time policy, or 0, for non-real-time processes. + RTPriority uint + // Scheduling policy (see sched_setscheduler(2)). Decode using the SCHED_* + // constants in linux/sched.h. + Policy uint + // Aggregated block I/O delays, measured in clock ticks (centiseconds). + DelayAcctBlkIOTicks uint64 + // Guest time of the process (time spent running a virtual CPU for a guest + // operating system), measured in clock ticks. + GuestTime int + // Guest time of the process's children, measured in clock ticks. + CGuestTime int +} + +// 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(fhs.Proc, strconv.Itoa(s.PID), "exe")) + + // When the executable has been deleted then Readlink returns a + // path appended with " (deleted)". + return strings.TrimSuffix(path, " (deleted)"), err +} + +// 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(fhs.Proc, strconv.Itoa(s.PID)), stat) + if err != nil { + err = os.NewSyscallError("stat", err) + } + return +} + +// Args reads arguments of the process referred to by s. +func (s *Stat) Args() ([]string, error) { + p, err := os.ReadFile(filepath.Join(fhs.Proc, strconv.Itoa(s.PID), "cmdline")) + if err != nil { + return nil, err + } + a := bytes.Split(p, []byte{0}) + if len(a) > 0 && len(a[len(a)-1]) == 0 { + a = a[:len(a)-1] + } + + args := make([]string, len(a)) + for i, arg := range a { + args[i] = unsafe.String(unsafe.SliceData(arg), len(arg)) + } + return args, nil +} + +// ErrBadDelimiters is returned by [Stat.UnmarshalText] if one or both bytes of +// the comm delimiter pair were missing or misplaced. +var ErrBadDelimiters = errors.New("missing comm delimiters") + +// UnmarshalText populates the structure pointed to by s from text. +func (s *Stat) UnmarshalText(text []byte) (err error) { + var ( + discard uint64 + _uint64 = &discard + _int64 = (*int64)(unsafe.Pointer(&discard)) + + ld = bytes.Index(text, []byte("(")) + rd = bytes.LastIndex(text, []byte(")")) + ) + + if ld <= 0 || rd < 0 { + return ErrBadDelimiters + } + + if s.PID, err = strconv.Atoi( + unsafe.String(unsafe.SliceData(text), ld-1), + ); err != nil { + return + } + + s.Comm = string(text[ld+1 : rd]) + + var ( + n int + + state string + ) + n, err = fmt.Fscan( + bytes.NewBuffer(text[rd+2:]), + &state, + &s.PPID, + &s.PGRP, + &s.Session, + &s.TTYNR, + &s.TPGID, + &s.Flags, + &s.MinFlt, + &s.CMinFlt, + &s.MajFlt, + &s.CMajFlt, + &s.UTime, + &s.STime, + &s.CUTime, + &s.CSTime, + &s.Priority, + &s.Nice, + &s.NumThreads, + _int64, + &s.StartTime, + &s.VSize, + &s.RSS, + &s.RSSLim, + &s.StartCode, + &s.EndCode, + &s.StartStack, + &s.KSTKESP, + &s.KSTKEIP, + _uint64, + _uint64, + _uint64, + _uint64, + &s.WChan, + _uint64, + _uint64, + &s.ExitSignal, + &s.Processor, + &s.RTPriority, + &s.Policy, + &s.DelayAcctBlkIOTicks, + &s.GuestTime, + &s.CGuestTime, + ) + if err != nil { + err = fmt.Errorf("field %d: %w", n, err) + } else if len(state) != 1 { + err = fmt.Errorf("invalid state %q", state) + } else { + s.State = state[0] + } + return +} + +// A StatScanner continuously scans the proc filesystem for process status +// information in /proc/pid/stat. +type StatScanner struct { + // Current entry. + stat Stat + // Cached top-level /proc entries. + dents []os.DirEntry + // Current progress through dents. + i int + // Whether the previous call to Scan had repopulated dents. + wrapped bool + // First stored error: a non-nil err disables the scanner. + 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. +func (s *StatScanner) Scan() bool { + if s.err != nil { + return false + } + + if s.wrapped = s.i == len(s.dents); s.wrapped { + if s.dents, s.err = os.ReadDir(fhs.Proc); s.err != nil { + return false + } + s.i = 0 + if len(s.dents) == 0 { + s.err = syscall.ENOTRECOVERABLE + return false + } + } + + for s.i < len(s.dents) { + dent := s.dents[s.i] + s.i++ + if !dent.IsDir() { + continue + } + + pid, err := strconv.Atoi(dent.Name()) + if err != nil { + continue + } + + var p []byte + p, err = os.ReadFile(filepath.Join(fhs.Proc, dent.Name(), "stat")) + if err != nil { + if IsNotExist(err) { + continue + } + s.err = err + return false + } + + s.err = s.stat.UnmarshalText(p) + if s.err == nil && pid != s.stat.PID { + s.err = fmt.Errorf( + "bad status information: dent=%d, stat=%d", + pid, s.stat.PID, + ) + } + return s.err == nil + } + return s.Scan() +} + +// Stat returns the address of the [Stat] structure populated by the last call +// to Scan. +func (s *StatScanner) Stat() *Stat { return &s.stat } + +// Err returns the stored error value. +func (s *StatScanner) Err() error { return s.err } + +// Repopulated returns whether the last Scan call had re-read the proc filesystem. +func (s *StatScanner) Repopulated() bool { return s.wrapped } diff --git a/internal/testsuite/proc_test.go b/internal/testsuite/proc_test.go new file mode 100644 index 00000000..73693972 --- /dev/null +++ b/internal/testsuite/proc_test.go @@ -0,0 +1,33 @@ +package testsuite_test + +import ( + "testing" + + "hakurei.app/internal/testsuite" +) + +func BenchmarkStatScanner(b *testing.B) { + var s testsuite.StatScanner + + for b.Loop() { + if !s.Scan() { + b.Fatal(s.Err()) + } + } +} + +func BenchmarkStatScannerFull(b *testing.B) { + var s testsuite.StatScanner + + for b.Loop() { + for s.Scan() { + if s.Repopulated() { + break + } + } + + if err := s.Err(); err != nil { + b.Fatal(err) + } + } +} diff --git a/internal/testsuite/ptrace.go b/internal/testsuite/ptrace.go new file mode 100644 index 00000000..ccf0900c --- /dev/null +++ b/internal/testsuite/ptrace.go @@ -0,0 +1,144 @@ +package testsuite + +import ( + "crypto/sha512" + "encoding/base64" + "errors" + "fmt" + "os" + "syscall" + "unsafe" +) + +const ( + // _PTRACE_ATTACH attaches to the process specified in pid. + _PTRACE_ATTACH = 16 + // _PTRACE_DETACH restarts the stopped tracee as for PTRACE_CONT, but first + // detaches from it. + _PTRACE_DETACH = 17 + + // _PTRACE_SECCOMP_GET_FILTER allows the tracer to dump the tracee's classic + // BPF filters. + _PTRACE_SECCOMP_GET_FILTER = 0x420c +) + +// ptrace wraps the ptrace syscall. +func ptrace( + op uintptr, + pid, addr int, + data unsafe.Pointer, +) (r uintptr, errno syscall.Errno) { + r, _, errno = syscall.Syscall6( + syscall.SYS_PTRACE, + op, + uintptr(pid), + uintptr(addr), + uintptr(data), + 0, 0, + ) + return +} + +// ptraceAttach attaches to the process referred to by pid. +func ptraceAttach(pid int) error { + if _, errno := ptrace(_PTRACE_ATTACH, pid, 0, nil); errno != 0 { + return os.NewSyscallError("PTRACE_ATTACH", errno) + } + + var status syscall.WaitStatus + for { + if _, err := syscall.Wait4( + pid, + &status, + syscall.WALL, + nil, + ); err != nil { + if errors.Is(err, syscall.EINTR) { + continue + } + return os.NewSyscallError("wait4", err) + } + switch { + case status.Stopped(): + 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. +func ptraceDetach(pid int) error { + if _, errno := ptrace(_PTRACE_DETACH, pid, 0, nil); errno != 0 { + return os.NewSyscallError("PTRACE_DETACH", errno) + } + return nil +} + +// getFilter dumps the specified tracee's cBPF filter at the specified index +// and returns the resulting payload. T must be eight bytes long and must not +// contain pointers. +func getFilter(pid, index int) ([]syscall.SockFilter, error) { + var buf []syscall.SockFilter + if n, errno := ptrace( + _PTRACE_SECCOMP_GET_FILTER, + pid, index, nil, + ); errno != 0 { + return nil, os.NewSyscallError("PTRACE_SECCOMP_GET_FILTER", errno) + } else { + buf = make([]syscall.SockFilter, n) + } + if _, errno := ptrace( + _PTRACE_SECCOMP_GET_FILTER, + pid, index, unsafe.Pointer(&buf[0]), + ); errno != 0 { + return nil, os.NewSyscallError("PTRACE_SECCOMP_GET_FILTER", errno) + } + return buf, nil +} + +// CheckFilter checks the process at pid to have its first filter's contents +// match the specified sha512 checksum. +func CheckFilter(pid, index int, sum [sha512.Size]byte) (err error) { + if err = ptraceAttach(pid); err != nil { + return + } + defer func() { + if detachErr := ptraceDetach(pid); err == nil { + err = detachErr + } + }() + + var buf []syscall.SockFilter + h := sha512.New() + if buf, err = getFilter(pid, index); err != nil { + return + } else { + h.Write(unsafe.Slice( + (*byte)(unsafe.Pointer(&buf[0])), + uintptr(len(buf))*unsafe.Sizeof(buf[0]), + )) + } + + if got := h.Sum(nil); string(got) != string(sum[:]) { + return fmt.Errorf( + "bad filter\n\t got: %s\n\twant: %s", + base64.StdEncoding.EncodeToString(got), + base64.StdEncoding.EncodeToString(sum[:]), + ) + } + return +} diff --git a/internal/testsuite/ptrace_test.go b/internal/testsuite/ptrace_test.go new file mode 100644 index 00000000..18ad3f31 --- /dev/null +++ b/internal/testsuite/ptrace_test.go @@ -0,0 +1,13 @@ +package testsuite_test + +import ( + "syscall" + "testing" + "unsafe" +) + +func TestBlockSize(t *testing.T) { + if sz := unsafe.Sizeof(syscall.SockFilter{}); sz != 8 { + t.Fatalf("invalid filter block size %d", sz) + } +} diff --git a/internal/testsuite/testsuite.go b/internal/testsuite/testsuite.go new file mode 100644 index 00000000..456c36b4 --- /dev/null +++ b/internal/testsuite/testsuite.go @@ -0,0 +1,290 @@ +// Package testsuite provides many quick-and-dirty integration testing utilities. +// +// Attempting to import this package outside testing causes the resulting +// program to panic. +package testsuite + +import ( + "bufio" + "context" + "crypto/sha512" + "errors" + "log" + "os" + "os/exec" + "os/signal" + "sync" + "syscall" + "time" +) + +// ReceiveSignals blocks until a termination signal arrives, and terminates. +func ReceiveSignals() { + s := make(chan os.Signal, 3) + signal.Notify(s, os.Interrupt, syscall.SIGTERM, syscall.SIGHUP) + log.Fatalf("terminating on signal %s", <-s) +} + +// MustRun runs command and terminates the testsuite on error. +func MustRun(cred *syscall.Credential, extraEnv []string, command ...string) { + cmd := exec.Command(command[0], command[1:]...) + cmd.Stdout, cmd.Stderr = os.Stdout, os.Stderr + cmd.SysProcAttr = &syscall.SysProcAttr{ + Pdeathsig: syscall.SIGKILL, + Credential: cred, + } + if len(extraEnv) != 0 { + cmd.Env = append(cmd.Environ(), extraEnv...) + } + if err := cmd.Run(); err != nil { + log.Fatal(err) + } +} + +// 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(cred *syscall.Credential, extraEnv []string, command ...string) { + cmd := exec.Command(command[0], command[1:]...) + cmd.Stdout, cmd.Stderr = os.Stdout, os.Stderr + cmd.SysProcAttr = &syscall.SysProcAttr{ + Pdeathsig: syscall.SIGKILL, + Credential: cred, + } + if len(extraEnv) != 0 { + cmd.Env = append(cmd.Environ(), extraEnv...) + } + 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) + } +} + +// 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 +} + +// MustStartWith wraps [MustStart] and creates the [exec.Cmd] object internally. +func MustStartWith( + ctx context.Context, + cred *syscall.Credential, + extraEnv []string, + files []*os.File, + command ...string, +) (proc *os.Process, done <-chan error) { + cmd := exec.CommandContext(ctx, command[0], command[1:]...) + cmd.Stdout, cmd.Stderr = os.Stdout, os.Stderr + cmd.ExtraFiles = files + cmd.SysProcAttr = &syscall.SysProcAttr{ + Pdeathsig: syscall.SIGTERM, + Credential: cred, + } + if len(extraEnv) != 0 { + cmd.Env = append(cmd.Environ(), extraEnv...) + } + 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, + cred *syscall.Credential, + extraEnv []string, + command ...string, +) { + for range time.NewTicker(d).C { + cmd := exec.Command(command[0], command[1:]...) + cmd.SysProcAttr = &syscall.SysProcAttr{ + Pdeathsig: syscall.SIGKILL, + Credential: cred, + } + if len(extraEnv) != 0 { + cmd.Env = append(cmd.Environ(), extraEnv...) + } + 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(cred *syscall.Credential) (dbusEnv string) { + r, w, err := os.Pipe() + if err != nil { + log.Fatal(err) + } + + // this is never explicitly terminated + _, done := MustStartWith( + context.Background(), cred, nil, []*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, + cred *syscall.Credential, + dbusEnv string, +) { + wg.Go(func() { + // this is terminated via swaymsg + _, done := MustStartWith( + context.Background(), cred, []string{ + "WLR_BACKENDS=headless", + XDGRuntimeEnv, + SwayEnv, + dbusEnv, + }, nil, + "sway", + ) + if err := <-done; err != nil { + log.Fatal(err) + } + }) + + Poll( + 50*time.Millisecond, + cred, + []string{SwayEnv}, + "swaymsg", + ) + log.Printf("sway available via %s", SwayEnv) +} + +// TerminateSway requests for the sway server to terminate via sway IPC. +func TerminateSway(cred *syscall.Credential) { + MustFail(cred, []string{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(cred *syscall.Credential, dbusEnv string) { + // this is never explicitly terminated + _, done := MustStartWith( + context.Background(), cred, []string{ + XDGRuntimeEnv, + dbusEnv, + }, nil, + "pipewire", + ) + + go func() { + if _err := <-done; _err != nil { + log.Fatal(_err) + } + log.Fatal("pipewire terminated unexpectedly") + }() + + Poll(50*time.Millisecond, cred, []string{ + XDGRuntimeEnv, + dbusEnv, + }, + "wpctl", + "status", + ) + + _, _done := MustStartWith( + context.Background(), cred, []string{ + XDGRuntimeEnv, + dbusEnv, + }, nil, + "wireplumber", + ) + + go func() { + if _err := <-_done; _err != nil { + log.Fatal(_err) + } + log.Fatal("wireplumber terminated unexpectedly") + }() +} diff --git a/internal/testsuite/testsuite_guard.go b/internal/testsuite/testsuite_guard.go new file mode 100644 index 00000000..0859d17f --- /dev/null +++ b/internal/testsuite/testsuite_guard.go @@ -0,0 +1,15 @@ +//go:build !testsuite && !tester + +package testsuite + +import ( + "os" + "testing" +) + +func init() { + if !testing.Testing() { + println("package testsuite imported in non-testsuite program") + os.Exit(1) + } +} diff --git a/internal/testsuite/testsuite_root.go b/internal/testsuite/testsuite_root.go new file mode 100644 index 00000000..11f503ef --- /dev/null +++ b/internal/testsuite/testsuite_root.go @@ -0,0 +1,17 @@ +//go:build testsuite + +package testsuite + +import ( + "log" + "os" +) + +func init() { + if os.Geteuid() != 0 { + log.Fatal("this program must run as root") + } + + log.SetFlags(0) + log.SetPrefix("testsuite: ") +} |
