diff --git a/virtio-devices/src/vhost_user/blk.rs b/virtio-devices/src/vhost_user/blk.rs index 26d0a7328..7e887deef 100644 --- a/virtio-devices/src/vhost_user/blk.rs +++ b/virtio-devices/src/vhost_user/blk.rs @@ -60,6 +60,7 @@ pub struct Blk { socket_path: String, epoll_thread: Option>, vu_num_queues: usize, + migration_started: bool, } impl Blk { @@ -147,6 +148,7 @@ impl Blk { socket_path: vu_cfg.socket, epoll_thread: None, vu_num_queues: num_queues, + migration_started: false, }) } @@ -409,7 +411,13 @@ impl Snapshottable for Blk { } fn snapshot(&mut self) -> std::result::Result { - Snapshot::new_from_versioned_state(&self.id(), &self.state()) + let snapshot = Snapshot::new_from_versioned_state(&self.id(), &self.state())?; + + if self.migration_started { + self.shutdown(); + } + + Ok(snapshot) } fn restore(&mut self, snapshot: Snapshot) -> std::result::Result<(), MigratableError> { @@ -421,6 +429,7 @@ impl Transportable for Blk {} impl Migratable for Blk { fn start_dirty_log(&mut self) -> std::result::Result<(), MigratableError> { + self.migration_started = true; if let Some(vu) = &self.vu { if let Some(guest_memory) = &self.guest_memory { let last_ram_addr = guest_memory.memory().last_addr().raw_value(); @@ -444,6 +453,7 @@ impl Migratable for Blk { } fn stop_dirty_log(&mut self) -> std::result::Result<(), MigratableError> { + self.migration_started = false; if let Some(vu) = &self.vu { vu.lock().unwrap().stop_dirty_log().map_err(|e| { MigratableError::StopDirtyLog(anyhow!( diff --git a/virtio-devices/src/vhost_user/fs.rs b/virtio-devices/src/vhost_user/fs.rs index b9f7952b1..93c8603cb 100644 --- a/virtio-devices/src/vhost_user/fs.rs +++ b/virtio-devices/src/vhost_user/fs.rs @@ -310,6 +310,7 @@ pub struct Fs { socket_path: String, epoll_thread: Option>, vu_num_queues: usize, + migration_started: bool, } impl Fs { @@ -398,6 +399,7 @@ impl Fs { socket_path: path.to_string(), epoll_thread: None, vu_num_queues: num_queues, + migration_started: false, }) } @@ -689,7 +691,13 @@ impl Snapshottable for Fs { } fn snapshot(&mut self) -> std::result::Result { - Snapshot::new_from_versioned_state(&self.id(), &self.state()) + let snapshot = Snapshot::new_from_versioned_state(&self.id(), &self.state())?; + + if self.migration_started { + self.shutdown(); + } + + Ok(snapshot) } fn restore(&mut self, snapshot: Snapshot) -> std::result::Result<(), MigratableError> { @@ -701,6 +709,7 @@ impl Transportable for Fs {} impl Migratable for Fs { fn start_dirty_log(&mut self) -> std::result::Result<(), MigratableError> { + self.migration_started = true; if let Some(vu) = &self.vu { if let Some(guest_memory) = &self.guest_memory { let last_ram_addr = guest_memory.memory().last_addr().raw_value(); @@ -724,6 +733,7 @@ impl Migratable for Fs { } fn stop_dirty_log(&mut self) -> std::result::Result<(), MigratableError> { + self.migration_started = false; if let Some(vu) = &self.vu { vu.lock().unwrap().stop_dirty_log().map_err(|e| { MigratableError::StopDirtyLog(anyhow!( diff --git a/virtio-devices/src/vhost_user/net.rs b/virtio-devices/src/vhost_user/net.rs index 3ce2865fc..78eba4e0a 100644 --- a/virtio-devices/src/vhost_user/net.rs +++ b/virtio-devices/src/vhost_user/net.rs @@ -116,6 +116,7 @@ pub struct Net { epoll_thread: Option>, seccomp_action: SeccompAction, vu_num_queues: usize, + migration_started: bool, } impl Net { @@ -208,6 +209,7 @@ impl Net { epoll_thread: None, seccomp_action, vu_num_queues, + migration_started: false, }) } @@ -495,7 +497,13 @@ impl Snapshottable for Net { } fn snapshot(&mut self) -> std::result::Result { - Snapshot::new_from_versioned_state(&self.id(), &self.state()) + let snapshot = Snapshot::new_from_versioned_state(&self.id(), &self.state())?; + + if self.migration_started { + self.shutdown(); + } + + Ok(snapshot) } fn restore(&mut self, snapshot: Snapshot) -> std::result::Result<(), MigratableError> { @@ -507,6 +515,7 @@ impl Transportable for Net {} impl Migratable for Net { fn start_dirty_log(&mut self) -> std::result::Result<(), MigratableError> { + self.migration_started = true; if let Some(vu) = &self.vu { if let Some(guest_memory) = &self.guest_memory { let last_ram_addr = guest_memory.memory().last_addr().raw_value(); @@ -530,6 +539,7 @@ impl Migratable for Net { } fn stop_dirty_log(&mut self) -> std::result::Result<(), MigratableError> { + self.migration_started = false; if let Some(vu) = &self.vu { vu.lock().unwrap().stop_dirty_log().map_err(|e| { MigratableError::StopDirtyLog(anyhow!(