diff options
| -rw-r--r-- | cmd/hakurei/testsuite/sandbox/main.go | 75 | ||||
| -rw-r--r-- | container/dispatcher.go | 16 | ||||
| -rw-r--r-- | container/errors.go | 27 |
3 files changed, 107 insertions, 11 deletions
diff --git a/cmd/hakurei/testsuite/sandbox/main.go b/cmd/hakurei/testsuite/sandbox/main.go index 820b7e2f..a2d1cb5e 100644 --- a/cmd/hakurei/testsuite/sandbox/main.go +++ b/cmd/hakurei/testsuite/sandbox/main.go @@ -12,6 +12,7 @@ import ( "log" "os" "os/exec" + "os/signal" "path/filepath" "slices" "strconv" @@ -19,6 +20,7 @@ import ( "sync" "sync/atomic" "syscall" + "time" "hakurei.app/check" "hakurei.app/fhs" @@ -49,7 +51,7 @@ func mustScanFor(f func(ps *testsuite.StatScanner) bool) int { // the container. This process must be terminated by the caller. func mustStart( ctx context.Context, - serial uint64, + serial uint64, identity int, cred *syscall.Credential, files ...*os.File, ) (pid int, done <-chan error) { @@ -57,6 +59,7 @@ func mustStart( _, done = testsuite.MustStartWith( ctx, cred, nil, files, "hakurei", "exec", + "-a", strconv.Itoa(identity), "sleep", "infinity", _serial, ) @@ -106,13 +109,74 @@ func mustStart( return } +// serial is the previous value returned by newSerial. +var serial atomic.Uint64 + +// newSerial returns a unique number. +func newSerial() uint64 { defer serial.Add(1); return serial.Load() } + func main() { + cred := syscall.Credential{Uid: 1000, Gid: 100} + + if p, ok := os.LookupEnv("HAKUREI_TESTSUITE_EXERCISE"); ok { + n, err := strconv.Atoi(p) + if err != nil { + log.Fatalf("invalid stress concurrency %q", p) + } + + ctx, stop := signal.NotifyContext(context.Background(), + os.Interrupt, + syscall.SIGTERM, + ) + defer stop() + + const identCount = 8 + + var wg sync.WaitGroup + in := identCount + var identity int + + wg.Add(n) + for range n { + in-- + if in < 0 { + identity++ + in = identCount + } + go func(identity int) { + defer wg.Done() + + t := time.NewTicker(500 * time.Millisecond) + start: + pid, done := mustStart(ctx, newSerial(), identity, &cred) + select { + case <-t.C: + if _err := syscall.Kill(pid, syscall.SIGTERM); _err != nil { + log.Fatal(_err) + } + + case <-ctx.Done(): + if _err := syscall.Kill(pid, syscall.SIGTERM); _err != nil { + log.Fatal(_err) + } + } + if _err := testsuite.FilterTerminated(<-done); _err != nil { + log.Fatal(_err) + } + if ctx.Err() == nil { + goto start + } + }(identity) + } + log.Printf("exercising cmd/hakurei with %d concurrent instances", n) + wg.Wait() + } + go testsuite.ReceiveSignals() // the signal handler does not wait for termination ctx := context.Background() - cred := syscall.Credential{Uid: 1000, Gid: 100} if err := os.MkdirAll("/opt/test-helper/bin", 0755); err != nil { log.Fatal(err) } @@ -133,9 +197,6 @@ func main() { var wg sync.WaitGroup defer wg.Wait() - var serial atomic.Uint64 - newSerial := func() uint64 { serial.Add(1); return serial.Load() } - testsuite.MustRun( &cred, nil, "hakurei", "exec", "capsh", "--print", @@ -166,7 +227,7 @@ func main() { c, cancel := context.WithCancel(ctx) defer cancel() - pid, done := mustStart(c, newSerial(), &cred) + pid, done := mustStart(c, newSerial(), 0, &cred) testsuite.MustCheckFilter(pid, testdata.SumPD) if err := testsuite.FilterTerminated(<-done); err != nil { log.Fatal(err) @@ -179,7 +240,7 @@ func main() { c, cancel := context.WithCancel(ctx) defer cancel() - pid, done := mustStart(c, newSerial(), &cred, os.Stdin, os.Stdout, os.Stderr) + pid, done := mustStart(c, newSerial(), 0, &cred, os.Stdin, os.Stdout, os.Stderr) prefix := filepath.Join(fhs.Proc, strconv.Itoa(pid), "fd") var fail bool diff --git a/container/dispatcher.go b/container/dispatcher.go index 4636a469..1bad7527 100644 --- a/container/dispatcher.go +++ b/container/dispatcher.go @@ -228,11 +228,19 @@ func (direct) seccompLoad(rules []std.NativeRule, flags seccomp.ExportFlag) erro func (direct) notify(c chan<- os.Signal, sig ...os.Signal) { signal.Notify(c, sig...) } func (direct) start(c *exec.Cmd) error { return c.Start() } func (direct) signal(c *exec.Cmd, sig os.Signal) error { return c.Process.Signal(sig) } -func (direct) evalSymlinks(path string) (string, error) { return filepath.EvalSymlinks(path) } +func (direct) evalSymlinks(path string) (string, error) { + return retryExp(func() (string, error) { + return filepath.EvalSymlinks(path) + }, syscall.EACCES) +} -func (direct) exit(code int) { os.Exit(code) } -func (direct) getpid() int { return os.Getpid() } -func (direct) stat(name string) (os.FileInfo, error) { return os.Stat(name) } +func (direct) exit(code int) { os.Exit(code) } +func (direct) getpid() int { return os.Getpid() } +func (direct) stat(name string) (os.FileInfo, error) { + return retryExp(func() (os.FileInfo, error) { + return os.Stat(name) + }, syscall.EACCES) +} func (direct) mkdir(name string, perm os.FileMode) error { return os.Mkdir(name, perm) } func (direct) mkdirTemp(dir, pattern string) (string, error) { return os.MkdirTemp(dir, pattern) } func (direct) mkdirAll(path string, perm os.FileMode) error { return os.MkdirAll(path, perm) } diff --git a/container/errors.go b/container/errors.go index 053bdeb7..8975dec4 100644 --- a/container/errors.go +++ b/container/errors.go @@ -4,12 +4,39 @@ import ( "errors" "os" "syscall" + "time" "hakurei.app/check" "hakurei.app/message" "hakurei.app/vfs" ) +// retryExp retries f if the resulting error is equivalent to e according to +// [errors.Is], up to 8 times. +func retryExp[T any](f func() (T, error), e error) (v T, err error) { + const statRetry = 8 // 64 ms + var n int + d := 256 * time.Microsecond + +retry: + v, err = f() + if err != nil { + if !errors.Is(err, e) { + return + } + + n++ + if n == statRetry { + return + } + + time.Sleep(d) + d *= 2 + goto retry + } + return +} + // messageFromError returns a printable error message for a supported concrete type. func messageFromError(err error) (m string, ok bool) { if m, ok = messagePrefixP[MountError]("cannot ", err); ok { |
