diff options
Diffstat (limited to 'internal/testsuite')
| -rw-r--r-- | internal/testsuite/testsuite.go | 81 |
1 files changed, 61 insertions, 20 deletions
diff --git a/internal/testsuite/testsuite.go b/internal/testsuite/testsuite.go index 456c36b4..6e48eca1 100644 --- a/internal/testsuite/testsuite.go +++ b/internal/testsuite/testsuite.go @@ -13,6 +13,9 @@ import ( "os" "os/exec" "os/signal" + "os/user" + "strconv" + "strings" "sync" "syscall" "time" @@ -25,6 +28,52 @@ func ReceiveSignals() { log.Fatalf("terminating on signal %s", <-s) } +var ( + // users caches [user.LookupId] calls. + users = make(map[string]*user.User) + // usersMu synchronises access to users. + usersMu sync.RWMutex +) + +// mustLookupId is like [user.LookupId], but the first result is cached. +func mustLookupId(uid string) *user.User { + usersMu.RLock() + v, ok := users[uid] + usersMu.RUnlock() + if ok { + return v + } + + u, err := user.LookupId(uid) + if err != nil { + log.Fatal(err) + } + + usersMu.Lock() + users[uid] = u + usersMu.Unlock() + return u +} + +// MustAppendEnv adds extra environment variables to cmd. +func MustAppendEnv(cmd *exec.Cmd, env ...string) { + if len(cmd.Env) == 0 { + var cred syscall.Credential + if cmd.SysProcAttr != nil && cmd.SysProcAttr.Credential != nil { + cred = *cmd.SysProcAttr.Credential + } + u := mustLookupId(strconv.Itoa(int(cred.Uid))) + + cmd.Env = append(cmd.Env, + "PATH="+os.Getenv("PATH"), + "HOME="+u.HomeDir, + "USER="+u.Username, + "USERNAME="+u.Username, + ) + } + cmd.Env = append(cmd.Env, env...) +} + // 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:]...) @@ -33,11 +82,9 @@ func MustRun(cred *syscall.Credential, extraEnv []string, command ...string) { Pdeathsig: syscall.SIGKILL, Credential: cred, } - if len(extraEnv) != 0 { - cmd.Env = append(cmd.Environ(), extraEnv...) - } + MustAppendEnv(cmd, extraEnv...) if err := cmd.Run(); err != nil { - log.Fatal(err) + log.Fatalf("must run %s: %v", strings.Join(command, " "), err) } } @@ -54,15 +101,13 @@ func MustFail(cred *syscall.Credential, extraEnv []string, command ...string) { Pdeathsig: syscall.SIGKILL, Credential: cred, } - if len(extraEnv) != 0 { - cmd.Env = append(cmd.Environ(), extraEnv...) - } + MustAppendEnv(cmd, extraEnv...) if err := cmd.Run(); err == nil { log.Fatal(ErrUnexpectedSuccess) } else if e, ok := errors.AsType[*exec.ExitError](err); !ok { - log.Fatal(err) + log.Fatalf("must fail %s: %v", strings.Join(command, " "), err) } else if !e.Exited() { - log.Fatal(e) + log.Fatalf("must fail %s: %v", strings.Join(command, " "), e) } } @@ -91,9 +136,7 @@ func MustStartWith( Pdeathsig: syscall.SIGTERM, Credential: cred, } - if len(extraEnv) != 0 { - cmd.Env = append(cmd.Environ(), extraEnv...) - } + MustAppendEnv(cmd, extraEnv...) return cmd.Process, MustStart(cmd) } @@ -140,14 +183,12 @@ func Poll( Pdeathsig: syscall.SIGKILL, Credential: cred, } - if len(extraEnv) != 0 { - cmd.Env = append(cmd.Environ(), extraEnv...) - } + MustAppendEnv(cmd, extraEnv...) if err := cmd.Run(); err != nil { if e, ok := errors.AsType[*exec.ExitError](err); ok && e.Exited() { continue } - log.Fatal(err) + log.Fatalf("poll %s: %v", strings.Join(command, " "), err) } break } @@ -183,7 +224,7 @@ func MustStartSessionBus(cred *syscall.Credential) (dbusEnv string) { go func() { if _err := <-done; _err != nil { - log.Fatal(_err) + log.Fatalf("session bus terminated unexpectedly: %v", _err) } log.Fatal("session bus terminated unexpectedly") }() @@ -228,7 +269,7 @@ func MustStartSway( "sway", ) if err := <-done; err != nil { - log.Fatal(err) + log.Fatalf("sway terminated unexpectedly: %v", err) } }) @@ -260,7 +301,7 @@ func MustStartPipeWire(cred *syscall.Credential, dbusEnv string) { go func() { if _err := <-done; _err != nil { - log.Fatal(_err) + log.Fatalf("pipewire terminated unexpectedly: %v", _err) } log.Fatal("pipewire terminated unexpectedly") }() @@ -283,7 +324,7 @@ func MustStartPipeWire(cred *syscall.Credential, dbusEnv string) { go func() { if _err := <-_done; _err != nil { - log.Fatal(_err) + log.Fatalf("wireplumber terminated unexpectedly: %v", _err) } log.Fatal("wireplumber terminated unexpectedly") }() |
