diff --git a/resources/seccomp/aarch64-unknown-linux-musl.json b/resources/seccomp/aarch64-unknown-linux-musl.json index 26dd661e46b..df3938de108 100644 --- a/resources/seccomp/aarch64-unknown-linux-musl.json +++ b/resources/seccomp/aarch64-unknown-linux-musl.json @@ -1094,5 +1094,209 @@ ] } ] + }, + "block_io": { + "default_action": "trap", + "filter_action": "allow", + "filter": [ + { + "syscall": "exit" + }, + { + "syscall": "exit_group" + }, + { + "syscall": "read", + "comment": "Reads from the backing file" + }, + { + "syscall": "write", + "comment": "Writes to the backing file, and signals completions through an eventfd" + }, + { + "syscall": "fsync", + "comment": "Flush requests" + }, + { + "syscall": "close", + "comment": "Closes the previous backing file after a drive update" + }, + { + "syscall": "brk", + "comment": "Called for expanding the heap" + }, + { + "syscall": "clock_gettime", + "comment": "Used for metrics and logging, via the helpers in utils/src/time.rs. It's not called on some platforms, because of vdso optimisations." + }, + { + "syscall": "lseek", + "comment": "Positions the backing file before each read or write" + }, + { + "syscall": "mremap", + "comment": "Used for re-allocating large memory regions, for example vectors" + }, + { + "syscall": "munmap", + "comment": "Used for freeing memory" + }, + { + "syscall": "rt_sigprocmask", + "comment": "rt_sigprocmask is used by libc::abort during a panic to block and unblock signals" + }, + { + "syscall": "rt_sigreturn", + "comment": "rt_sigreturn is needed in case a fault does occur, so that the signal handler can return. Otherwise we get stuck in a fault loop." + }, + { + "syscall": "sigaltstack", + "comment": "sigaltstack is used by Rust stdlib to remove alternative signal stack during thread teardown." + }, + { + "syscall": "futex", + "comment": "Used for synchronization (during thread teardown when joining multiple vcpu threads at once)", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 0, + "comment": "FUTEX_WAIT" + } + ] + }, + { + "syscall": "futex", + "comment": "Used for synchronization (during thread teardown)", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 1, + "comment": "FUTEX_WAKE" + } + ] + }, + { + "syscall": "futex", + "comment": "Used for synchronization", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 128, + "comment": "FUTEX_WAIT_PRIVATE" + } + ] + }, + { + "syscall": "futex", + "comment": "Used for synchronization", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 137, + "comment": "FUTEX_WAIT_BITSET_PRIVATE" + } + ] + }, + { + "syscall": "futex", + "comment": "Used for synchronization", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 129, + "comment": "FUTEX_WAKE_PRIVATE" + } + ] + }, + { + "syscall": "madvise", + "comment": "Used by the VirtIO balloon device and by musl for some customer workloads. It is also used by aws-lc during random number generation. They setup a memory page that mark with MADV_WIPEONFORK to be able to detect forks. They also call it with -1 to see if madvise is supported in certain platforms." + }, + { + "syscall": "mmap", + "comment": "Used by the allocator", + "args": [ + { + "index": 3, + "type": "dword", + "op": "eq", + "val": 34, + "comment": "libc::MAP_ANONYMOUS | libc::MAP_PRIVATE" + } + ] + }, + { + "syscall": "tkill", + "comment": "tkill is used by libc::abort during a panic to raise SIGABRT", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 6, + "comment": "SIGABRT" + } + ] + }, + { + "syscall": "sched_yield", + "comment": "Used by the rust standard library in std::sync::mpmc. Firecracker uses mpsc channels from this module for inter-thread communication" + }, + { + "syscall": "restart_syscall", + "comment": "automatically issued by the kernel when specific timing-related syscalls (e.g. nanosleep) get interrupted by SIGSTOP" + }, + { + "syscall": "preadv", + "comment": "Reads from the backing file, into the buffers of a request" + }, + { + "syscall": "pwritev", + "comment": "Writes to the backing file, from the buffers of a request" + }, + { + "syscall": "epoll_ctl", + "comment": "Watches the queue and rate limiter events while serving the queue" + }, + { + "syscall": "epoll_pwait", + "comment": "Waits for the queue, rate limiter and control events" + }, + { + "syscall": "timerfd_settime", + "comment": "Arms the refill timer of the rate limiter", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 0 + } + ] + }, + { + "syscall": "prctl", + "comment": "Names the thread after its drive", + "args": [ + { + "index": 0, + "type": "dword", + "op": "eq", + "val": 15, + "comment": "PR_SET_NAME" + } + ] + } + ] } } diff --git a/resources/seccomp/unimplemented.json b/resources/seccomp/unimplemented.json index a919df15519..f733949931f 100644 --- a/resources/seccomp/unimplemented.json +++ b/resources/seccomp/unimplemented.json @@ -13,5 +13,10 @@ "default_action": "allow", "filter_action": "trap", "filter": [] + }, + "block_io": { + "default_action": "allow", + "filter_action": "trap", + "filter": [] } } diff --git a/resources/seccomp/x86_64-unknown-linux-musl.json b/resources/seccomp/x86_64-unknown-linux-musl.json index dcd6753a4c5..93f82d2c4e9 100644 --- a/resources/seccomp/x86_64-unknown-linux-musl.json +++ b/resources/seccomp/x86_64-unknown-linux-musl.json @@ -1226,5 +1226,209 @@ ] } ] + }, + "block_io": { + "default_action": "trap", + "filter_action": "allow", + "filter": [ + { + "syscall": "exit" + }, + { + "syscall": "exit_group" + }, + { + "syscall": "read", + "comment": "Reads from the backing file" + }, + { + "syscall": "write", + "comment": "Writes to the backing file, and signals completions through an eventfd" + }, + { + "syscall": "fsync", + "comment": "Flush requests" + }, + { + "syscall": "close", + "comment": "Closes the previous backing file after a drive update" + }, + { + "syscall": "brk", + "comment": "Called for expanding the heap" + }, + { + "syscall": "clock_gettime", + "comment": "Used for metrics and logging, via the helpers in utils/src/time.rs. It's not called on some platforms, because of vdso optimisations." + }, + { + "syscall": "lseek", + "comment": "Positions the backing file before each read or write" + }, + { + "syscall": "mremap", + "comment": "Used for re-allocating large memory regions, for example vectors" + }, + { + "syscall": "munmap", + "comment": "Used for freeing memory" + }, + { + "syscall": "rt_sigprocmask", + "comment": "rt_sigprocmask is used by libc::abort during a panic to block and unblock signals" + }, + { + "syscall": "rt_sigreturn", + "comment": "rt_sigreturn is needed in case a fault does occur, so that the signal handler can return. Otherwise we get stuck in a fault loop." + }, + { + "syscall": "sigaltstack", + "comment": "sigaltstack is used by Rust stdlib to remove alternative signal stack during thread teardown." + }, + { + "syscall": "futex", + "comment": "Used for synchronization (during thread teardown when joining multiple vcpu threads at once)", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 0, + "comment": "FUTEX_WAIT" + } + ] + }, + { + "syscall": "futex", + "comment": "Used for synchronization (during thread teardown)", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 1, + "comment": "FUTEX_WAKE" + } + ] + }, + { + "syscall": "futex", + "comment": "Used for synchronization", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 128, + "comment": "FUTEX_WAIT_PRIVATE" + } + ] + }, + { + "syscall": "futex", + "comment": "Used for synchronization", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 137, + "comment": "FUTEX_WAIT_BITSET_PRIVATE" + } + ] + }, + { + "syscall": "futex", + "comment": "Used for synchronization", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 129, + "comment": "FUTEX_WAKE_PRIVATE" + } + ] + }, + { + "syscall": "madvise", + "comment": "Used by the VirtIO balloon device and by musl for some customer workloads. It is also used by aws-lc during random number generation. They setup a memory page that mark with MADV_WIPEONFORK to be able to detect forks. They also call it with -1 to see if madvise is supported in certain platforms." + }, + { + "syscall": "mmap", + "comment": "Used by the allocator", + "args": [ + { + "index": 3, + "type": "dword", + "op": "eq", + "val": 34, + "comment": "libc::MAP_ANONYMOUS | libc::MAP_PRIVATE" + } + ] + }, + { + "syscall": "tkill", + "comment": "tkill is used by libc::abort during a panic to raise SIGABRT", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 6, + "comment": "SIGABRT" + } + ] + }, + { + "syscall": "sched_yield", + "comment": "Used by the rust standard library in std::sync::mpmc. Firecracker uses mpsc channels from this module for inter-thread communication" + }, + { + "syscall": "restart_syscall", + "comment": "automatically issued by the kernel when specific timing-related syscalls (e.g. nanosleep) get interrupted by SIGSTOP" + }, + { + "syscall": "preadv", + "comment": "Reads from the backing file, into the buffers of a request" + }, + { + "syscall": "pwritev", + "comment": "Writes to the backing file, from the buffers of a request" + }, + { + "syscall": "epoll_ctl", + "comment": "Watches the queue and rate limiter events while serving the queue" + }, + { + "syscall": "epoll_pwait", + "comment": "Waits for the queue, rate limiter and control events" + }, + { + "syscall": "timerfd_settime", + "comment": "Arms the refill timer of the rate limiter", + "args": [ + { + "index": 1, + "type": "dword", + "op": "eq", + "val": 0 + } + ] + }, + { + "syscall": "prctl", + "comment": "Names the thread after its drive", + "args": [ + { + "index": 0, + "type": "dword", + "op": "eq", + "val": 15, + "comment": "PR_SET_NAME" + } + ] + } + ] } } diff --git a/src/firecracker/Cargo.toml b/src/firecracker/Cargo.toml index 68f22554e77..0550b2ed3df 100644 --- a/src/firecracker/Cargo.toml +++ b/src/firecracker/Cargo.toml @@ -42,6 +42,7 @@ serde_json = "1.0.145" [dev-dependencies] cargo_toml = "0.22.3" libc = "0.2.177" +seccompiler = { path = "../seccompiler" } regex = { version = "1.12.2", default-features = false, features = [ "std", "unicode-perl", diff --git a/src/firecracker/src/main.rs b/src/firecracker/src/main.rs index 739214999a4..b02c4d22f00 100644 --- a/src/firecracker/src/main.rs +++ b/src/firecracker/src/main.rs @@ -363,6 +363,14 @@ fn main_exec() -> Result<(), MainError> { .and_then(seccomp::get_filters) .map_err(MainError::SeccompFilter)?; + // Threaded block IO engine workers install their own filter. + if let Some(filter) = seccomp_filters + .get("block_io") + .or_else(|| seccomp_filters.get("vmm")) + { + vmm::devices::virtio::block::virtio::set_worker_seccomp_filter(filter.clone()); + } + let vmm_config_json = arguments .single_value("config-file") .map(fs::read_to_string) diff --git a/src/firecracker/src/seccomp.rs b/src/firecracker/src/seccomp.rs index 421220a7b5f..f1095106a8c 100644 --- a/src/firecracker/src/seccomp.rs +++ b/src/firecracker/src/seccomp.rs @@ -8,6 +8,9 @@ use std::path::Path; use vmm::seccomp::{BpfThreadMap, DeserializationError, deserialize_binary, get_empty_filters}; const THREAD_CATEGORIES: [&str; 3] = ["vmm", "api", "vcpu"]; +/// Thread categories a filter file may leave out. A missing "block_io" filter, for the threaded +/// block IO engine's workers, falls back to the "vmm" one, which allows the same IO. +const OPTIONAL_THREAD_CATEGORIES: [&str; 1] = ["block_io"]; /// Error retrieving seccomp filters. #[derive(Debug, thiserror::Error, displaydoc::Display)] @@ -78,9 +81,11 @@ fn get_custom_filters(reader: R) -> Result Result { - let (filters, invalid_filters): (BpfThreadMap, BpfThreadMap) = map - .into_iter() - .partition(|(k, _)| THREAD_CATEGORIES.contains(&k.as_str())); + let (filters, invalid_filters): (BpfThreadMap, BpfThreadMap) = + map.into_iter().partition(|(k, _)| { + THREAD_CATEGORIES.contains(&k.as_str()) + || OPTIONAL_THREAD_CATEGORIES.contains(&k.as_str()) + }); if !invalid_filters.is_empty() { // build the error message let mut thread_categories_string = @@ -143,6 +148,15 @@ mod tests { assert_eq!(filter_thread_categories(map).unwrap().len(), 3); + // correct categories, including the optional ones + let mut map = BpfThreadMap::new(); + map.insert("vcpu".to_string(), Arc::new(vec![])); + map.insert("vmm".to_string(), Arc::new(vec![])); + map.insert("api".to_string(), Arc::new(vec![])); + map.insert("block_io".to_string(), Arc::new(vec![])); + + assert_eq!(filter_thread_categories(map).unwrap().len(), 4); + // invalid categories let mut map = BpfThreadMap::new(); map.insert("vcpu".to_string(), Arc::new(vec![])); @@ -168,6 +182,15 @@ mod tests { } } + #[test] + fn test_default_filters() { + // Debug builds compile an empty policy, so only release builds check the real one. + let filters = get_default_filters().unwrap(); + for category in THREAD_CATEGORIES.iter().chain(&OPTIONAL_THREAD_CATEGORIES) { + assert!(filters.contains_key(*category), "missing {category}"); + } + } + #[test] fn test_seccomp_config() { assert!(matches!( diff --git a/src/firecracker/swagger/firecracker.yaml b/src/firecracker/swagger/firecracker.yaml index 828bb086198..663ac236c4e 100644 --- a/src/firecracker/swagger/firecracker.yaml +++ b/src/firecracker/swagger/firecracker.yaml @@ -1220,9 +1220,10 @@ definitions: type: string description: Type of the IO engine used by the device. "Async" is supported on - host kernels newer than 5.10.51. + host kernels newer than 5.10.51. "Threaded" does blocking IO like + "Sync", but from a worker thread per drive. This field is optional for virtio-block config and should be omitted for vhost-user-block configuration. - enum: ["Sync", "Async"] + enum: ["Sync", "Async", "Threaded"] default: "Sync" # VhostUserBlock specific parameters diff --git a/src/firecracker/tests/threaded_block_seccomp.rs b/src/firecracker/tests/threaded_block_seccomp.rs new file mode 100644 index 00000000000..ff561c5a03f --- /dev/null +++ b/src/firecracker/tests/threaded_block_seccomp.rs @@ -0,0 +1,172 @@ +// Copyright 2026 Fly.io, Inc. +// SPDX-License-Identifier: Apache-2.0 + +//! Everything the VMM thread does with a Threaded drive after boot, under its seccomp filter. +//! +//! The VMM thread runs under the "vmm" filter of the default policy once the microVM boots, but +//! unit tests never apply it, so a system call the filter does not allow only shows in +//! production. A forbidden call kills this test with SIGSYS, and a handler names it. +//! +//! The default policy is written for musl, so this only runs on musl targets, which is where the +//! unit tests run in CI. + +#![allow(clippy::tests_outside_test_module)] +#![cfg(target_env = "musl")] + +use std::sync::Arc; +use std::thread; +use std::time::{Duration, Instant}; + +use vmm::devices::virtio::block::virtio::device::{FileEngineType, VirtioBlock}; +use vmm::devices::virtio::block::virtio::test_utils::{default_block, set_queue}; +use vmm::devices::virtio::block::virtio::{RequestHeader, VIRTIO_BLK_S_OK, VIRTIO_BLK_T_OUT}; +use vmm::devices::virtio::device::VirtioDevice; +use vmm::devices::virtio::queue::{VIRTQ_DESC_F_NEXT, VIRTQ_DESC_F_WRITE}; +use vmm::devices::virtio::test_utils::{VirtQueue, default_interrupt, default_mem}; +use vmm::rate_limiter::{BucketUpdate, RateLimiter}; +use vmm::seccomp::BpfProgram; +use vmm::vstate::memory::{Bytes, GuestAddress}; +use vmm_sys_util::tempfile::TempFile; + +/// The "vmm" filter of the default seccomp policy for this target. +fn vmm_seccomp_filter() -> BpfProgram { + let json = format!( + "{}/../../resources/seccomp/{}-unknown-linux-musl.json", + env!("CARGO_MANIFEST_DIR"), + std::env::consts::ARCH + ); + let mut policy: serde_json::Value = + serde_json::from_str(&std::fs::read_to_string(json).unwrap()).unwrap(); + if cfg!(debug_assertions) { + // Debug builds of the standard library check that a file descriptor is open before + // closing it, with fcntl(F_GETFD). Release builds, which the policy is for, don't. + policy["vmm"]["filter"] + .as_array_mut() + .unwrap() + .push(serde_json::json!({ + "syscall": "fcntl", + "args": [{"index": 1, "type": "dword", "op": "eq", "val": libc::F_GETFD}] + })); + } + let json = TempFile::new().unwrap(); + std::fs::write(json.as_path(), policy.to_string()).unwrap(); + let bpf = TempFile::new().unwrap(); + seccompiler::compile_bpf( + json.as_path().to_str().unwrap(), + std::env::consts::ARCH, + bpf.as_path().to_str().unwrap(), + false, + ) + .unwrap(); + let mut filters = vmm::seccomp::deserialize_binary(bpf.into_file()).unwrap(); + Arc::into_inner(filters.remove("vmm").unwrap()).unwrap() +} + +/// Report which system call a seccomp filter trapped, so that a failure names it. +extern "C" fn report_sigsys(_: libc::c_int, info: *mut libc::siginfo_t, _: *mut libc::c_void) { + // SAFETY: the kernel passes a valid siginfo_t, and for SIGSYS its `_sigsys` member holds the + // system call number at this offset on 64-bit Linux. + let nr = unsafe { *info.cast::().add(24).cast::() }; + let mut msg = *b"seccomp trapped system call \n"; + let mut n = u32::try_from(nr).unwrap_or(0); + for i in (29..35).rev() { + msg[i] = b'0' + u8::try_from(n % 10).unwrap(); + n /= 10; + } + // SAFETY: writing a local buffer to stderr, then exiting, both async-signal-safe. + unsafe { + libc::write(2, msg.as_ptr().cast(), msg.len()); + libc::_exit(1); + } +} + +fn wait_for(what: &str, cond: impl Fn() -> bool) { + let deadline = Instant::now() + Duration::from_secs(5); + while !cond() { + assert!(Instant::now() < deadline, "timed out waiting for {what}"); + // Not sleep: the filter does not allow it. + thread::yield_now(); + } +} + +/// Put a write of the one-sector buffer at 0x2000 to `sector` in the queue, as its `n`th +/// request, and notify the device. Returns where its status goes. +fn write_request(block: &VirtioBlock, vq: &VirtQueue, n: u16, sector: u64) -> GuestAddress { + let (header, data, status) = (0x1000, 0x2000, 0x3000 + u64::from(n)); + let first = n * 3; + vq.memory() + .write_obj( + RequestHeader::new(VIRTIO_BLK_T_OUT, sector), + GuestAddress(header), + ) + .unwrap(); + vq.dtable[usize::from(first)].set(header, 16, VIRTQ_DESC_F_NEXT, first + 1); + vq.dtable[usize::from(first + 1)].set(data, 512, VIRTQ_DESC_F_NEXT, first + 2); + vq.dtable[usize::from(first + 2)].set(status, 1, VIRTQ_DESC_F_WRITE, 0); + vq.memory().write_obj(0xffu8, GuestAddress(status)).unwrap(); + vq.avail.ring[usize::from(n)].set(first); + vq.avail.idx.set(n + 1); + block.queue_evts[0].write(1).unwrap(); + GuestAddress(status) +} + +#[test] +fn test_threaded_drive_under_vmm_seccomp() { + let filter = vmm_seccomp_filter(); + // SAFETY: installing a handler that only calls async-signal-safe functions. + unsafe { + let mut action: libc::sigaction = std::mem::zeroed(); + action.sa_sigaction = report_sigsys as usize; + action.sa_flags = libc::SA_SIGINFO; + libc::sigaction(libc::SIGSYS, &action, std::ptr::null_mut()); + } + + // Created before the filter, as drives are, before boot. + let mut block = default_block(FileEngineType::Threaded); + block.rate_limiter = RateLimiter::new(0, 0, 0, 1000, 0, 1000).unwrap(); + let mem = default_mem(); + let interrupt = default_interrupt(); + let new_file = TempFile::new().unwrap(); + new_file.as_file().set_len(0x2000).unwrap(); + let new_path = new_file.as_path().to_str().unwrap().to_string(); + + thread::spawn(move || { + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + set_queue(&mut block, 0, vq.create_queue()); + mem.write_slice(&[0x5a; 512], GuestAddress(0x2000)).unwrap(); + vmm::seccomp::apply_filter(&filter).unwrap(); + + // Activation, and the kick that hands the queue to the worker. + block.activate(mem.clone(), interrupt).unwrap(); + block.process_virtio_queues().unwrap(); + + // PATCH /drives, to a file twice the size, then a write past the end of the old one. + block.update_disk_image(new_path).unwrap(); + assert_eq!(block.config_space.capacity, 16); + let status = write_request(&block, &vq, 0, 12); + wait_for("the write", || vq.used.idx.get() == 1); + assert_eq!( + u32::from(mem.read_obj::(status).unwrap()), + VIRTIO_BLK_S_OK + ); + + // PATCH /drives with a rate limiter, and GET /vm/config. + block.update_rate_limiter(BucketUpdate::None, BucketUpdate::Disabled); + assert!(block.config().rate_limiter.is_none()); + + // A snapshot, then the kick on resume, which has the worker serve the queue again. + block.prepare_save(); + block.process_virtio_queues().unwrap(); + let status = write_request(&block, &vq, 1, 13); + wait_for("the write after resume", || vq.used.idx.get() == 2); + assert_eq!( + u32::from(mem.read_obj::(status).unwrap()), + VIRTIO_BLK_S_OK + ); + + // The drive going away. + drop(block); + }) + .join() + .unwrap(); +} diff --git a/src/vmm/src/devices/virtio/block/virtio/device.rs b/src/vmm/src/devices/virtio/block/virtio/device.rs index ecdd8ee4f6d..1f51feb49a0 100644 --- a/src/vmm/src/devices/virtio/block/virtio/device.rs +++ b/src/vmm/src/devices/virtio/block/virtio/device.rs @@ -21,7 +21,10 @@ use vmm_sys_util::eventfd::EventFd; use super::io::async_io; use super::request::*; -use super::{BLOCK_QUEUE_SIZES, SECTOR_SHIFT, SECTOR_SIZE, VirtioBlockError, io as block_io}; +use super::{ + BLOCK_QUEUE_SIZES, RATE_LIMITER_MIN_REFILL_DELAY, SECTOR_SHIFT, SECTOR_SIZE, VirtioBlockError, + io as block_io, +}; use crate::devices::virtio::ActivateError; use crate::devices::virtio::block::CacheType; use crate::devices::virtio::block::virtio::metrics::{BlockDeviceMetrics, BlockMetricsPerDevice}; @@ -50,6 +53,8 @@ pub enum FileEngineType { /// Use a Sync engine, based on blocking system calls. #[default] Sync, + /// Use a Threaded engine: blocking system calls, made from a worker thread per drive. + Threaded, } /// Helper object for setting up all `Block` fields derived from its backing file. @@ -272,7 +277,7 @@ macro_rules! unwrap_async_file_engine_or_return { ($file_engine: expr) => { match $file_engine { FileEngine::Async(engine) => engine, - FileEngine::Sync(_) => { + _ => { error!("The block device doesn't use an async IO engine"); return; } @@ -291,12 +296,13 @@ impl VirtioBlock { config.file_engine_type, )?; - let rate_limiter = config + let mut rate_limiter: RateLimiter = config .rate_limiter .map(RateLimiterConfig::try_into) .transpose() .map_err(VirtioBlockError::RateLimiter)? .unwrap_or_default(); + rate_limiter.set_min_refill_delay(RATE_LIMITER_MIN_REFILL_DELAY); let mut avail_features = (1u64 << VIRTIO_F_VERSION_1) | (1u64 << VIRTIO_RING_F_EVENT_IDX); @@ -308,6 +314,8 @@ impl VirtioBlock { avail_features |= 1u64 << VIRTIO_BLK_F_RO; }; + avail_features |= Self::threaded_features(config.file_engine_type); + let queue_evts = [EventFd::new(libc::EFD_NONBLOCK).map_err(VirtioBlockError::EventFd)?]; let queues = BLOCK_QUEUE_SIZES.iter().map(|&s| Queue::new(s)).collect(); @@ -341,7 +349,9 @@ impl VirtioBlock { /// Returns a copy of a device config pub fn config(&self) -> VirtioBlockConfig { - let rl: RateLimiterConfig = (&self.rate_limiter).into(); + let rl: RateLimiterConfig = self + .threaded_rate_limiter_config() + .unwrap_or_else(|| (&self.rate_limiter).into()); VirtioBlockConfig { drive_id: self.id.clone(), path_on_host: self.disk.file_path.clone(), @@ -374,6 +384,9 @@ impl VirtioBlock { /// Process device virtio queue(s). pub fn process_virtio_queues(&mut self) -> Result<(), InvalidAvailIdx> { + if self.threaded_kick() { + return Ok(()); + } self.process_queue(0) } @@ -537,6 +550,7 @@ impl VirtioBlock { /// Update the backing file and the config space of the block device. pub fn update_disk_image(&mut self, disk_image_path: String) -> Result<(), VirtioBlockError> { self.disk.update(disk_image_path, self.read_only)?; + self.threaded_disk_updated(); self.config_space.capacity = self.disk.nsectors.to_le(); // virtio_block_config_space(); // Kick the driver to pick up the changes. (But only if the device is already activated). @@ -552,7 +566,9 @@ impl VirtioBlock { /// Updates the parameters for the rate limiter pub fn update_rate_limiter(&mut self, bytes: BucketUpdate, ops: BucketUpdate) { - self.rate_limiter.update_buckets(bytes, ops); + if let Some((bytes, ops)) = self.threaded_update_rate_limiter(bytes, ops) { + self.rate_limiter.update_buckets(bytes, ops); + } } /// Retrieve the file engine type. @@ -560,6 +576,7 @@ impl VirtioBlock { match self.disk.file_engine { FileEngine::Sync(_) => FileEngineType::Sync, FileEngine::Async(_) => FileEngineType::Async, + FileEngine::Threaded(_) => FileEngineType::Threaded, } } @@ -575,6 +592,7 @@ impl VirtioBlock { return; } + self.threaded_stop(); self.drain_and_flush(false); if let FileEngine::Async(ref _engine) = self.disk.file_engine { self.process_async_completion_queue(); @@ -618,6 +636,9 @@ impl VirtioDevice for VirtioBlock { } fn read_config(&self, offset: u64, data: &mut [u8]) { + if self.threaded_read_config(offset, data) { + return; + } if let Some(config_space_bytes) = self.config_space.as_slice().get(u64_to_usize(offset)..) { let len = config_space_bytes.len().min(data.len()); data[..len].copy_from_slice(&config_space_bytes[..len]); @@ -675,6 +696,7 @@ impl VirtioDevice for VirtioBlock { impl Drop for VirtioBlock { fn drop(&mut self) { + self.threaded_stop(); match self.cache_type { CacheType::Unsafe => { if let Err(err) = self.disk.file_engine.drain(true) { diff --git a/src/vmm/src/devices/virtio/block/virtio/event_handler.rs b/src/vmm/src/devices/virtio/block/virtio/event_handler.rs index 9f02862f814..68295d44786 100644 --- a/src/vmm/src/devices/virtio/block/virtio/event_handler.rs +++ b/src/vmm/src/devices/virtio/block/virtio/event_handler.rs @@ -15,6 +15,10 @@ impl VirtioBlock { const PROCESS_ASYNC_COMPLETION: u32 = 3; fn register_runtime_events(&self, ops: &mut EventOps) { + // The threaded engine's worker serves the queue, and has its events. + if matches!(self.disk.file_engine, FileEngine::Threaded(_)) { + return; + } if let Err(err) = ops.add(Events::with_data( &self.queue_evts[0], Self::PROCESS_QUEUE, @@ -84,7 +88,10 @@ impl MutEventSubscriber for VirtioBlock { if self.is_activated() { match source { - Self::PROCESS_ACTIVATE => self.process_activate_event(ops), + Self::PROCESS_ACTIVATE => { + self.process_activate_event(ops); + self.threaded_start(); + } Self::PROCESS_QUEUE => self.process_queue_event(), Self::PROCESS_RATE_LIMITER => self.process_rate_limiter_event(), Self::PROCESS_ASYNC_COMPLETION => self.process_async_completion_event(), @@ -105,6 +112,7 @@ impl MutEventSubscriber for VirtioBlock { // - on device restore from snapshot. if self.is_activated() { self.register_runtime_events(ops); + self.threaded_start(); } else { self.register_activate_event(ops); } diff --git a/src/vmm/src/devices/virtio/block/virtio/io/mod.rs b/src/vmm/src/devices/virtio/block/virtio/io/mod.rs index b7aa8061d76..0121fc7bc4f 100644 --- a/src/vmm/src/devices/virtio/block/virtio/io/mod.rs +++ b/src/vmm/src/devices/virtio/block/virtio/io/mod.rs @@ -3,12 +3,14 @@ pub mod async_io; pub mod sync_io; +pub mod threaded_io; use std::fmt::Debug; use std::fs::File; pub use self::async_io::{AsyncFileEngine, AsyncIoError}; pub use self::sync_io::{SyncFileEngine, SyncIoError}; +pub use self::threaded_io::{ThreadedFileEngine, ThreadedIoError}; use crate::devices::virtio::block::virtio::PendingRequest; use crate::devices::virtio::block::virtio::device::FileEngineType; use crate::vstate::memory::{GuestAddress, GuestMemoryMmap}; @@ -31,6 +33,8 @@ pub enum BlockIoError { Sync(SyncIoError), /// Async error: {0} Async(AsyncIoError), + /// Threaded error: {0} + Threaded(ThreadedIoError), } impl BlockIoError { @@ -54,6 +58,7 @@ pub enum FileEngine { #[allow(unused)] Async(AsyncFileEngine), Sync(SyncFileEngine), + Threaded(ThreadedFileEngine), } impl FileEngine { @@ -63,6 +68,9 @@ impl FileEngine { AsyncFileEngine::from_file(file).map_err(BlockIoError::Async)?, )), FileEngineType::Sync => Ok(FileEngine::Sync(SyncFileEngine::from_file(file))), + FileEngineType::Threaded => Ok(FileEngine::Threaded( + ThreadedFileEngine::from_file(file).map_err(BlockIoError::Threaded)?, + )), } } @@ -70,6 +78,9 @@ impl FileEngine { match self { FileEngine::Async(engine) => engine.update_file(file).map_err(BlockIoError::Async)?, FileEngine::Sync(engine) => engine.update_file(file), + FileEngine::Threaded(engine) => { + engine.update_file(file).map_err(BlockIoError::Threaded)? + } }; Ok(()) @@ -80,6 +91,7 @@ impl FileEngine { match self { FileEngine::Async(engine) => engine.file(), FileEngine::Sync(engine) => engine.file(), + FileEngine::Threaded(engine) => engine.file(), } } @@ -99,6 +111,13 @@ impl FileEngine { error: BlockIoError::Async(err.error), }), }, + FileEngine::Threaded(_) => { + let err = ThreadedFileEngine::refuse(req); + Err(RequestError { + req: err.req, + error: BlockIoError::Threaded(err.error), + }) + } FileEngine::Sync(engine) => match engine.read(offset, mem, addr, count) { Ok(count) => Ok(FileEngineOk::Executed(RequestOk { req, count })), Err(err) => Err(RequestError { @@ -125,6 +144,13 @@ impl FileEngine { error: BlockIoError::Async(err.error), }), }, + FileEngine::Threaded(_) => { + let err = ThreadedFileEngine::refuse(req); + Err(RequestError { + req: err.req, + error: BlockIoError::Threaded(err.error), + }) + } FileEngine::Sync(engine) => match engine.write(offset, mem, addr, count) { Ok(count) => Ok(FileEngineOk::Executed(RequestOk { req, count })), Err(err) => Err(RequestError { @@ -147,6 +173,13 @@ impl FileEngine { error: BlockIoError::Async(err.error), }), }, + FileEngine::Threaded(_) => { + let err = ThreadedFileEngine::refuse(req); + Err(RequestError { + req: err.req, + error: BlockIoError::Threaded(err.error), + }) + } FileEngine::Sync(engine) => match engine.flush() { Ok(_) => Ok(FileEngineOk::Executed(RequestOk { req, count: 0 })), Err(err) => Err(RequestError { @@ -161,6 +194,7 @@ impl FileEngine { match self { FileEngine::Async(engine) => engine.drain(discard).map_err(BlockIoError::Async), FileEngine::Sync(_engine) => Ok(()), + FileEngine::Threaded(engine) => engine.drain(discard).map_err(BlockIoError::Threaded), } } @@ -170,6 +204,9 @@ impl FileEngine { engine.drain_and_flush(discard).map_err(BlockIoError::Async) } FileEngine::Sync(engine) => engine.flush().map_err(BlockIoError::Sync), + FileEngine::Threaded(engine) => engine + .drain_and_flush(discard) + .map_err(BlockIoError::Threaded), } } } diff --git a/src/vmm/src/devices/virtio/block/virtio/io/threaded_io.rs b/src/vmm/src/devices/virtio/block/virtio/io/threaded_io.rs new file mode 100644 index 00000000000..5eccd8ad316 --- /dev/null +++ b/src/vmm/src/devices/virtio/block/virtio/io/threaded_io.rs @@ -0,0 +1,884 @@ +// Copyright 2026 Fly.io, Inc. +// SPDX-License-Identifier: Apache-2.0 + +//! Threaded file engine. +//! +//! Each drive gets a worker thread that serves its virtqueue on its own, with blocking IO: it +//! waits on the queue's eventfd and the rate limiter's timer, parses and rate-limits requests, +//! reads and writes the backing file with `preadv`/`pwritev`, fills the used ring and raises the +//! guest interrupt. A slow backing file then stalls only its own worker, not the event loop and +//! every other device serviced there, and no io_uring support is needed. +//! +//! The event loop only sends the worker control messages: start and stop serving the queue, kick +//! it, swap the backing file, and wait for it. Requests carry up to [`THREADED_SEG_MAX`] guest +//! buffers each. + +use std::fs::File; +use std::mem::ManuallyDrop; +use std::os::unix::io::{AsRawFd, FromRawFd, RawFd}; +use std::sync::{Arc, Mutex, OnceLock, mpsc}; +use std::thread; + +use vm_memory::GuestMemoryError; +use vmm_sys_util::epoll::{ControlOperation, Epoll, EpollEvent, EventSet}; +use vmm_sys_util::eventfd::EventFd; + +use super::{BlockIoError, RequestError}; +use crate::devices::virtio::block::virtio::metrics::BlockDeviceMetrics; +use crate::devices::virtio::block::virtio::{ + FinishedRequest, IoErr, PendingRequest, Request, RequestType, SECTOR_SHIFT, VIRTIO_BLK_ID_BYTES, +}; +use crate::devices::virtio::queue::Queue; +use crate::devices::virtio::transport::{VirtioInterrupt, VirtioInterruptType}; +use crate::logger::{IncMetric, error}; +use crate::rate_limiter::{BucketUpdate, RateLimiter}; +use crate::seccomp::{BpfProgram, BpfProgramRef}; +use crate::vstate::memory::{ + Bytes, GuestAddress, GuestMemory, GuestMemoryExtension, GuestMemoryMmap, +}; + +/// Maximum number of data buffers in one request, advertised to the guest as `seg_max`. +/// +/// Without it a guest must send one request per physically contiguous buffer, which for page +/// cache writeback means one per 4 KiB page. Indirect descriptors are not supported, so a +/// request takes this many entries of the 256-entry queue, plus two. +pub const THREADED_SEG_MAX: u32 = 32; + +/// A guest buffer: its address and length. +pub type Segment = (GuestAddress, u32); + +/// Filter the worker threads apply to themselves, see [`set_worker_seccomp_filter`]. +static WORKER_SECCOMP_FILTER: OnceLock> = OnceLock::new(); + +/// Set the seccomp filter each worker thread applies to itself when it starts. +/// +/// Workers are started whenever a drive is created, which is before the VMM thread installs its +/// own filter, so they cannot rely on inheriting it. Only the first call has an effect. +pub fn set_worker_seccomp_filter(filter: Arc) { + // Ignoring the error is what makes later calls no-ops. + let _ = WORKER_SECCOMP_FILTER.set(filter); +} + +fn worker_seccomp_filter() -> Option> { + WORKER_SECCOMP_FILTER.get().map(|filter| filter.as_slice()) +} + +/// The prefix of the worker thread names; the rest is the start of the drive id. Linux keeps 15 +/// bytes of a thread name, which leaves 11 for the drive id. +const THREAD_NAME_PREFIX: &str = "blk_"; + +#[derive(Debug, thiserror::Error, displaydoc::Display)] +pub enum ThreadedIoError { + /// Read: {0} + Read(std::io::Error), + /// Write: {0} + Write(std::io::Error), + /// SyncAll: {0} + SyncAll(std::io::Error), + /// Guest memory: {0} + GuestMemory(GuestMemoryError), + /// EventFd: {0} + EventFd(std::io::Error), + /// Epoll: {0} + Epoll(std::io::Error), + /// Cloning the backing file: {0} + FileClone(std::io::Error), + /// Spawning the IO worker thread: {0} + Spawn(std::io::Error), + /// The IO worker thread is gone + WorkerGone, + /// Requests are served by the IO worker thread + ServedByWorker, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Direction { + Read, + Write, +} + +/// Transfer `segments` from or to `file` at `offset` with positioned, vectored IO: one system +/// call, unless the kernel transfers less than asked. +fn transfer( + file: &File, + direction: Direction, + mut offset: u64, + mem: &GuestMemoryMmap, + segments: &[Segment], +) -> Result { + let mut iovecs = Vec::with_capacity(segments.len()); + let mut total: u32 = 0; + for &(addr, len) in segments { + let slice = mem + .get_slice(addr, len as usize) + .map_err(ThreadedIoError::GuestMemory)?; + iovecs.push(libc::iovec { + iov_base: slice.ptr_guard_mut().as_ptr().cast(), + iov_len: len as usize, + }); + total = total.checked_add(len).ok_or(ThreadedIoError::GuestMemory( + GuestMemoryError::GuestAddressOverflow, + ))?; + } + + let io_error = |err| match direction { + Direction::Read => ThreadedIoError::Read(err), + Direction::Write => ThreadedIoError::Write(err), + }; + + let mut iovecs = iovecs.as_mut_slice(); + while !iovecs.is_empty() { + let iovcnt = libc::c_int::try_from(iovecs.len()).unwrap_or(libc::c_int::MAX); + let offset_arg = + libc::off_t::try_from(offset).map_err(|_| io_error(libc::EOVERFLOW.into_io()))?; + // SAFETY: the iovecs point into guest memory that `mem`, which outlives the call, keeps + // mapped, each within the bounds `get_slice` checked. + let ret = unsafe { + match direction { + Direction::Read => { + libc::preadv(file.as_raw_fd(), iovecs.as_ptr(), iovcnt, offset_arg) + } + Direction::Write => { + libc::pwritev(file.as_raw_fd(), iovecs.as_ptr(), iovcnt, offset_arg) + } + } + }; + let mut done = match usize::try_from(ret) { + Ok(0) => { + return Err(io_error(match direction { + Direction::Read => std::io::ErrorKind::UnexpectedEof.into(), + Direction::Write => std::io::ErrorKind::WriteZero.into(), + })); + } + Ok(done) => done, + Err(_) => { + let err = std::io::Error::last_os_error(); + if err.kind() == std::io::ErrorKind::Interrupted { + continue; + } + return Err(io_error(err)); + } + }; + offset += done as u64; + + // Skip what was transferred: whole buffers first, then the start of the next one. + while let Some(first) = iovecs.first_mut() { + if done < first.iov_len { + // SAFETY: `done` is within the buffer. + first.iov_base = unsafe { first.iov_base.add(done) }; + first.iov_len -= done; + break; + } + done -= first.iov_len; + iovecs = &mut iovecs[1..]; + } + } + + Ok(total) +} + +trait IntoIoError { + fn into_io(self) -> std::io::Error; +} + +impl IntoIoError for libc::c_int { + fn into_io(self) -> std::io::Error { + std::io::Error::from_raw_os_error(self) + } +} + +/// What the worker needs to serve a device's queue, from its activation on. +#[derive(Debug)] +pub struct Serving { + /// The virtqueue. The worker owns it until it stops, and hands it back then. + pub queue: Queue, + /// The queue's eventfd. It stays the device's: the worker only waits on it, and stops before + /// the device closes it. + pub queue_evt: RawFd, + pub mem: GuestMemoryMmap, + pub interrupt: Arc, + pub metrics: Arc, + pub nsectors: u64, + pub image_id: [u8; VIRTIO_BLK_ID_BYTES as usize], + /// Whether requests may carry several data buffers. + pub segmented: bool, + pub drive_id: String, +} + +#[derive(Debug)] +enum Control { + Start(Box, Arc>), + Stop, + Kick, + UpdateDisk { + file: File, + nsectors: u64, + image_id: [u8; VIRTIO_BLK_ID_BYTES as usize], + }, + SyncAll, + Barrier, + Exit, +} + +#[derive(Debug)] +enum Reply { + Stopped(Option), + Synced(Result<(), std::io::Error>), + Done, +} + +// Epoll tokens. +const CONTROL: u64 = 0; +const QUEUE: u64 = 1; +const RATE_LIMITER: u64 = 2; + +/// The worker's state while it serves the queue. +#[derive(Debug)] +struct Active { + queue: Queue, + queue_evt: ManuallyDrop, + mem: GuestMemoryMmap, + interrupt: Arc, + rate_limiter: Arc>, + rate_limiter_fd: RawFd, + metrics: Arc, + nsectors: u64, + image_id: [u8; VIRTIO_BLK_ID_BYTES as usize], + segmented: bool, +} + +#[derive(Debug)] +struct Worker { + file: File, + epoll: Epoll, + control_evt: EventFd, + control: mpsc::Receiver, + replies: mpsc::Sender, + active: Option, +} + +impl Worker { + fn run(mut self) { + if let Some(filter) = worker_seccomp_filter() + && let Err(err) = crate::seccomp::apply_filter(filter) + { + panic!("Failed to set the requested seccomp filters on the block IO worker: {err}"); + } + + let mut events = [EpollEvent::default(); 4]; + loop { + let count = match self.epoll.wait(-1, &mut events) { + Ok(count) => count, + Err(err) if err.kind() == std::io::ErrorKind::Interrupted => continue, + Err(err) => { + error!("Block IO worker failed to wait for events: {:?}", err); + return; + } + }; + for event in &events[..count] { + match event.data() { + CONTROL => { + let _ = self.control_evt.read(); + while let Ok(control) = self.control.try_recv() { + if !self.handle(control) { + return; + } + } + } + QUEUE => self.queue_event(), + RATE_LIMITER => self.rate_limiter_event(), + _ => {} + } + } + } + } + + /// Handle a control message. Returns false when the worker has to exit. + fn handle(&mut self, control: Control) -> bool { + match control { + Control::Start(serving, rate_limiter) => self.start(*serving, rate_limiter), + Control::Stop => { + let queue = self.stop(); + let _ = self.replies.send(Reply::Stopped(queue)); + } + Control::Kick => self.process_queue(), + Control::UpdateDisk { + file, + nsectors, + image_id, + } => { + self.file = file; + if let Some(active) = self.active.as_mut() { + active.nsectors = nsectors; + active.image_id = image_id; + } + let _ = self.replies.send(Reply::Done); + } + Control::SyncAll => { + let _ = self.replies.send(Reply::Synced(self.file.sync_all())); + } + Control::Barrier => { + let _ = self.replies.send(Reply::Done); + } + Control::Exit => { + self.stop(); + return false; + } + } + true + } + + fn start(&mut self, serving: Serving, rate_limiter: Arc>) { + let name = format!("{THREAD_NAME_PREFIX}{}", serving.drive_id); + let mut name = name.into_bytes(); + name.truncate(15); + name.push(0); + // SAFETY: `name` is NUL terminated, and at most 16 bytes long, as PR_SET_NAME expects. + unsafe { libc::prctl(libc::PR_SET_NAME, name.as_ptr()) }; + + let rate_limiter_fd = rate_limiter + .lock() + .expect("Poisoned block rate limiter lock") + .as_raw_fd(); + for (fd, token) in [(serving.queue_evt, QUEUE), (rate_limiter_fd, RATE_LIMITER)] { + if let Err(err) = self.epoll.ctl( + ControlOperation::Add, + fd, + EpollEvent::new(EventSet::IN, token), + ) { + error!("Block IO worker failed to watch fd {}: {:?}", fd, err); + } + } + + self.active = Some(Active { + queue: serving.queue, + // SAFETY: the fd is the device's, which keeps it open until it has stopped the + // worker. `ManuallyDrop` keeps the worker from closing it. + queue_evt: ManuallyDrop::new(unsafe { EventFd::from_raw_fd(serving.queue_evt) }), + mem: serving.mem, + interrupt: serving.interrupt, + rate_limiter, + rate_limiter_fd, + metrics: serving.metrics, + nsectors: serving.nsectors, + image_id: serving.image_id, + segmented: serving.segmented, + }); + // Serve what the guest queued before the worker was watching. + self.process_queue(); + } + + /// Stop serving the queue, and return it. + fn stop(&mut self) -> Option { + let active = self.active.take()?; + for fd in [active.queue_evt.as_raw_fd(), active.rate_limiter_fd] { + let _ = self + .epoll + .ctl(ControlOperation::Delete, fd, EpollEvent::default()); + } + // Dropping the rest, the rate limiter included, before the device gets the queue. + Some(active.queue) + } + + fn queue_event(&mut self) { + let Some(active) = self.active.as_mut() else { + return; + }; + active.metrics.queue_event_count.inc(); + if let Err(err) = active.queue_evt.read() { + error!("Failed to get queue event: {:?}", err); + active.metrics.event_fails.inc(); + } else if active + .rate_limiter + .lock() + .expect("Poisoned block rate limiter lock") + .is_blocked() + { + active.metrics.rate_limiter_throttled_events.inc(); + } else { + self.process_queue(); + } + } + + fn rate_limiter_event(&mut self) { + let Some(active) = self.active.as_mut() else { + return; + }; + active.metrics.rate_limiter_event_count.inc(); + let refilled = active + .rate_limiter + .lock() + .expect("Poisoned block rate limiter lock") + .event_handler() + .is_ok(); + if refilled { + self.process_queue(); + } + } + + /// Serve every request available, until the queue is empty or the rate limiter blocks. + fn process_queue(&mut self) { + let Some(active) = self.active.as_mut() else { + return; + }; + let mut used_any = false; + + loop { + let head = match active.queue.pop_or_enable_notification() { + Ok(Some(head)) => head, + Ok(None) => break, + Err(err) => panic!("Block queue is corrupt: {err}"), + }; + active + .metrics + .remaining_reqs_count + .add(active.queue.len().into()); + + let parsed = Request::parse_any(&head, &active.mem, active.nsectors, active.segmented); + let finished = match parsed { + Ok(request) => { + let limited = request.rate_limit( + &mut active + .rate_limiter + .lock() + .expect("Poisoned block rate limiter lock"), + ); + if limited { + // Leave it in the avail ring until the rate limiter's timer fires. + active.queue.undo_pop(); + active.metrics.rate_limiter_throttled_events.inc(); + break; + } + execute(&self.file, active, request, head.index) + } + Err(err) => { + error!("Failed to parse available descriptor chain: {:?}", err); + active.metrics.execute_fails.inc(); + FinishedRequest { + num_bytes_to_mem: 0, + desc_idx: head.index, + } + } + }; + + used_any = true; + active + .queue + .add_used(finished.desc_idx, finished.num_bytes_to_mem) + .unwrap_or_else(|err| { + error!( + "Failed to add available descriptor head {}: {}", + finished.desc_idx, err + ) + }); + } + active.queue.advance_used_ring_idx(); + + if used_any && active.queue.prepare_kick() { + active + .interrupt + .trigger(VirtioInterruptType::Queue(0)) + .unwrap_or_else(|_| { + active.metrics.event_fails.inc(); + }); + } + if !used_any { + active.metrics.no_avail_buffer.inc(); + } + } +} + +/// Perform a request with blocking IO, write its status, and return what to put in the used +/// ring. +fn execute(file: &File, active: &Active, mut request: Request, desc_idx: u16) -> FinishedRequest { + let (mem, metrics) = (&active.mem, &*active.metrics); + let pending = request.to_pending_request(desc_idx); + let file_error = |err| IoErr::FileEngine(BlockIoError::Threaded(err)); + + let result = match request.r#type { + RequestType::In | RequestType::Out => { + let segments = if request.segments.is_empty() { + vec![(request.data_addr, request.data_len)] + } else { + std::mem::take(&mut request.segments) + }; + let offset = request.sector << SECTOR_SHIFT; + if request.r#type == RequestType::In { + let _metric = metrics.read_agg.record_latency_metrics(); + let result = transfer(file, Direction::Read, offset, mem, &segments); + if result.is_ok() { + // The guest memory was written from this thread, so account for it in the + // dirty bitmap before the guest gets to see the completion. + for &(addr, len) in &segments { + mem.mark_dirty(addr, len as usize); + } + } + result.map_err(file_error) + } else { + let _metric = metrics.write_agg.record_latency_metrics(); + transfer(file, Direction::Write, offset, mem, &segments).map_err(file_error) + } + } + // Sync data out to physical media on host. + RequestType::Flush => file + .sync_all() + .map(|_| 0) + .map_err(|err| file_error(ThreadedIoError::SyncAll(err))), + RequestType::GetDeviceID => mem + .write_slice(&active.image_id, request.data_addr) + .map(|_| VIRTIO_BLK_ID_BYTES) + .map_err(IoErr::GetId), + RequestType::Unsupported(_) => Ok(0), + }; + + pending.finish(mem, result, metrics) +} + +/// Front end of the threaded engine, used from the thread that owns the device. +#[derive(Debug)] +pub struct ThreadedFileEngine { + control: mpsc::Sender, + control_evt: EventFd, + replies: mpsc::Receiver, + // While the worker serves the queue: the rate limiter it shares with the device. + rate_limiter: Option>>, + // Takes the device's rate limiter's place while the worker has it. Made here, before the VMM + // thread's seccomp filter forbids creating the timerfd a rate limiter needs. + spare_rate_limiter: Option, + // A new backing file, until the device hands it over with its properties. + new_file: Option, + #[cfg(test)] + file: File, + worker: Option>, +} + +impl ThreadedFileEngine { + pub fn from_file(file: File) -> Result { + let control_evt = EventFd::new(libc::EFD_NONBLOCK).map_err(ThreadedIoError::EventFd)?; + let worker_control_evt = control_evt.try_clone().map_err(ThreadedIoError::EventFd)?; + // Before the worker applies its seccomp filter, which does not allow creating it. + let epoll = Epoll::new().map_err(ThreadedIoError::Epoll)?; + epoll + .ctl( + ControlOperation::Add, + worker_control_evt.as_raw_fd(), + EpollEvent::new(EventSet::IN, CONTROL), + ) + .map_err(ThreadedIoError::Epoll)?; + #[cfg(test)] + let test_file = file.try_clone().map_err(ThreadedIoError::FileClone)?; + let (control, worker_control) = mpsc::channel(); + let (worker_replies, replies) = mpsc::channel(); + + let worker = Worker { + file, + epoll, + control_evt: worker_control_evt, + control: worker_control, + replies: worker_replies, + active: None, + }; + let worker = thread::Builder::new() + .name("fc_blk_io".to_string()) + .spawn(move || worker.run()) + .map_err(ThreadedIoError::Spawn)?; + + Ok(ThreadedFileEngine { + control, + control_evt, + replies, + rate_limiter: None, + spare_rate_limiter: Some(RateLimiter::default()), + new_file: None, + #[cfg(test)] + file: test_file, + worker: Some(worker), + }) + } + + /// The backing file the engine was created with. + #[cfg(test)] + pub fn file(&self) -> &File { + &self.file + } + + fn send(&self, control: Control) -> Result<(), ThreadedIoError> { + self.control + .send(control) + .map_err(|_| ThreadedIoError::WorkerGone)?; + self.control_evt.write(1).map_err(ThreadedIoError::EventFd) + } + + fn reply(&self) -> Result { + self.replies.recv().map_err(|_| ThreadedIoError::WorkerGone) + } + + /// Whether the worker serves the device's queue. + pub fn is_serving(&self) -> bool { + self.rate_limiter.is_some() + } + + /// Hand the device's queue and rate limiter to the worker, which serves the queue from then + /// on. The device gets a spare rate limiter in exchange, until [`Self::stop`]. + pub fn start( + &mut self, + serving: Serving, + device_rate_limiter: &mut RateLimiter, + ) -> Result<(), ThreadedIoError> { + let Some(mut spare) = self.spare_rate_limiter.take() else { + // Already serving. + return Ok(()); + }; + std::mem::swap(device_rate_limiter, &mut spare); + let rate_limiter = Arc::new(Mutex::new(spare)); + self.send(Control::Start(Box::new(serving), rate_limiter.clone()))?; + self.rate_limiter = Some(rate_limiter); + Ok(()) + } + + /// Have the worker stop serving the queue. Returns the queue, which the device gets back, + /// and puts its rate limiter back. + pub fn stop( + &mut self, + device_rate_limiter: &mut RateLimiter, + ) -> Result, ThreadedIoError> { + let Some(rate_limiter) = self.rate_limiter.take() else { + return Ok(None); + }; + self.send(Control::Stop)?; + let queue = match self.reply()? { + Reply::Stopped(queue) => queue, + reply => panic!("Unexpected block IO worker reply to stop: {reply:?}"), + }; + // The worker dropped its reference before replying. + let mut rate_limiter = Arc::try_unwrap(rate_limiter) + .expect("Block IO worker kept the rate limiter") + .into_inner() + .expect("Poisoned block rate limiter lock"); + std::mem::swap(device_rate_limiter, &mut rate_limiter); + self.spare_rate_limiter = Some(rate_limiter); + Ok(queue) + } + + /// Have the worker look at the queue, as if the guest had notified it. + pub fn kick(&self) -> Result<(), ThreadedIoError> { + self.send(Control::Kick) + } + + /// The rate limiter the worker uses, while it serves the queue. + pub fn rate_limiter(&self) -> Option<&Arc>> { + self.rate_limiter.as_ref() + } + + /// Update the rate limiter the worker uses. Returns the updates back when the worker does + /// not serve the queue, for the device to apply to its own. + pub fn update_rate_limiter( + &self, + bytes: BucketUpdate, + ops: BucketUpdate, + ) -> Option<(BucketUpdate, BucketUpdate)> { + match self.rate_limiter.as_ref() { + Some(rate_limiter) => { + rate_limiter + .lock() + .expect("Poisoned block rate limiter lock") + .update_buckets(bytes, ops); + None + } + None => Some((bytes, ops)), + } + } + + /// Take a new backing file. The worker switches to it with [`Self::update_disk`]. + pub fn update_file(&mut self, file: File) -> Result<(), ThreadedIoError> { + // No clone of it: the VMM thread's seccomp filter does not allow one after boot. + self.new_file = Some(file); + Ok(()) + } + + /// Switch the worker to the backing file given to [`Self::update_file`], with its size and + /// id, at once. Requests the worker served before use the old file, later ones the new one. + pub fn update_disk( + &mut self, + nsectors: u64, + image_id: [u8; VIRTIO_BLK_ID_BYTES as usize], + ) -> Result<(), ThreadedIoError> { + let Some(file) = self.new_file.take() else { + return Ok(()); + }; + self.send(Control::UpdateDisk { + file, + nsectors, + image_id, + })?; + match self.reply()? { + Reply::Done => Ok(()), + reply => panic!("Unexpected block IO worker reply to a disk update: {reply:?}"), + } + } + + /// Wait for the worker to finish the request it is serving, if any. + pub fn drain(&mut self, _discard: bool) -> Result<(), ThreadedIoError> { + self.send(Control::Barrier)?; + match self.reply()? { + Reply::Done => Ok(()), + reply => panic!("Unexpected block IO worker reply to a barrier: {reply:?}"), + } + } + + pub fn drain_and_flush(&mut self, _discard: bool) -> Result<(), ThreadedIoError> { + // Sync data out to physical media on host, from the worker, which has the backing file. + // It finishes what it is serving first. + self.send(Control::SyncAll)?; + match self.reply()? { + Reply::Synced(result) => result.map_err(ThreadedIoError::SyncAll), + reply => panic!("Unexpected block IO worker reply to a sync: {reply:?}"), + } + } + + /// Requests are served by the worker, from the queue. Should one get here, it fails. + pub fn refuse(req: PendingRequest) -> RequestError { + RequestError { + req, + error: ThreadedIoError::ServedByWorker, + } + } +} + +impl Drop for ThreadedFileEngine { + fn drop(&mut self) { + let _ = self.send(Control::Exit); + if let Some(worker) = self.worker.take() + && worker.join().is_err() + { + error!("The block IO worker thread panicked"); + } + } +} + +#[cfg(test)] +mod tests { + #![allow(clippy::undocumented_unsafe_blocks)] + use vm_memory::GuestMemoryRegion; + use vmm_sys_util::tempfile::TempFile; + + use super::*; + use crate::utils::u64_to_usize; + use crate::vmm_config::machine_config::HugePageConfig; + use crate::vstate::memory; + use crate::vstate::memory::{Bitmap, GuestRegionMmapExt}; + + const MEM_LEN: usize = 8192; + + fn create_mem() -> GuestMemoryMmap { + GuestMemoryMmap::from_regions( + memory::anonymous( + [(GuestAddress(0), MEM_LEN)].into_iter(), + true, + HugePageConfig::None, + ) + .unwrap() + .into_iter() + .map(|region| GuestRegionMmapExt::dram_from_mmap_region(region, 0)) + .collect(), + ) + .unwrap() + } + + fn is_dirty(mem: &GuestMemoryMmap, addr: GuestAddress) -> bool { + mem.find_region(addr) + .unwrap() + .bitmap() + .dirty_at(u64_to_usize(addr.0)) + } + + #[test] + fn test_transfer() { + let file = TempFile::new().unwrap().into_file(); + let data: Vec = (0..3072u32).map(|i| (i % 251) as u8).collect(); + + // Write three buffers that are neither adjacent nor in order in guest memory. + let mem = create_mem(); + mem.write_slice(&data[..1024], GuestAddress(4096)).unwrap(); + mem.write_slice(&data[1024..1536], GuestAddress(0)).unwrap(); + mem.write_slice(&data[1536..], GuestAddress(2048)).unwrap(); + let segments = [ + (GuestAddress(4096), 1024), + (GuestAddress(0), 512), + (GuestAddress(2048), 1536), + ]; + assert_eq!( + transfer(&file, Direction::Write, 512, &mem, &segments).unwrap(), + 3072 + ); + + // Read them back into different buffers. + let mem = create_mem(); + let segments = [(GuestAddress(1024), 2048), (GuestAddress(6144), 1024)]; + assert_eq!( + transfer(&file, Direction::Read, 512, &mem, &segments).unwrap(), + 3072 + ); + let mut buf = vec![0u8; 3072]; + mem.read_slice(&mut buf[..2048], GuestAddress(1024)) + .unwrap(); + mem.read_slice(&mut buf[2048..], GuestAddress(6144)) + .unwrap(); + assert_eq!(buf, data); + } + + #[test] + fn test_transfer_errors() { + let file = TempFile::new().unwrap().into_file(); + file.set_len(4096).unwrap(); + let mem = create_mem(); + + // One bad buffer fails the whole request, before any of it is transferred. + let segments = [(GuestAddress(0), 512), (GuestAddress(MEM_LEN as u64), 512)]; + assert!(matches!( + transfer(&file, Direction::Read, 0, &mem, &segments), + Err(ThreadedIoError::GuestMemory(_)) + )); + assert!(!is_dirty(&mem, GuestAddress(0))); + + // Reading past the end of the file is an error, not a short read. + assert!(matches!( + transfer( + &file, + Direction::Read, + 3584, + &mem, + &[(GuestAddress(0), 1024)] + ), + Err(ThreadedIoError::Read(_)) + )); + } + + #[test] + fn test_engine_before_activation() { + let old = TempFile::new().unwrap(); + let new = TempFile::new().unwrap(); + let mut engine = ThreadedFileEngine::from_file(old.as_file().try_clone().unwrap()).unwrap(); + assert!(!engine.is_serving()); + assert!(engine.rate_limiter().is_none()); + + // Without a queue to serve, the worker still answers, and takes a new backing file. + engine.drain(true).unwrap(); + engine.drain_and_flush(true).unwrap(); + engine + .update_file(new.as_file().try_clone().unwrap()) + .unwrap(); + engine + .update_disk(0, [0; VIRTIO_BLK_ID_BYTES as usize]) + .unwrap(); + engine.drain(false).unwrap(); + + // Rate limiter updates are the device's to apply. + let updates = engine.update_rate_limiter(BucketUpdate::None, BucketUpdate::Disabled); + assert!(updates.is_some()); + + // Stopping a worker that serves nothing returns nothing. + let mut rate_limiter = RateLimiter::default(); + assert!(engine.stop(&mut rate_limiter).unwrap().is_none()); + } +} diff --git a/src/vmm/src/devices/virtio/block/virtio/mod.rs b/src/vmm/src/devices/virtio/block/virtio/mod.rs index 9e97d6d3897..62f80f7027f 100644 --- a/src/vmm/src/devices/virtio/block/virtio/mod.rs +++ b/src/vmm/src/devices/virtio/block/virtio/mod.rs @@ -10,6 +10,9 @@ pub mod metrics; pub mod persist; pub mod request; pub mod test_utils; +mod threaded; + +pub use self::io::threaded_io::set_worker_seccomp_filter; use vm_memory::GuestMemoryError; @@ -29,6 +32,9 @@ pub const BLOCK_QUEUE_SIZES: [u16; BLOCK_NUM_QUEUES] = [FIRECRACKER_MAX_QUEUE_SI // So we can use 128 IO_URING entries without ever triggering a FullSq Error. /// Maximum number of io uring entries we allow in the queue. pub const IO_URING_NUM_ENTRIES: u16 = 128; +/// Minimum wait before retrying a depleted rate limiter. The limiter otherwise waits +/// until the failed request's tokens have refilled, rather than a fixed 100ms. +pub const RATE_LIMITER_MIN_REFILL_DELAY: std::time::Duration = std::time::Duration::from_millis(5); /// Errors the block device can trigger. #[derive(Debug, thiserror::Error, displaydoc::Display)] @@ -47,6 +53,8 @@ pub enum VirtioBlockError { InvalidDataLength, /// The requested operation would cause a seek beyond disk end. InvalidOffset, + /// Guest gave us more data descriptors than the device allows. + TooManySegments, /// Guest gave us a read only descriptor that protocol says to write to. UnexpectedReadOnlyDescriptor, /// Guest gave us a write only descriptor that protocol says to read from. diff --git a/src/vmm/src/devices/virtio/block/virtio/persist.rs b/src/vmm/src/devices/virtio/block/virtio/persist.rs index 380fe1de0e8..cde43b0fecf 100644 --- a/src/vmm/src/devices/virtio/block/virtio/persist.rs +++ b/src/vmm/src/devices/virtio/block/virtio/persist.rs @@ -30,6 +30,8 @@ pub enum FileEngineTypeState { Sync, /// Async File Engine. Async, + /// Threaded File Engine. + Threaded, } impl From for FileEngineTypeState { @@ -37,6 +39,7 @@ impl From for FileEngineTypeState { match file_engine_type { FileEngineType::Sync => FileEngineTypeState::Sync, FileEngineType::Async => FileEngineTypeState::Async, + FileEngineType::Threaded => FileEngineTypeState::Threaded, } } } @@ -46,6 +49,7 @@ impl From for FileEngineType { match file_engine_type_state { FileEngineTypeState::Sync => FileEngineType::Sync, FileEngineTypeState::Async => FileEngineType::Async, + FileEngineTypeState::Threaded => FileEngineType::Threaded, } } } @@ -87,8 +91,9 @@ impl Persist<'_> for VirtioBlock { state: &Self::State, ) -> Result { let is_read_only = state.virtio_state.avail_features & (1u64 << VIRTIO_BLK_F_RO) != 0; - let rate_limiter = RateLimiter::restore((), &state.rate_limiter_state) + let mut rate_limiter = RateLimiter::restore((), &state.rate_limiter_state) .map_err(VirtioBlockError::RateLimiter)?; + rate_limiter.set_min_refill_delay(RATE_LIMITER_MIN_REFILL_DELAY); let disk_properties = DiskProperties::new( state.disk_path.clone(), @@ -239,5 +244,14 @@ mod tests { // Test that block specific fields are the same. assert_eq!(restored_block.disk.file_path, block.disk.file_path); + // The adaptive refill delay is device policy and must be reapplied on restore. + assert_eq!( + block.rate_limiter.min_refill_delay(), + Some(RATE_LIMITER_MIN_REFILL_DELAY) + ); + assert_eq!( + restored_block.rate_limiter.min_refill_delay(), + Some(RATE_LIMITER_MIN_REFILL_DELAY) + ); } } diff --git a/src/vmm/src/devices/virtio/block/virtio/request.rs b/src/vmm/src/devices/virtio/block/virtio/request.rs index 8fc83cf43da..322d1d3101f 100644 --- a/src/vmm/src/devices/virtio/block/virtio/request.rs +++ b/src/vmm/src/devices/virtio/block/virtio/request.rs @@ -208,6 +208,10 @@ pub struct RequestHeader { unsafe impl ByteValued for RequestHeader {} impl RequestHeader { + pub(super) fn type_and_sector(&self) -> (RequestType, u64) { + (RequestType::from(self.request_type), self.sector) + } + pub fn new(request_type: u32, sector: u64) -> RequestHeader { RequestHeader { request_type, @@ -236,8 +240,10 @@ pub struct Request { pub r#type: RequestType, pub data_len: u32, pub status_addr: GuestAddress, - sector: u64, - data_addr: GuestAddress, + pub(super) sector: u64, + pub(super) data_addr: GuestAddress, + /// The data buffers of a request parsed with `parse_segmented`, `data_len` bytes in total. + pub(super) segments: Vec<(GuestAddress, u32)>, } impl Request { @@ -258,6 +264,7 @@ impl Request { data_addr: GuestAddress(0), data_len: 0, status_addr: GuestAddress(0), + segments: Vec::new(), }; let data_desc; @@ -354,7 +361,7 @@ impl Request { self.sector << SECTOR_SHIFT } - fn to_pending_request(&self, desc_idx: u16) -> PendingRequest { + pub(super) fn to_pending_request(&self, desc_idx: u16) -> PendingRequest { PendingRequest { r#type: self.r#type, data_len: self.data_len, @@ -834,6 +841,7 @@ mod tests { status_addr, sector: sector & (NUM_DISK_SECTORS - sectors_len), data_addr, + segments: Vec::new(), }; let mut request_header = RequestHeader::new(virtio_request_id, request.sector); diff --git a/src/vmm/src/devices/virtio/block/virtio/test_utils.rs b/src/vmm/src/devices/virtio/block/virtio/test_utils.rs index e4f23c6a038..0584a38f9e5 100644 --- a/src/vmm/src/devices/virtio/block/virtio/test_utils.rs +++ b/src/vmm/src/devices/virtio/block/virtio/test_utils.rs @@ -122,6 +122,10 @@ pub fn simulate_queue_and_async_completion_events(b: &mut VirtioBlock, expected_ FileEngine::Sync(_) => { simulate_queue_event(b, Some(expected_irq)); } + // The threaded engine's worker serves its queue on its own. + FileEngine::Threaded(_) => { + simulate_queue_event(b, Some(expected_irq)); + } } } diff --git a/src/vmm/src/devices/virtio/block/virtio/threaded.rs b/src/vmm/src/devices/virtio/block/virtio/threaded.rs new file mode 100644 index 00000000000..b5cb941ac90 --- /dev/null +++ b/src/vmm/src/devices/virtio/block/virtio/threaded.rs @@ -0,0 +1,743 @@ +// Copyright 2026 Fly.io, Inc. +// SPDX-License-Identifier: Apache-2.0 + +//! Block device support for the threaded IO engine: requests with several data buffers, and +//! handing the queue to the engine's worker thread, which serves it from activation on. + +use vm_memory::ByteValued; + +use std::os::unix::io::AsRawFd; + +use super::device::{FileEngineType, VirtioBlock}; +use super::io::FileEngine; +use super::io::threaded_io::{Serving, THREADED_SEG_MAX}; +use super::request::{Request, RequestType}; +use super::{SECTOR_SHIFT, SECTOR_SIZE, VirtioBlockError}; +use crate::devices::virtio::device::VirtioDevice; +use crate::devices::virtio::generated::virtio_blk::VIRTIO_BLK_F_SEG_MAX; +use crate::devices::virtio::queue::DescriptorChain; +use crate::logger::{IncMetric, error}; +use crate::rate_limiter::BucketUpdate; +use crate::utils::u64_to_usize; +use crate::vmm_config::RateLimiterConfig; +use crate::vstate::memory::GuestMemoryMmap; + +/// The start of `struct virtio_blk_config`, up to `seg_max`. The device's own config space ends +/// after `capacity`. +#[derive(Debug, Default, Clone, Copy)] +#[repr(C)] +struct SegmentedConfigSpace { + capacity: u64, + size_max: u32, + seg_max: u32, +} + +// SAFETY: `SegmentedConfigSpace` contains only PODs in `repr(C)`, without padding. +unsafe impl ByteValued for SegmentedConfigSpace {} + +impl Request { + /// Parse a request, with any number of data buffers if `segmented`. + pub fn parse_any( + avail_desc: &DescriptorChain, + mem: &GuestMemoryMmap, + num_disk_sectors: u64, + segmented: bool, + ) -> Result { + if segmented { + Self::parse_segmented(avail_desc, mem, num_disk_sectors) + } else { + Self::parse(avail_desc, mem, num_disk_sectors) + } + } + + /// Parse a request whose data may come in several buffers: every descriptor between the + /// header and the status is one. Requests that are not reads or writes are left to `parse`. + pub fn parse_segmented( + avail_desc: &DescriptorChain, + mem: &GuestMemoryMmap, + num_disk_sectors: u64, + ) -> Result { + // The head contains the request type which MUST be readable. + if avail_desc.is_write_only() { + return Err(VirtioBlockError::UnexpectedWriteOnlyDescriptor); + } + let header: super::RequestHeader = { + use crate::vstate::memory::Bytes; + mem.read_obj(avail_desc.addr) + .map_err(VirtioBlockError::GuestMemory)? + }; + let (r#type, sector) = header.type_and_sector(); + if r#type != RequestType::In && r#type != RequestType::Out { + return Self::parse(avail_desc, mem, num_disk_sectors); + } + + let mut segments = Vec::new(); + let mut data_len: u32 = 0; + let mut desc = avail_desc + .next_descriptor() + .ok_or(VirtioBlockError::DescriptorChainTooShort)?; + // The last descriptor is the status, the ones before it are data. + while let Some(next) = desc.next_descriptor() { + // This also bounds the walk, should the chain loop. + if segments.len() >= u64_to_usize(u64::from(THREADED_SEG_MAX)) { + return Err(VirtioBlockError::TooManySegments); + } + if desc.is_write_only() && r#type == RequestType::Out { + return Err(VirtioBlockError::UnexpectedWriteOnlyDescriptor); + } + if !desc.is_write_only() && r#type == RequestType::In { + return Err(VirtioBlockError::UnexpectedReadOnlyDescriptor); + } + data_len = data_len + .checked_add(desc.len) + .ok_or(VirtioBlockError::InvalidDataLength)?; + segments.push((desc.addr, desc.len)); + desc = next; + } + let status_desc = desc; + if segments.is_empty() { + return Err(VirtioBlockError::DescriptorChainTooShort); + } + + // Check that the data length is a multiple of 512 as specified in the virtio standard. + if data_len % SECTOR_SIZE != 0 { + return Err(VirtioBlockError::InvalidDataLength); + } + let top_sector = sector + .checked_add(u64::from(data_len) >> SECTOR_SHIFT) + .ok_or(VirtioBlockError::InvalidOffset)?; + if top_sector > num_disk_sectors { + return Err(VirtioBlockError::InvalidOffset); + } + + // The status MUST always be writable. + if !status_desc.is_write_only() { + return Err(VirtioBlockError::UnexpectedReadOnlyDescriptor); + } + if status_desc.len < 1 { + return Err(VirtioBlockError::DescriptorLengthTooSmall); + } + + Ok(Request { + r#type, + data_len, + status_addr: status_desc.addr, + sector, + data_addr: segments[0].0, + segments, + }) + } +} + +impl VirtioBlock { + /// Features only the threaded engine offers. + pub(super) fn threaded_features(engine_type: FileEngineType) -> u64 { + match engine_type { + FileEngineType::Threaded => 1u64 << VIRTIO_BLK_F_SEG_MAX, + _ => 0, + } + } + + /// Whether requests may have several data buffers. Decided by what the device offered, not + /// by what the driver accepted: a driver that did not accept it sends one buffer anyway. + pub(super) fn accepts_segments(&self) -> bool { + self.avail_features & (1u64 << VIRTIO_BLK_F_SEG_MAX) != 0 + && matches!(self.disk.file_engine, FileEngine::Threaded(_)) + } + + /// Serve a config space read from the longer config space of a device that takes several + /// data buffers per request. Returns false if this device does not. + pub(super) fn threaded_read_config(&self, offset: u64, data: &mut [u8]) -> bool { + if !self.accepts_segments() { + return false; + } + + let config_space = SegmentedConfigSpace { + capacity: self.config_space.capacity, + size_max: 0, + seg_max: THREADED_SEG_MAX.to_le(), + }; + if let Some(bytes) = config_space.as_slice().get(u64_to_usize(offset)..) { + let len = bytes.len().min(data.len()); + data[..len].copy_from_slice(&bytes[..len]); + } else { + error!("Failed to read config space"); + self.metrics.cfg_fails.inc(); + } + true + } + + /// Hand the queue to the worker, which serves it from then on. Called once the device is + /// activated, and again after `threaded_stop`, when the device is kicked. + pub(super) fn threaded_start(&mut self) { + let FileEngine::Threaded(engine) = &self.disk.file_engine else { + return; + }; + if engine.is_serving() { + return; + } + let Some(active_state) = self.device_state.active_state() else { + return; + }; + let serving = Serving { + // The device keeps its copy for the transport, which reads the queue's + // configuration. The worker's copy is the one in use until it hands it back. + queue: self.queues[0].clone(), + queue_evt: self.queue_evts[0].as_raw_fd(), + mem: active_state.mem.clone(), + interrupt: active_state.interrupt.clone(), + metrics: self.metrics.clone(), + nsectors: self.disk.nsectors, + image_id: self.disk.image_id, + segmented: self.accepts_segments(), + drive_id: self.id.clone(), + }; + let FileEngine::Threaded(engine) = &mut self.disk.file_engine else { + return; + }; + if let Err(err) = engine.start(serving, &mut self.rate_limiter) { + error!("Failed to start the block IO worker: {:?}", err); + self.metrics.event_fails.inc(); + } + } + + /// Take the queue and the rate limiter back from the worker, for the device to save them, + /// or before it goes away. + pub(super) fn threaded_stop(&mut self) { + let FileEngine::Threaded(engine) = &mut self.disk.file_engine else { + return; + }; + match engine.stop(&mut self.rate_limiter) { + Ok(Some(queue)) => self.queues[0] = queue, + Ok(None) => {} + Err(err) => error!("Failed to stop the block IO worker: {:?}", err), + } + } + + /// Have the worker look at the queue, starting it if it is not serving. Returns false if the + /// device does not use the threaded engine. + pub(super) fn threaded_kick(&mut self) -> bool { + let FileEngine::Threaded(engine) = &self.disk.file_engine else { + return false; + }; + if !engine.is_serving() { + // Starting serves what is queued. + self.threaded_start(); + } else if let Err(err) = engine.kick() { + error!("Failed to kick the block IO worker: {:?}", err); + } + true + } + + /// The configuration of the rate limiter the worker uses, while it serves the queue. + pub(super) fn threaded_rate_limiter_config(&self) -> Option { + let FileEngine::Threaded(engine) = &self.disk.file_engine else { + return None; + }; + let rate_limiter = engine.rate_limiter()?; + let rate_limiter = rate_limiter + .lock() + .expect("Poisoned block rate limiter lock"); + Some((&*rate_limiter).into()) + } + + /// Update the rate limiter the worker uses, while it serves the queue. Returns the updates + /// back otherwise. + pub(super) fn threaded_update_rate_limiter( + &self, + bytes: BucketUpdate, + ops: BucketUpdate, + ) -> Option<(BucketUpdate, BucketUpdate)> { + match &self.disk.file_engine { + FileEngine::Threaded(engine) => engine.update_rate_limiter(bytes, ops), + _ => Some((bytes, ops)), + } + } + + /// Switch the worker to the backing file the disk was just updated with. + pub(super) fn threaded_disk_updated(&mut self) { + let (nsectors, image_id) = (self.disk.nsectors, self.disk.image_id); + if let FileEngine::Threaded(engine) = &mut self.disk.file_engine + && let Err(err) = engine.update_disk(nsectors, image_id) + { + error!("Failed to update the block IO worker's disk: {:?}", err); + } + } +} + +#[cfg(test)] +mod tests { + use std::sync::{Arc, Mutex}; + use std::thread; + use std::time::{Duration, Instant}; + + use event_manager::{EventManager, SubscriberOps}; + use vmm_sys_util::tempfile::TempFile; + + use super::*; + use crate::devices::virtio::block::persist::BlockConstructorArgs; + use crate::devices::virtio::block::virtio::device::VirtioBlockConfig; + use crate::devices::virtio::block::virtio::test_utils::{default_block, set_queue}; + use crate::devices::virtio::block::virtio::{ + CacheType, RequestHeader, VIRTIO_BLK_S_OK, VIRTIO_BLK_T_FLUSH, VIRTIO_BLK_T_IN, + VIRTIO_BLK_T_OUT, + }; + use crate::devices::virtio::queue::{VIRTQ_DESC_F_NEXT, VIRTQ_DESC_F_WRITE}; + use crate::devices::virtio::test_utils::{VirtQueue, default_interrupt, default_mem}; + use crate::devices::virtio::transport::VirtioInterruptType; + use crate::rate_limiter::RateLimiter; + use crate::snapshot::Persist; + use crate::vstate::memory::{Address, Bytes, GuestAddress}; + + fn add_flush_requests_batch(block: &mut VirtioBlock, vq: &VirtQueue, count: u16) { + let mem = vq.memory(); + vq.avail.idx.set(0); + vq.used.idx.set(0); + set_queue(block, 0, vq.create_queue()); + + let hdr_addr = vq + .end() + .checked_align_up(std::mem::align_of::() as u64) + .unwrap(); + mem.write_obj(RequestHeader::new(VIRTIO_BLK_T_FLUSH, 0), hdr_addr) + .unwrap(); + let mut status_addr = hdr_addr + .checked_add(std::mem::size_of::() as u64) + .unwrap() + .checked_align_up(4) + .unwrap(); + + for i in 0..count { + let idx = i * 2; + let hdr_desc = &vq.dtable[idx as usize]; + hdr_desc.addr.set(hdr_addr.0); + hdr_desc.flags.set(VIRTQ_DESC_F_NEXT); + hdr_desc.next.set(idx + 1); + + let status_desc = &vq.dtable[idx as usize + 1]; + status_desc.addr.set(status_addr.0); + status_desc.flags.set(VIRTQ_DESC_F_WRITE); + status_desc.len.set(4); + status_addr = status_addr.checked_add(4).unwrap(); + + vq.avail.ring[i as usize].set(idx); + vq.avail.idx.set(i + 1); + } + } + + fn check_flush_requests_batch(count: u16, vq: &VirtQueue) { + assert_eq!(vq.used.idx.get(), count); + for i in 0..count { + let used = vq.used.ring[i as usize].get(); + let status_addr = vq.dtable[used.id as usize + 1].addr.get(); + assert_eq!(used.len, 1); + assert_eq!( + u32::from( + vq.memory() + .read_obj::(GuestAddress(status_addr)) + .unwrap() + ), + VIRTIO_BLK_S_OK + ); + } + } + + fn wait_for(what: &str, cond: impl Fn() -> bool) { + let deadline = Instant::now() + Duration::from_secs(5); + while !cond() { + assert!(Instant::now() < deadline, "timed out waiting for {what}"); + thread::sleep(Duration::from_millis(1)); + } + } + + fn engine( + block: &VirtioBlock, + ) -> &crate::devices::virtio::block::virtio::io::ThreadedFileEngine { + match &block.disk.file_engine { + FileEngine::Threaded(engine) => engine, + _ => panic!("not a threaded block device"), + } + } + + /// Activate the device and hand its queue to the worker, as the event loop does. + fn serve(block: &mut VirtioBlock, mem: &GuestMemoryMmap) { + block.activate(mem.clone(), default_interrupt()).unwrap(); + block.threaded_start(); + assert!(engine(block).is_serving()); + } + + fn notify(block: &VirtioBlock) { + block.queue_evts[0].write(1).unwrap(); + } + + fn interrupted(block: &VirtioBlock) -> bool { + block + .interrupt_trigger() + .has_pending_interrupt(VirtioInterruptType::Queue(0)) + } + + #[test] + fn test_engine_type() { + let block = default_block(FileEngineType::Threaded); + assert!(matches!(block.disk.file_engine, FileEngine::Threaded(_))); + assert_eq!(block.file_engine_type(), FileEngineType::Threaded); + assert_eq!(block.config().file_engine_type, FileEngineType::Threaded); + } + + /// Chain a header, `segments` data buffers and a status, from descriptor 0 on. + fn set_request( + vq: &VirtQueue, + request_type: u32, + sector: u64, + segments: &[(u64, u32)], + ) -> GuestAddress { + let (header_addr, status_addr) = (0x1000, 0x1800); + vq.memory() + .write_obj( + RequestHeader::new(request_type, sector), + GuestAddress(header_addr), + ) + .unwrap(); + vq.dtable[0].set(header_addr, 16, VIRTQ_DESC_F_NEXT, 1); + let data_flags = match request_type { + VIRTIO_BLK_T_IN => VIRTQ_DESC_F_NEXT | VIRTQ_DESC_F_WRITE, + _ => VIRTQ_DESC_F_NEXT, + }; + let mut idx = 1; + for &(addr, len) in segments { + vq.dtable[usize::from(idx)].set(addr, len, data_flags, idx + 1); + idx += 1; + } + vq.dtable[usize::from(idx)].set(status_addr, 1, VIRTQ_DESC_F_WRITE, 0); + vq.memory() + .write_obj(0xffu8, GuestAddress(status_addr)) + .unwrap(); + vq.avail.ring[0].set(0); + vq.avail.idx.set(1); + vq.used.idx.set(0); + GuestAddress(status_addr) + } + + #[test] + fn test_features_and_config_space() { + let block = default_block(FileEngineType::Threaded); + assert_ne!(block.avail_features() & (1 << VIRTIO_BLK_F_SEG_MAX), 0); + assert!(block.accepts_segments()); + + // capacity (8 sectors), size_max, seg_max + let mut config = [0xffu8; 16]; + block.read_config(0, &mut config); + assert_eq!(config[..8], 8u64.to_le_bytes()); + assert_eq!(config[8..12], 0u32.to_le_bytes()); + assert_eq!(config[12..], THREADED_SEG_MAX.to_le_bytes()); + // Partial reads, as the guest does them. + let mut seg_max = [0u8; 4]; + block.read_config(12, &mut seg_max); + assert_eq!(seg_max, THREADED_SEG_MAX.to_le_bytes()); + + // The other engines offer neither the feature nor the longer config space. + for engine in [FileEngineType::Sync, FileEngineType::Async] { + let block = default_block(engine); + assert_eq!(block.avail_features() & (1 << VIRTIO_BLK_F_SEG_MAX), 0); + assert!(!block.accepts_segments()); + let mut config = [0xffu8; 16]; + block.read_config(0, &mut config); + assert_eq!(config[..8], 8u64.to_le_bytes()); + assert_eq!(config[8..], [0xff; 8]); + } + } + + #[test] + fn test_read_write() { + let mut block = default_block(FileEngineType::Threaded); + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + set_queue(&mut block, 0, vq.create_queue()); + serve(&mut block, &mem); + + // The backing file is 8 sectors. Write 4 of them from three scattered buffers. + let data: Vec = (0..2048u32).map(|i| (i % 251) as u8).collect(); + mem.write_slice(&data[..512], GuestAddress(0x4000)).unwrap(); + mem.write_slice(&data[512..1536], GuestAddress(0x2000)) + .unwrap(); + mem.write_slice(&data[1536..], GuestAddress(0x6000)) + .unwrap(); + let status_addr = set_request( + &vq, + VIRTIO_BLK_T_OUT, + 2, + &[(0x4000, 512), (0x2000, 1024), (0x6000, 512)], + ); + // The worker serves the queue from the guest's notification, without the event loop. + notify(&block); + wait_for("the write", || vq.used.idx.get() == 1); + wait_for("the interrupt", || interrupted(&block)); + assert_eq!(vq.used.ring[0].get().id, 0); + assert_eq!(vq.used.ring[0].get().len, 1); + assert_eq!( + u32::from(mem.read_obj::(status_addr).unwrap()), + VIRTIO_BLK_S_OK + ); + + // Read them back as two buffers. + let status_addr = set_request(&vq, VIRTIO_BLK_T_IN, 2, &[(0x8000, 1536), (0xa000, 512)]); + vq.avail.ring[1].set(0); + vq.avail.idx.set(2); + vq.used.idx.set(1); + notify(&block); + wait_for("the read", || vq.used.idx.get() == 2); + assert_eq!(vq.used.ring[1].get().len, 2049); + assert_eq!( + u32::from(mem.read_obj::(status_addr).unwrap()), + VIRTIO_BLK_S_OK + ); + let mut buf = vec![0u8; 2048]; + mem.read_slice(&mut buf[..1536], GuestAddress(0x8000)) + .unwrap(); + mem.read_slice(&mut buf[1536..], GuestAddress(0xa000)) + .unwrap(); + assert_eq!(buf, data); + } + + #[test] + fn test_serves_what_was_queued_before() { + let mut block = default_block(FileEngineType::Threaded); + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + add_flush_requests_batch(&mut block, &vq, 5); + // No notification: the worker looks at the queue when it starts. + serve(&mut block, &mem); + wait_for("the flushes", || vq.used.idx.get() == 5); + check_flush_requests_batch(5, &vq); + } + + #[test] + fn test_rate_limiter() { + let mut block = default_block(FileEngineType::Threaded); + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + add_flush_requests_batch(&mut block, &vq, 4); + // One op per 100 ms. + block.rate_limiter = RateLimiter::new(0, 0, 0, 1, 0, 100).unwrap(); + serve(&mut block, &mem); + + wait_for("the first flush", || vq.used.idx.get() >= 1); + // The worker has the rate limiter, and the device reports it. + let config = block.config().rate_limiter.unwrap(); + assert_eq!(config.ops.unwrap().size, 1); + // Its timer lets the rest through, one at a time. + wait_for("the second flush", || vq.used.idx.get() >= 2); + + // An update reaches the worker's rate limiter. + block.update_rate_limiter(BucketUpdate::None, BucketUpdate::Disabled); + assert!(block.config().rate_limiter.is_none()); + block.process_virtio_queues().unwrap(); + wait_for("the other flushes", || vq.used.idx.get() == 4); + check_flush_requests_batch(4, &vq); + } + + #[test] + fn test_prepare_save() { + let mut block = default_block(FileEngineType::Threaded); + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + add_flush_requests_batch(&mut block, &vq, 5); + block.rate_limiter = RateLimiter::new(0, 0, 0, 100, 0, 1000).unwrap(); + serve(&mut block, &mem); + wait_for("the flushes", || vq.used.idx.get() == 5); + + // The device gets the queue, as the worker left it, and its rate limiter back. + block.prepare_save(); + assert!(!engine(&block).is_serving()); + assert_eq!(block.queues[0].next_avail.0, 5); + assert_eq!(block.queues[0].next_used.0, 5); + let config: RateLimiterConfig = (&block.rate_limiter).into(); + assert_eq!(config.ops.unwrap().size, 100); + + // A kick, as when the VM resumes, has the worker serve the queue again. + block.process_virtio_queues().unwrap(); + assert!(engine(&block).is_serving()); + } + + #[test] + fn test_update_disk_image() { + let mut block = default_block(FileEngineType::Threaded); + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + set_queue(&mut block, 0, vq.create_queue()); + serve(&mut block, &mem); + + // A new backing file of 16 sectors, twice the old one. + let new_file = TempFile::new().unwrap(); + new_file.as_file().set_len(0x2000).unwrap(); + block + .update_disk_image(new_file.as_path().to_str().unwrap().to_string()) + .unwrap(); + assert_eq!(block.config_space.capacity, 16); + + // The worker writes to it, past the end of the old one. + mem.write_slice(&[0x5a; 512], GuestAddress(0x2000)).unwrap(); + let status_addr = set_request(&vq, VIRTIO_BLK_T_OUT, 12, &[(0x2000, 512)]); + notify(&block); + wait_for("the write", || vq.used.idx.get() == 1); + assert_eq!( + u32::from(mem.read_obj::(status_addr).unwrap()), + VIRTIO_BLK_S_OK + ); + let mut buf = [0u8; 512]; + use std::os::unix::fs::FileExt; + new_file + .as_file() + .read_exact_at(&mut buf, 12 * 512) + .unwrap(); + assert_eq!(buf, [0x5a; 512]); + } + + #[test] + fn test_event_handler() { + let mut event_manager = EventManager::new().unwrap(); + let block = default_block(FileEngineType::Threaded); + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + let block = Arc::new(Mutex::new(block)); + set_queue(&mut block.lock().unwrap(), 0, vq.create_queue()); + event_manager.add_subscriber(block.clone()); + block + .lock() + .unwrap() + .activate(mem.clone(), default_interrupt()) + .unwrap(); + // The activation event hands the queue to the worker. + assert_eq!(event_manager.run_with_timeout(50).unwrap(), 1); + assert!(engine(&block.lock().unwrap()).is_serving()); + + // Which serves the queue from then on: the event loop gets no event for it. + let status_addr = set_request(&vq, VIRTIO_BLK_T_OUT, 0, &[(0x2000, 512)]); + notify(&block.lock().unwrap()); + assert_eq!(event_manager.run_with_timeout(50).unwrap(), 0); + wait_for("the write", || vq.used.idx.get() == 1); + assert_eq!( + u32::from(mem.read_obj::(status_addr).unwrap()), + VIRTIO_BLK_S_OK + ); + } + + #[test] + fn test_worker_thread_name() { + let mut block = default_block(FileEngineType::Threaded); + block.id = "vol_x7k2p9q4mz81".to_string(); + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + set_queue(&mut block, 0, vq.create_queue()); + serve(&mut block, &mem); + + // The prefix and the drive id, cut to the 15 bytes Linux keeps. + let names = || { + std::fs::read_dir("/proc/self/task") + .unwrap() + .filter_map(|task| std::fs::read_to_string(task.unwrap().path().join("comm")).ok()) + .map(|name| name.trim_end().to_string()) + .collect::>() + }; + wait_for("the worker's name", || { + names().contains(&"blk_vol_x7k2p9q".to_string()) + }); + } + + #[test] + fn test_drop_while_serving() { + let mut block = default_block(FileEngineType::Threaded); + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 16); + add_flush_requests_batch(&mut block, &vq, 4); + // Leave requests waiting on the rate limiter. + block.rate_limiter = RateLimiter::new(0, 0, 0, 1, 0, 10_000).unwrap(); + serve(&mut block, &mem); + wait_for("the first flush", || vq.used.idx.get() == 1); + drop(block); + } + + #[test] + fn test_segmented_parse_failures() { + let mem = default_mem(); + let vq = VirtQueue::new(GuestAddress(0), &mem, 64); + let parse = |request_type, sector, segments: &[(u64, u32)]| { + set_request(&vq, request_type, sector, segments); + let mut queue = vq.create_queue(); + let head = queue.pop().unwrap().unwrap(); + Request::parse_segmented(&head, &mem, 8) + }; + + // As many buffers as advertised, but no more. + let max = usize::try_from(THREADED_SEG_MAX).unwrap(); + let segments: Vec<(u64, u32)> = (0..=max).map(|i| (0x2000 + 512 * i as u64, 0)).collect(); + let request = parse(VIRTIO_BLK_T_OUT, 0, &segments[..max]).unwrap(); + assert_eq!(request.segments.len(), max); + assert!(matches!( + parse(VIRTIO_BLK_T_OUT, 0, &segments), + Err(VirtioBlockError::TooManySegments) + )); + + // The total length counts: a multiple of the sector size, within the disk. + parse(VIRTIO_BLK_T_OUT, 0, &[(0x2000, 256), (0x3000, 256)]).unwrap(); + assert!(matches!( + parse(VIRTIO_BLK_T_OUT, 0, &[(0x2000, 512), (0x3000, 256)]), + Err(VirtioBlockError::InvalidDataLength) + )); + assert!(matches!( + parse(VIRTIO_BLK_T_OUT, 6, &[(0x2000, 512), (0x3000, 1024)]), + Err(VirtioBlockError::InvalidOffset) + )); + assert!(matches!( + parse(VIRTIO_BLK_T_OUT, 0, &[(0x2000, u32::MAX), (0x3000, 512)]), + Err(VirtioBlockError::InvalidDataLength) + )); + // No data at all. + assert!(matches!( + parse(VIRTIO_BLK_T_IN, 0, &[]), + Err(VirtioBlockError::DescriptorChainTooShort) + )); + + // Every buffer has to have the direction of the request. + set_request(&vq, VIRTIO_BLK_T_IN, 0, &[(0x2000, 512), (0x3000, 512)]); + vq.dtable[2].flags.set(VIRTQ_DESC_F_NEXT); + let mut queue = vq.create_queue(); + let head = queue.pop().unwrap().unwrap(); + assert!(matches!( + Request::parse_segmented(&head, &mem, 8), + Err(VirtioBlockError::UnexpectedReadOnlyDescriptor) + )); + } + + #[test] + fn test_persistence() { + let f = TempFile::new().unwrap(); + f.as_file().set_len(0x1000).unwrap(); + let block = VirtioBlock::new(VirtioBlockConfig { + drive_id: "threaded".to_string(), + path_on_host: f.as_path().to_str().unwrap().to_string(), + is_root_device: false, + partuuid: None, + is_read_only: false, + cache_type: CacheType::Unsafe, + rate_limiter: None, + file_engine_type: FileEngineType::Threaded, + }) + .unwrap(); + + let mut snapshot = vec![0; 4096]; + crate::snapshot::Snapshot::new(block.save()) + .save(&mut snapshot.as_mut_slice()) + .unwrap(); + let restored = VirtioBlock::restore( + BlockConstructorArgs { mem: default_mem() }, + &crate::snapshot::Snapshot::load_without_crc_check(snapshot.as_slice()) + .unwrap() + .data, + ) + .unwrap(); + + assert_eq!(restored.file_engine_type(), FileEngineType::Threaded); + assert!(matches!(restored.disk.file_engine, FileEngine::Threaded(_))); + } +} diff --git a/src/vmm/src/rate_limiter/mod.rs b/src/vmm/src/rate_limiter/mod.rs index 97ebad51fcc..d97543b6433 100644 --- a/src/vmm/src/rate_limiter/mod.rs +++ b/src/vmm/src/rate_limiter/mod.rs @@ -236,6 +236,24 @@ impl TokenBucket { self.budget = std::cmp::min(self.budget.saturating_add(tokens), self.size); } + /// Returns how long until the budget covers `tokens`, assuming no other consumption. + /// + /// Meant to be called right after `reduce()` failed, when the budget has just been + /// replenished up to `last_update`. + fn refill_delay(&self, tokens: u64) -> Duration { + let deficit = u128::from(tokens.saturating_sub(self.budget)); + if deficit == 0 { + return Duration::ZERO; + } + let processed_capacity = u128::from(self.processed_capacity); + let processed_refill_time = u128::from(self.processed_refill_time); + // Round up so the timer never fires before the tokens are available. + let needed_ns = (deficit * processed_refill_time).div_ceil(processed_capacity); + // `auto_replenish()` carries sub-token time in `last_update`; it counts towards the wait. + let needed_ns = needed_ns.saturating_sub(self.last_update.elapsed().as_nanos()); + Duration::from_nanos(u64::try_from(needed_ns).unwrap_or(u64::MAX)) + } + /// Returns the capacity of the token bucket. pub fn capacity(&self) -> u64 { self.size @@ -304,6 +322,9 @@ pub struct RateLimiter { timer_fd: TimerFd, // Internal flag that quickly determines timer state. timer_active: bool, + // When set, a depleted bucket arms the timer for the time its tokens take to refill, + // bounded below by this value. When unset, the timer always waits the fixed interval. + min_refill_delay: Option, } impl PartialEq for RateLimiter { @@ -374,9 +395,23 @@ impl RateLimiter { ops: ops_token_bucket, timer_fd, timer_active: false, + min_refill_delay: None, }) } + /// Arms the refill timer for the time a depleted bucket needs to cover the failed request, + /// but no less than `min_delay` and no more than the fixed refill interval. + /// + /// Without this, the limiter always waits the fixed refill interval once a bucket is empty. + pub fn set_min_refill_delay(&mut self, min_delay: Duration) { + self.min_refill_delay = Some(min_delay); + } + + /// Returns the minimum refill delay, if adaptive refill timing is enabled. + pub fn min_refill_delay(&self) -> Option { + self.min_refill_delay + } + // Arm the timer of the rate limiter with the provided `TimerState`. fn activate_timer(&mut self, timer_state: TimerState) { // Register the timer; don't care about its previous state @@ -407,7 +442,16 @@ impl RateLimiter { // make sure there is only one running timer for this limiter. BucketReduction::Failure => { if !self.timer_active { - self.activate_timer(TIMER_REFILL_STATE); + let timer_state = match self.min_refill_delay { + Some(min_delay) => TimerState::Oneshot( + bucket + .refill_delay(tokens) + .max(min_delay) + .min(Duration::from_millis(REFILL_TIMER_INTERVAL_MS)), + ), + None => TIMER_REFILL_STATE, + }; + self.activate_timer(timer_state); } false } @@ -996,6 +1040,104 @@ pub(crate) mod tests { } } + #[test] + fn test_token_bucket_refill_delay() { + // 1 token per millisecond. + let mut tb = TokenBucket::new(1000, 0, 1000).unwrap(); + assert_eq!(tb.refill_delay(1000), Duration::ZERO); + + assert_eq!(tb.reduce(1000), BucketReduction::Success); + tb.last_update = Instant::now(); + let delay = tb.refill_delay(10); + assert!(delay <= Duration::from_millis(10), "{delay:?}"); + assert!(delay > Duration::from_millis(9), "{delay:?}"); + + // Time already accrued towards the next tokens shortens the wait. + tb.last_update = Instant::now() - Duration::from_millis(4); + let delay = tb.refill_delay(10); + assert!(delay <= Duration::from_millis(6), "{delay:?}"); + assert!(delay > Duration::from_millis(5), "{delay:?}"); + + // Only the deficit over the current budget needs to refill. + tb.force_replenish(8); + tb.last_update = Instant::now(); + let delay = tb.refill_delay(10); + assert!(delay <= Duration::from_millis(2), "{delay:?}"); + assert!(delay > Duration::from_millis(1), "{delay:?}"); + + // Rounds up to the nanosecond the token becomes available: 3 tokens per millisecond. + let mut tb = TokenBucket::new(3, 0, 1).unwrap(); + assert_eq!(tb.reduce(3), BucketReduction::Success); + // A future `last_update` means no time has accrued yet. + tb.last_update = Instant::now() + Duration::from_secs(1); + assert_eq!(tb.refill_delay(1), Duration::from_nanos(333_334)); + } + + fn armed_delay(l: &RateLimiter) -> Duration { + match l.timer_fd.get_state() { + TimerState::Oneshot(delay) => delay, + state => panic!("expected an armed oneshot timer, got {state:?}"), + } + } + + #[test] + fn test_rate_limiter_adaptive_refill_timer() { + let min_delay = Duration::from_millis(5); + let max_delay = Duration::from_millis(REFILL_TIMER_INTERVAL_MS); + + // Bandwidth of 1 byte per millisecond: 20 missing bytes refill in 20ms. + let mut l = RateLimiter::new(1000, 0, 1000, 0, 0, 0).unwrap(); + assert_eq!(l.min_refill_delay(), None); + l.set_min_refill_delay(min_delay); + assert_eq!(l.min_refill_delay(), Some(min_delay)); + assert!(l.consume(1000, TokenType::Bytes)); + assert!(!l.consume(20, TokenType::Bytes)); + assert!(l.is_blocked()); + let delay = armed_delay(&l); + assert!(delay <= Duration::from_millis(20), "{delay:?}"); + assert!(delay > Duration::from_millis(15), "{delay:?}"); + // The limiter unblocks once those tokens are available, well before 100ms. + thread::sleep(Duration::from_millis(30)); + l.event_handler().unwrap(); + assert!(!l.is_blocked()); + assert!(l.consume(20, TokenType::Bytes)); + + // A tiny deficit still waits for the minimum delay. + let mut l = RateLimiter::new(1000, 0, 1000, 0, 0, 0).unwrap(); + l.set_min_refill_delay(min_delay); + assert!(l.consume(1000, TokenType::Bytes)); + assert!(!l.consume(1, TokenType::Bytes)); + let delay = armed_delay(&l); + assert!(delay <= min_delay, "{delay:?}"); + assert!(delay > min_delay - Duration::from_millis(2), "{delay:?}"); + + // A large deficit waits no longer than the fixed refill interval. + let mut l = RateLimiter::new(1000, 0, 1000, 0, 0, 0).unwrap(); + l.set_min_refill_delay(min_delay); + assert!(l.consume(1000, TokenType::Bytes)); + assert!(!l.consume(500, TokenType::Bytes)); + let delay = armed_delay(&l); + assert!(delay <= max_delay, "{delay:?}"); + assert!(delay > max_delay - Duration::from_millis(5), "{delay:?}"); + + // The ops bucket arms the timer from its own refill rate. + let mut l = RateLimiter::new(0, 0, 0, 1000, 0, 1000).unwrap(); + l.set_min_refill_delay(min_delay); + assert!(l.consume(1000, TokenType::Ops)); + assert!(!l.consume(20, TokenType::Ops)); + let delay = armed_delay(&l); + assert!(delay <= Duration::from_millis(20), "{delay:?}"); + assert!(delay > Duration::from_millis(15), "{delay:?}"); + + // Without a minimum refill delay, the fixed interval is kept. + let mut l = RateLimiter::new(1000, 0, 1000, 0, 0, 0).unwrap(); + assert!(l.consume(1000, TokenType::Bytes)); + assert!(!l.consume(1, TokenType::Bytes)); + let delay = armed_delay(&l); + assert!(delay <= max_delay, "{delay:?}"); + assert!(delay > max_delay - Duration::from_millis(5), "{delay:?}"); + } + #[test] fn test_rate_limiter_bandwidth() { // rate limiter with limit of 1000 bytes/s diff --git a/src/vmm/src/rate_limiter/persist.rs b/src/vmm/src/rate_limiter/persist.rs index 6c9e5052ecf..ee400a52850 100644 --- a/src/vmm/src/rate_limiter/persist.rs +++ b/src/vmm/src/rate_limiter/persist.rs @@ -84,6 +84,8 @@ impl Persist<'_> for RateLimiter { }, timer_fd: TimerFd::new_custom(ClockId::Monotonic, true, true)?, timer_active: false, + // Device policy, not limiter state: the owning device sets it after restore. + min_refill_delay: None, }; Ok(rate_limiter)