virtio-devices: block: drain async I/O before pausing

During pause the block backend's async I/O path can have unfinished I/O
requests. A snapshot or migration RAM copy taken after pause returns can
then race with kernel writes and capture torn pages.

Since vCPUs are already paused, the VMM thread can stop new block
submissions and wait for the worker to drain before parking the worker
threads.

Assisted-by: Codex:GPT-5
Signed-off-by: Dylan Reid <dgreid@fb.com>
This commit is contained in:
Dylan Reid
2026-06-11 11:16:06 -07:00
committed by Rob Bradford
parent 859bce5cae
commit 50f2fd369f

View File

@@ -16,9 +16,10 @@ use std::num::Wrapping;
use std::ops::Deref; use std::ops::Deref;
use std::os::unix::io::AsRawFd; use std::os::unix::io::AsRawFd;
use std::path::PathBuf; use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, Ordering}; use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Barrier}; use std::sync::{Arc, Barrier};
use std::{io, result}; use std::time::{Duration, Instant};
use std::{io, result, thread};
use anyhow::anyhow; use anyhow::anyhow;
use block::async_io::{AsyncIo, AsyncIoError}; use block::async_io::{AsyncIo, AsyncIoError};
@@ -149,6 +150,26 @@ impl Default for BlockCounters {
} }
} }
/// Releases one active request count when dropped.
struct ActiveRequestGuard {
counter: Arc<AtomicUsize>,
}
impl ActiveRequestGuard {
fn new(counter: &Arc<AtomicUsize>) -> Self {
Self {
counter: counter.clone(),
}
}
}
impl Drop for ActiveRequestGuard {
fn drop(&mut self) {
let previous = self.counter.fetch_sub(1, Ordering::SeqCst);
debug_assert!(previous > 0);
}
}
struct BlockEpollHandler { struct BlockEpollHandler {
queue_index: u16, queue_index: u16,
queue: Queue, queue: Queue,
@@ -163,6 +184,11 @@ struct BlockEpollHandler {
counters: BlockCounters, counters: BlockCounters,
queue_evt: EventFd, queue_evt: EventFd,
inflight_requests: VecDeque<(u16, Request)>, inflight_requests: VecDeque<(u16, Request)>,
// The active count includes `inflight_requests` plus requests in transition
// from the queue to inflight or inflight to completion.
active_request_count: Arc<AtomicUsize>,
// True when draining before pause.
draining_active_requests: Arc<AtomicBool>,
rate_limiter: Option<RateLimiterGroupHandle>, rate_limiter: Option<RateLimiterGroupHandle>,
access_platform: Option<Arc<dyn AccessPlatform>>, access_platform: Option<Arc<dyn AccessPlatform>>,
host_cpus: Option<Box<[usize]>>, host_cpus: Option<Box<[usize]>>,
@@ -215,6 +241,16 @@ impl BlockEpollHandler {
} }
fn process_queue_submit(&mut self) -> Result<()> { fn process_queue_submit(&mut self) -> Result<()> {
// Artificially bump the active counter while submitting so pause doesn't
// race and read a zero.
self.active_request_count.fetch_add(1, Ordering::SeqCst);
let _active_request = ActiveRequestGuard::new(&self.active_request_count);
// Clone the Arc so the `self.queue` mutable borrow is allowed.
let draining_active_requests = self.draining_active_requests.clone();
if draining_active_requests.load(Ordering::SeqCst) {
return Ok(());
}
let queue = &mut self.queue; let queue = &mut self.queue;
let queue_size = queue.size(); let queue_size = queue.size();
let mut batch_requests = Vec::new(); let mut batch_requests = Vec::new();
@@ -227,6 +263,9 @@ impl BlockEpollHandler {
if processed >= queue_size { if processed >= queue_size {
break; break;
} }
if draining_active_requests.load(Ordering::SeqCst) {
break;
}
processed += 1; processed += 1;
let mut desc_chain = match queue let mut desc_chain = match queue
.iter(self.mem.memory()) .iter(self.mem.memory())
@@ -323,6 +362,7 @@ impl BlockEpollHandler {
} else { } else {
self.inflight_requests self.inflight_requests
.push_back((desc_chain.head_index(), request)); .push_back((desc_chain.head_index(), request));
self.active_request_count.fetch_add(1, Ordering::SeqCst);
} }
} else { } else {
let status = match result { let status = match result {
@@ -359,7 +399,10 @@ impl BlockEpollHandler {
if !batch_requests.is_empty() { if !batch_requests.is_empty() {
match self.disk_image.submit_batch_requests(batch_requests) { match self.disk_image.submit_batch_requests(batch_requests) {
Ok(()) => { Ok(()) => {
let batch_len = batch_inflight_requests.len();
self.inflight_requests.extend(batch_inflight_requests); self.inflight_requests.extend(batch_inflight_requests);
self.active_request_count
.fetch_add(batch_len, Ordering::SeqCst);
} }
Err(e) => { Err(e) => {
// If batch submission fails, report VIRTIO_BLK_S_IOERR for all requests. // If batch submission fails, report VIRTIO_BLK_S_IOERR for all requests.
@@ -451,6 +494,7 @@ impl BlockEpollHandler {
let desc_index = completion.user_data as u16; let desc_index = completion.user_data as u16;
let mut request = self.find_inflight_request(desc_index)?; let mut request = self.find_inflight_request(desc_index)?;
let _active_request = ActiveRequestGuard::new(&self.active_request_count);
request request
.complete_async(&mem, &mut completion) .complete_async(&mem, &mut completion)
@@ -724,6 +768,8 @@ pub struct Block {
disable_sector0_writes: bool, disable_sector0_writes: bool,
lock_granularity_choice: LockGranularityChoice, lock_granularity_choice: LockGranularityChoice,
device_status: Arc<AtomicU8>, device_status: Arc<AtomicU8>,
active_request_count: Arc<AtomicUsize>,
draining_active_requests: Arc<AtomicBool>,
} }
#[derive(Serialize, Deserialize)] #[derive(Serialize, Deserialize)]
@@ -888,9 +934,39 @@ impl Block {
disable_sector0_writes, disable_sector0_writes,
lock_granularity_choice: lock_granularity, lock_granularity_choice: lock_granularity,
device_status: Arc::new(AtomicU8::new(0)), device_status: Arc::new(AtomicU8::new(0)),
active_request_count: Arc::new(AtomicUsize::new(0)),
draining_active_requests: Arc::new(AtomicBool::new(false)),
}) })
} }
fn wait_for_active_requests(&self) -> result::Result<(), anyhow::Error> {
const BLOCK_PAUSE_DRAIN_TIMEOUT: Duration = Duration::from_secs(30);
const BLOCK_PAUSE_FIRST_DRAIN_WARNING: Duration = Duration::from_secs(1);
const BLOCK_PAUSE_DRAIN_WARNING_INTERVAL: Duration = Duration::from_secs(5);
let started = Instant::now();
let mut next_warning = BLOCK_PAUSE_FIRST_DRAIN_WARNING;
loop {
let active = self.active_request_count.load(Ordering::SeqCst);
if active == 0 {
return Ok(());
}
let elapsed = started.elapsed();
if elapsed >= BLOCK_PAUSE_DRAIN_TIMEOUT {
return Err(anyhow!("timed out draining block requests"));
}
if elapsed >= next_warning {
warn!("pause: still waiting for {active} active block requests after {elapsed:?}");
next_warning += BLOCK_PAUSE_DRAIN_WARNING_INTERVAL;
}
thread::yield_now();
}
}
fn read_only(&self) -> bool { fn read_only(&self) -> bool {
has_feature(self.features(), VIRTIO_BLK_F_RO.into()) has_feature(self.features(), VIRTIO_BLK_F_RO.into())
} }
@@ -1140,6 +1216,8 @@ impl VirtioDevice for Block {
host_cpus: self.queue_affinity.get(&queue_idx).cloned(), host_cpus: self.queue_affinity.get(&queue_idx).cloned(),
acked_features: self.common.acked_features, acked_features: self.common.acked_features,
disable_sector0_writes: self.disable_sector0_writes, disable_sector0_writes: self.disable_sector0_writes,
active_request_count: self.active_request_count.clone(),
draining_active_requests: self.draining_active_requests.clone(),
}; };
let paused = self.common.paused.clone(); let paused = self.common.paused.clone();
@@ -1163,6 +1241,8 @@ impl VirtioDevice for Block {
fn reset(&mut self) { fn reset(&mut self) {
self.common.reset(); self.common.reset();
self.draining_active_requests.store(false, Ordering::SeqCst);
self.active_request_count.store(0, Ordering::SeqCst);
self.set_writeback_mode(true); self.set_writeback_mode(true);
event!("virtio-device", "reset", "id", &self.id); event!("virtio-device", "reset", "id", &self.id);
} }
@@ -1229,10 +1309,22 @@ impl VirtioDevice for Block {
impl Pausable for Block { impl Pausable for Block {
fn pause(&mut self) -> result::Result<(), MigratableError> { fn pause(&mut self) -> result::Result<(), MigratableError> {
self.common.pause() self.draining_active_requests.store(true, Ordering::SeqCst);
// Drain before parking the worker threads: the workers are what
// complete in-flight I/O, so they must keep running until the count
// reaches zero. Roll back the drain flag if any step fails.
let result = self
.wait_for_active_requests()
.map_err(MigratableError::Pause)
.and_then(|()| self.common.pause());
self.draining_active_requests.store(false, Ordering::SeqCst);
result
} }
fn resume(&mut self) -> result::Result<(), MigratableError> { fn resume(&mut self) -> result::Result<(), MigratableError> {
self.draining_active_requests.store(false, Ordering::SeqCst);
self.common.resume() self.common.resume()
} }
} }