diff --git a/pkg/sentry/kernel/BUILD b/pkg/sentry/kernel/BUILD index 968be2809..f9c561d95 100644 --- a/pkg/sentry/kernel/BUILD +++ b/pkg/sentry/kernel/BUILD @@ -338,6 +338,7 @@ go_library( "//pkg/errors", "//pkg/errors/linuxerr", "//pkg/eventchannel", + "//pkg/fd", "//pkg/fspath", "//pkg/goid", "//pkg/hostarch", diff --git a/pkg/sentry/kernel/kernel.go b/pkg/sentry/kernel/kernel.go index 1e31df435..8c63ef9ee 100644 --- a/pkg/sentry/kernel/kernel.go +++ b/pkg/sentry/kernel/kernel.go @@ -47,6 +47,7 @@ import ( "gvisor.dev/gvisor/pkg/devutil" "gvisor.dev/gvisor/pkg/errors/linuxerr" "gvisor.dev/gvisor/pkg/eventchannel" + "gvisor.dev/gvisor/pkg/fd" "gvisor.dev/gvisor/pkg/fspath" "gvisor.dev/gvisor/pkg/log" "gvisor.dev/gvisor/pkg/refs" @@ -75,6 +76,7 @@ import ( "gvisor.dev/gvisor/pkg/sentry/uniqueid" "gvisor.dev/gvisor/pkg/sentry/vfs" "gvisor.dev/gvisor/pkg/state" + "gvisor.dev/gvisor/pkg/state/statefile" "gvisor.dev/gvisor/pkg/sync" "gvisor.dev/gvisor/pkg/tcpip" ) @@ -544,7 +546,7 @@ func savePrivateMFs(ctx context.Context, w io.Writer, pw io.Writer, mfsToSave ma return nil } -func loadPrivateMFs(ctx context.Context, r io.Reader, pr io.Reader) error { +func loadPrivateMFs(ctx context.Context, r io.Reader, pr *statefile.AsyncReader) error { // Load the metadata. var meta privateMemoryFileMetadata if _, err := state.Load(ctx, r, &meta); err != nil { @@ -682,7 +684,7 @@ func (k *Kernel) invalidateUnsavableMappings(ctx context.Context) error { } // LoadFrom returns a new Kernel loaded from args. -func (k *Kernel) LoadFrom(ctx context.Context, r io.Reader, pagesFile io.Reader, timeReady chan struct{}, net inet.Stack, clocks sentrytime.Clocks, vfsOpts *vfs.CompleteRestoreOptions) error { +func (k *Kernel) LoadFrom(ctx context.Context, r io.Reader, pagesFile *fd.FD, timeReady chan struct{}, net inet.Stack, clocks sentrytime.Clocks, vfsOpts *vfs.CompleteRestoreOptions) error { loadStart := time.Now() k.runningTasksCond.L = &k.runningTasksMu @@ -724,9 +726,9 @@ func (k *Kernel) LoadFrom(ctx context.Context, r io.Reader, pagesFile io.Reader, // Load the memory files' state. memoryStart := time.Now() - pr := io.Reader(r) + var pr *statefile.AsyncReader if pagesFile != nil { - pr = pagesFile + pr = statefile.NewAsyncReader(pagesFile, 0 /* off */) } if err := k.mf.LoadFrom(ctx, r, pr); err != nil { return err @@ -734,6 +736,11 @@ func (k *Kernel) LoadFrom(ctx context.Context, r io.Reader, pagesFile io.Reader, if err := loadPrivateMFs(ctx, r, pr); err != nil { return err } + if pr != nil { + if err := pr.Close(); err != nil { + return err + } + } log.Infof("Memory files load took [%s].", time.Since(memoryStart)) log.Infof("Overall load took [%s]", time.Since(loadStart)) diff --git a/pkg/sentry/pgalloc/BUILD b/pkg/sentry/pgalloc/BUILD index 7bba423d8..bca426dba 100644 --- a/pkg/sentry/pgalloc/BUILD +++ b/pkg/sentry/pgalloc/BUILD @@ -116,6 +116,7 @@ go_library( "//pkg/sentry/memmap", "//pkg/sentry/usage", "//pkg/state", + "//pkg/state/statefile", "//pkg/state/wire", "//pkg/sync", "//pkg/sync/locking", diff --git a/pkg/sentry/pgalloc/save_restore.go b/pkg/sentry/pgalloc/save_restore.go index d77cf1c83..c211b0586 100644 --- a/pkg/sentry/pgalloc/save_restore.go +++ b/pkg/sentry/pgalloc/save_restore.go @@ -27,6 +27,7 @@ import ( "gvisor.dev/gvisor/pkg/log" "gvisor.dev/gvisor/pkg/sentry/usage" "gvisor.dev/gvisor/pkg/state" + "gvisor.dev/gvisor/pkg/state/statefile" ) // SaveTo writes f's state to the given stream. @@ -134,7 +135,7 @@ func (f *MemoryFile) RestoreID() string { } // LoadFrom loads MemoryFile state from the given stream. -func (f *MemoryFile) LoadFrom(ctx context.Context, r io.Reader, pr io.Reader) error { +func (f *MemoryFile) LoadFrom(ctx context.Context, r io.Reader, pr *statefile.AsyncReader) error { // Load metadata. if _, err := state.Load(ctx, r, &f.fileSize); err != nil { return err @@ -195,7 +196,11 @@ func (f *MemoryFile) LoadFrom(ctx context.Context, r io.Reader, pr io.Reader) er if ioErr != nil { return } - _, ioErr = io.ReadFull(pr, s) + if pr != nil { + pr.ReadAsync(s) + } else { + _, ioErr = io.ReadFull(r, s) + } }) if ioErr != nil { return ioErr diff --git a/pkg/sentry/state/BUILD b/pkg/sentry/state/BUILD index 21db8c49d..3e443644b 100644 --- a/pkg/sentry/state/BUILD +++ b/pkg/sentry/state/BUILD @@ -17,6 +17,7 @@ go_library( "//pkg/abi/linux", "//pkg/context", "//pkg/errors/linuxerr", + "//pkg/fd", "//pkg/log", "//pkg/sentry/inet", "//pkg/sentry/kernel", diff --git a/pkg/sentry/state/state.go b/pkg/sentry/state/state.go index 629d75fdc..2d4bf533b 100644 --- a/pkg/sentry/state/state.go +++ b/pkg/sentry/state/state.go @@ -22,6 +22,7 @@ import ( "gvisor.dev/gvisor/pkg/context" "gvisor.dev/gvisor/pkg/errors/linuxerr" + "gvisor.dev/gvisor/pkg/fd" "gvisor.dev/gvisor/pkg/log" "gvisor.dev/gvisor/pkg/sentry/inet" "gvisor.dev/gvisor/pkg/sentry/kernel" @@ -115,7 +116,7 @@ type LoadOpts struct { // PagesFile is the file in which all MemoryFile pages are stored if // PagesFile is non-nil. - PagesFile io.Reader + PagesFile *fd.FD // Key is used for state integrity check. Key []byte diff --git a/pkg/state/statefile/BUILD b/pkg/state/statefile/BUILD index 9e78098ca..bb56eee20 100644 --- a/pkg/state/statefile/BUILD +++ b/pkg/state/statefile/BUILD @@ -7,15 +7,29 @@ package( go_library( name = "statefile", - srcs = ["statefile.go"], + srcs = [ + "async_io.go", + "statefile.go", + ], visibility = ["//:sandbox"], - deps = ["//pkg/compressio"], + deps = [ + "//pkg/compressio", + "//pkg/fd", + "//pkg/sync", + ], ) go_test( name = "statefile_test", size = "small", - srcs = ["statefile_test.go"], + srcs = [ + "async_io_test.go", + "statefile_test.go", + ], library = ":statefile", - deps = ["//pkg/compressio"], + deps = [ + "//pkg/compressio", + "//pkg/fd", + "@org_golang_x_sys//unix:go_default_library", + ], ) diff --git a/pkg/state/statefile/async_io.go b/pkg/state/statefile/async_io.go new file mode 100644 index 000000000..78b655047 --- /dev/null +++ b/pkg/state/statefile/async_io.go @@ -0,0 +1,94 @@ +// Copyright 2024 The gVisor Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package statefile + +import ( + "runtime" + "sync/atomic" + + "gvisor.dev/gvisor/pkg/fd" + "gvisor.dev/gvisor/pkg/sync" +) + +type chunk struct { + dst []byte + off int64 +} + +// AsyncReader can be used to do reads asynchronously. It does not change the +// underlying file's offset. +type AsyncReader struct { + // in is the backing file which contains all pages. + in *fd.FD + // off is the offset being read. + off int64 + // q is the work queue. + q chan chunk + // err stores the latest IO error that occured during async read. + err atomic.Pointer[error] + // wg tracks all in flight work. + wg sync.WaitGroup +} + +// NewAsyncReader initializes a new AsyncReader. +func NewAsyncReader(in *fd.FD, off int64) *AsyncReader { + workers := runtime.GOMAXPROCS(0) + r := &AsyncReader{ + in: in, + off: off, + q: make(chan chunk, workers), + } + for i := 0; i < workers; i++ { + go r.work() + } + return r +} + +// ReadAsync schedules a read of len(p) bytes from current offset into p. +func (r *AsyncReader) ReadAsync(p []byte) { + r.wg.Add(1) + r.q <- chunk{off: r.off, dst: p} + r.off += int64(len(p)) +} + +// Wait blocks until all in flight work is complete and then returns any IO +// errors that occurred since the last call to Wait(). +func (r *AsyncReader) Wait() error { + r.wg.Wait() + if err := r.err.Swap(nil); err != nil { + return *err + } + return nil +} + +// Close calls Wait() and additionally cleans up all worker goroutines. +func (r *AsyncReader) Close() error { + err := r.Wait() + close(r.q) + return err +} + +func (r *AsyncReader) work() { + for { + c := <-r.q + if c.dst == nil { + return + } + if _, err := r.in.ReadAt(c.dst, c.off); err != nil { + r.err.Store(&err) + } + r.wg.Done() + } +} diff --git a/pkg/state/statefile/async_io_test.go b/pkg/state/statefile/async_io_test.go new file mode 100644 index 000000000..9e74c1128 --- /dev/null +++ b/pkg/state/statefile/async_io_test.go @@ -0,0 +1,59 @@ +// Copyright 2024 The gVisor Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package statefile + +import ( + "bytes" + "math/rand" + "os" + "testing" + + "golang.org/x/sys/unix" + "gvisor.dev/gvisor/pkg/fd" +) + +func TestAsyncReader(t *testing.T) { + const chunkSize = 4096 + const dataLen = 1024 * chunkSize + data := make([]byte, dataLen) + _, _ = rand.Read(data) + testFile, err := os.CreateTemp(t.TempDir(), "source") + if err != nil { + t.Fatalf("failed to create temp source file: %v", err) + } + if _, err := testFile.Write(data); err != nil { + t.Fatalf("failed to write temp source file: %v", err) + } + testFilePath := testFile.Name() + if err := testFile.Close(); err != nil { + t.Fatalf("failed to close temp source file: %v", err) + } + + 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 */) + 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 ret := bytes.Compare(p, data); ret != 0 { + t.Errorf("bytes differ") + } +} diff --git a/runsc/boot/restore.go b/runsc/boot/restore.go index 08823cf6c..3a1c15b08 100644 --- a/runsc/boot/restore.go +++ b/runsc/boot/restore.go @@ -48,7 +48,7 @@ const ( type restorer struct { container *containerInfo stateFile io.Reader - pagesFile io.Reader + pagesFile *fd.FD deviceFile *fd.FD }