aboutsummaryrefslogtreecommitdiffhomepage
path: root/container/container_test.go
diff options
context:
space:
mode:
authorOphestra <cat@gensokyo.uk>2025-07-28 01:44:31 +0900
committerOphestra <cat@gensokyo.uk>2025-07-28 01:45:38 +0900
commitd6b07f12ff5064682b722f77bbb4436451bfcdab (patch)
tree9eee3b43af5bcb3b68ac6872e3486e176a04465b /container/container_test.go
parent65fe09caf9bd8cbed9177022c765423885aed5be (diff)
container: forward context cancellation
This allows container processes to exit gracefully. Signed-off-by: Ophestra <cat@gensokyo.uk>
Diffstat (limited to 'container/container_test.go')
-rw-r--r--container/container_test.go104
1 files changed, 73 insertions, 31 deletions
diff --git a/container/container_test.go b/container/container_test.go
index 8e59e4eb..190fa5b5 100644
--- a/container/container_test.go
+++ b/container/container_test.go
@@ -8,6 +8,8 @@ import (
"fmt"
"log"
"os"
+ "os/exec"
+ "os/signal"
"strconv"
"strings"
"syscall"
@@ -90,41 +92,32 @@ func TestContainer(t *testing.T) {
t.Cleanup(func() { container.SetOutput(oldOutput) })
}
- t.Run("cancel", func(t *testing.T) {
- ctx, cancel := context.WithTimeout(t.Context(), helperDefaultTimeout)
-
- c := helperNewContainer(ctx, "block")
- c.Stdout, c.Stderr = os.Stdout, os.Stderr
- c.WaitDelay = helperDefaultTimeout
-
- ready := make(chan struct{})
- if r, w, err := os.Pipe(); err != nil {
- t.Fatalf("cannot pipe: %v", err)
- } else {
- c.ExtraFiles = append(c.ExtraFiles, w)
- go func() {
- defer close(ready)
- if _, err = r.Read(make([]byte, 1)); err != nil {
- panic(err.Error())
- }
- }()
- }
-
- if err := c.Start(); err != nil {
- hlog.PrintBaseError(err, "start:")
- t.Fatalf("cannot start container: %v", err)
- } else if err = c.Serve(); err != nil {
- hlog.PrintBaseError(err, "serve:")
- t.Errorf("cannot serve setup params: %v", err)
- }
- <-ready
- cancel()
+ t.Run("cancel", testContainerCancel(nil, func(t *testing.T, c *container.Container) {
wantErr := context.Canceled
+ wantExitCode := 0
if err := c.Wait(); !errors.Is(err, wantErr) {
hlog.PrintBaseError(err, "wait:")
- t.Fatalf("Wait: error = %v, want %v", err, wantErr)
+ t.Errorf("Wait: error = %v, want %v", err, wantErr)
}
- })
+ if ps := c.ProcessState(); ps == nil {
+ t.Errorf("ProcessState unexpectedly returned nil")
+ } else if code := ps.ExitCode(); code != wantExitCode {
+ t.Errorf("ExitCode: %d, want %d", code, wantExitCode)
+ }
+ }))
+
+ t.Run("forward", testContainerCancel(func(c *container.Container) {
+ c.ForwardCancel = true
+ }, func(t *testing.T, c *container.Container) {
+ var exitError *exec.ExitError
+ if err := c.Wait(); !errors.As(err, &exitError) {
+ hlog.PrintBaseError(err, "wait:")
+ t.Errorf("Wait: error = %v", err)
+ }
+ if code := exitError.ExitCode(); code != blockExitCodeInterrupt {
+ t.Errorf("ExitCode: %d, want %d", code, blockExitCodeInterrupt)
+ }
+ }))
for i, tc := range containerTestCases {
t.Run(tc.name, func(t *testing.T) {
@@ -214,6 +207,46 @@ func hostnameFromTestCase(name string) string {
return "test-" + strings.Join(strings.Fields(name), "-")
}
+func testContainerCancel(
+ containerExtra func(c *container.Container),
+ waitCheck func(t *testing.T, c *container.Container),
+) func(t *testing.T) {
+ return func(t *testing.T) {
+ ctx, cancel := context.WithTimeout(t.Context(), helperDefaultTimeout)
+
+ c := helperNewContainer(ctx, "block")
+ c.Stdout, c.Stderr = os.Stdout, os.Stderr
+ c.WaitDelay = helperDefaultTimeout
+ if containerExtra != nil {
+ containerExtra(c)
+ }
+
+ ready := make(chan struct{})
+ if r, w, err := os.Pipe(); err != nil {
+ t.Fatalf("cannot pipe: %v", err)
+ } else {
+ c.ExtraFiles = append(c.ExtraFiles, w)
+ go func() {
+ defer close(ready)
+ if _, err = r.Read(make([]byte, 1)); err != nil {
+ panic(err.Error())
+ }
+ }()
+ }
+
+ if err := c.Start(); err != nil {
+ hlog.PrintBaseError(err, "start:")
+ t.Fatalf("cannot start container: %v", err)
+ } else if err = c.Serve(); err != nil {
+ hlog.PrintBaseError(err, "serve:")
+ t.Errorf("cannot serve setup params: %v", err)
+ }
+ <-ready
+ cancel()
+ waitCheck(t, c)
+ }
+}
+
func TestContainerString(t *testing.T) {
c := container.New(t.Context(), "ldd", "/usr/bin/env")
c.SeccompFlags |= seccomp.AllowMultiarch
@@ -227,12 +260,21 @@ func TestContainerString(t *testing.T) {
}
}
+const (
+ blockExitCodeInterrupt = 2
+)
+
func init() {
helperCommands = append(helperCommands, func(c command.Command) {
c.Command("block", command.UsageInternal, func(args []string) error {
if _, err := os.NewFile(3, "sync").Write([]byte{0}); err != nil {
return fmt.Errorf("write to sync pipe: %v", err)
}
+ {
+ sig := make(chan os.Signal, 1)
+ signal.Notify(sig, os.Interrupt)
+ go func() { <-sig; os.Exit(blockExitCodeInterrupt) }()
+ }
select {}
})