diff --git a/pkg/sentry/devices/nvproxy/nvproxy.go b/pkg/sentry/devices/nvproxy/nvproxy.go index e580c26b1..032f04146 100644 --- a/pkg/sentry/devices/nvproxy/nvproxy.go +++ b/pkg/sentry/devices/nvproxy/nvproxy.go @@ -70,9 +70,9 @@ func Register(vfsObj *vfs.VirtualFilesystem, versionStr string, uvmDevMajor uint // +stateify savable type nvproxy struct { - objsMu objsMutex `state:"nosave"` - objsLive map[nvgpu.Handle]*object - abi *driverABI `state:"nosave"` + objsMu objsMutex `state:"nosave"` + objsLive map[nvgpu.Handle]*object `state:"nosave"` + abi *driverABI `state:"nosave"` version DriverVersion } diff --git a/pkg/sentry/devices/nvproxy/save_restore.go b/pkg/sentry/devices/nvproxy/save_restore.go index 31aa34e0c..0b810f4de 100644 --- a/pkg/sentry/devices/nvproxy/save_restore.go +++ b/pkg/sentry/devices/nvproxy/save_restore.go @@ -16,8 +16,20 @@ package nvproxy import ( "fmt" + + "gvisor.dev/gvisor/pkg/abi/nvgpu" + "gvisor.dev/gvisor/pkg/context" ) +func (n *nvproxy) beforeSave() { + n.objsMu.Lock() + defer n.objsMu.Unlock() + for _, o := range n.objsLive { + o.Release(context.Background()) + } + n.objsLive = nil +} + func (n *nvproxy) afterLoad() { Init() abiCons, ok := abis[n.version] @@ -25,4 +37,5 @@ func (n *nvproxy) afterLoad() { panic(fmt.Sprintf("driver version %q not found in abis map", n.version)) } n.abi = abiCons.cons() + n.objsLive = make(map[nvgpu.Handle]*object) }