diff --git a/pkg/sentry/kernel/kernel.go b/pkg/sentry/kernel/kernel.go index 7c23abeb7..783999f0f 100644 --- a/pkg/sentry/kernel/kernel.go +++ b/pkg/sentry/kernel/kernel.go @@ -806,6 +806,7 @@ func (k *Kernel) loadMemoryFiles(ctx context.Context, r io.Reader, pagesMetadata var pr *statefile.AsyncReader if pagesFile != nil { pr = statefile.NewAsyncReader(pagesFile, 0 /* off */) + defer pr.Close() } if err := k.mf.LoadFrom(ctx, pmr, pr); err != nil { return err @@ -814,7 +815,7 @@ func (k *Kernel) loadMemoryFiles(ctx context.Context, r io.Reader, pagesMetadata return err } if pr != nil { - if err := pr.Close(); err != nil { + if err := pr.Wait(); err != nil { return err } } diff --git a/pkg/state/statefile/async_io_test.go b/pkg/state/statefile/async_io_test.go index 9e74c1128..84adc1693 100644 --- a/pkg/state/statefile/async_io_test.go +++ b/pkg/state/statefile/async_io_test.go @@ -25,10 +25,13 @@ import ( ) func TestAsyncReader(t *testing.T) { + // Create random data. const chunkSize = 4096 const dataLen = 1024 * chunkSize data := make([]byte, dataLen) _, _ = rand.Read(data) + + // Create a temp file with the data. testFile, err := os.CreateTemp(t.TempDir(), "source") if err != nil { t.Fatalf("failed to create temp source file: %v", err) @@ -41,17 +44,19 @@ func TestAsyncReader(t *testing.T) { t.Fatalf("failed to close temp source file: %v", err) } + // Read the data from the file using async reads. sourceFD, err := fd.Open(testFilePath, unix.O_RDONLY, 0) if err != nil { t.Fatalf("failed to open source file %q: %v", testFilePath, err) } ar := NewAsyncReader(sourceFD, 0 /* off */) + defer ar.Close() p := make([]byte, dataLen) for i := 0; i < dataLen; i += chunkSize { ar.ReadAsync(p[i : i+chunkSize]) } - if err := ar.Close(); err != nil { - t.Fatalf("AsyncReader.Wait returned error: %v", err) + if err := ar.Wait(); err != nil { + t.Fatalf("AsyncReader.Wait failed: %v", err) } if ret := bytes.Compare(p, data); ret != 0 { t.Errorf("bytes differ")