Skip to main content

storvsp/
lib.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! VMBus SCSI controller emulator (StorVSP).
5//!
6//! StorVSP implements the Hyper-V synthetic SCSI protocol — a VMBus-based
7//! transport that carries SCSI CDBs between the guest's `storvsc` driver and
8//! the VMM. This is not a standard SCSI transport (like iSCSI or SAS); it's a
9//! Hyper-V-specific wire format defined in [`storvsp_protocol`].
10//!
11//! # Architecture
12//!
13//! The crate uses a multi-worker model. The primary VMBus channel handles
14//! protocol version negotiation (Win6 through Blue); sub-channels process I/O
15//! in parallel. Each worker owns a VMBus ring and processes packets
16//! concurrently via `FuturesUnordered`.
17//!
18//! StorVSP handles the transport (ring buffer management, GPADL setup, packet
19//! framing, sub-channel lifecycle) and a few SCSI control commands directly
20//! (`REPORT_LUNS`, `INQUIRY` for absent targets). All actual I/O is delegated
21//! to [`AsyncScsiDisk`] implementations — StorVSP
22//! never interprets SCSI data CDBs itself.
23//!
24//! For the channel/sub-channel model, CPU affinity, and performance
25//! characteristics, see the
26//! [StorVSP Channels & Subchannels](https://openvmm.dev/reference/devices/vmbus/storvsp_channels.html)
27//! page in the OpenVMM Guide.
28//!
29//! # Key types
30//!
31//! - [`StorageDevice`] — the VMBus device. Implements `VmbusDevice` and
32//!   `SaveRestoreVmbusDevice`.
33//! - [`ScsiController`] — manages attached disks by [`ScsiPath`]. Supports
34//!   runtime attach/remove.
35//! - [`ScsiControllerDisk`] — wraps `Arc<dyn AsyncScsiDisk>`.
36//!
37//! # Performance
38//!
39//! Poll-mode optimization: when pending I/O count exceeds
40//! `poll_mode_queue_depth`, the worker switches from interrupt-driven to
41//! busy-poll for new requests, reducing guest exit frequency. Future storage
42//! for SCSI request processing is pooled to avoid allocation on the hot path.
43
44#![expect(missing_docs)]
45#![forbid(unsafe_code)]
46
47#[cfg(feature = "ioperf")]
48pub mod ioperf;
49
50#[cfg(feature = "test")]
51pub mod test_helpers;
52
53#[cfg(not(feature = "test"))]
54mod test_helpers;
55
56pub mod resolver;
57mod save_restore;
58
59use crate::ring::gparange::MultiPagedRangeBuf;
60use anyhow::Context as _;
61use async_trait::async_trait;
62use fast_select::FastSelect;
63use futures::FutureExt;
64use futures::StreamExt;
65use futures::select_biased;
66use guestmem::AccessError;
67use guestmem::GuestMemory;
68use guestmem::MemoryRead;
69use guestmem::MemoryWrite;
70use guestmem::ranges::PagedRange;
71use guid::Guid;
72use inspect::Inspect;
73use inspect::InspectMut;
74use inspect_counters::Counter;
75use inspect_counters::Histogram;
76use oversized_box::OversizedBox;
77use parking_lot::Mutex;
78use parking_lot::RwLock;
79use ring::OutgoingPacketType;
80use scsi::AdditionalSenseCode;
81use scsi::ScsiOp;
82use scsi::ScsiStatus;
83use scsi::srb::SrbStatus;
84use scsi::srb::SrbStatusAndFlags;
85use scsi_buffers::RequestBuffers;
86use scsi_core::AsyncScsiDisk;
87use scsi_core::Request;
88use scsi_core::ScsiResult;
89use scsi_defs as scsi;
90use scsidisk::illegal_request_sense;
91use slab::Slab;
92use std::collections::hash_map::Entry;
93use std::collections::hash_map::HashMap;
94use std::fmt::Debug;
95use std::future::Future;
96use std::future::poll_fn;
97use std::pin::Pin;
98use std::sync::Arc;
99use std::sync::atomic::AtomicU32;
100use std::sync::atomic::Ordering::Relaxed;
101use std::task::Context;
102use std::task::Poll;
103use storvsp_resources::ScsiPath;
104use task_control::AsyncRun;
105use task_control::InspectTask;
106use task_control::StopTask;
107use task_control::TaskControl;
108use thiserror::Error;
109use tracing_helpers::ErrorValueExt;
110use unicycle::FuturesUnordered;
111use vmbus_async::queue;
112use vmbus_async::queue::ExternalDataError;
113use vmbus_async::queue::IncomingPacket;
114use vmbus_async::queue::OutgoingPacket;
115use vmbus_async::queue::Queue;
116use vmbus_channel::RawAsyncChannel;
117use vmbus_channel::bus::ChannelType;
118use vmbus_channel::bus::OfferParams;
119use vmbus_channel::bus::OpenRequest;
120use vmbus_channel::channel::ChannelControl;
121use vmbus_channel::channel::ChannelOpenError;
122use vmbus_channel::channel::DeviceResources;
123use vmbus_channel::channel::RestoreControl;
124use vmbus_channel::channel::SaveRestoreVmbusDevice;
125use vmbus_channel::channel::VmbusDevice;
126use vmbus_channel::gpadl_ring::GpadlRingMem;
127use vmbus_channel::gpadl_ring::gpadl_channel;
128use vmbus_core::protocol::UserDefinedData;
129use vmbus_ring as ring;
130use vmbus_ring::RingMem;
131use vmcore::save_restore::RestoreError;
132use vmcore::save_restore::SaveError;
133use vmcore::save_restore::SavedStateBlob;
134use vmcore::vm_task::VmTaskDriver;
135use vmcore::vm_task::VmTaskDriverSource;
136use zerocopy::FromBytes;
137use zerocopy::FromZeros;
138use zerocopy::Immutable;
139use zerocopy::IntoBytes;
140use zerocopy::KnownLayout;
141
142/// The IO queue depth at which the controller switches from guest-signal-driven
143/// to poll-mode operation. This optimization reduces the guest exit rate by
144/// relying on (typically-interrupt-driven) IO completions to drive polling for
145/// new IO requests.
146const DEFAULT_POLL_MODE_QUEUE_DEPTH: u32 = 1;
147
148pub struct StorageDevice {
149    instance_id: Guid,
150    ide_path: Option<ScsiPath>,
151    workers: Vec<WorkerAndDriver>,
152    controller: Arc<ScsiControllerState>,
153    resources: DeviceResources,
154    driver_source: VmTaskDriverSource,
155    max_sub_channel_count: u16,
156    protocol: Arc<Protocol>,
157    io_queue_depth: u32,
158}
159
160#[derive(Inspect)]
161struct WorkerAndDriver {
162    #[inspect(flatten)]
163    worker: TaskControl<WorkerState, Worker>,
164    driver: VmTaskDriver,
165}
166
167struct WorkerState;
168
169impl InspectMut for StorageDevice {
170    fn inspect_mut(&mut self, req: inspect::Request<'_>) {
171        let mut resp = req.respond();
172
173        let disks = self.controller.disks.read();
174        for (path, controller_disk) in disks.iter() {
175            resp.child(&format!("disks/{}", path), |req| {
176                controller_disk.disk.inspect(req);
177            });
178        }
179
180        resp.fields(
181            "channels",
182            self.workers
183                .iter()
184                .filter(|task| task.worker.has_state())
185                .enumerate(),
186        )
187        .field(
188            "poll_mode_queue_depth",
189            inspect::AtomicMut(&self.controller.poll_mode_queue_depth),
190        );
191    }
192}
193
194struct Worker<T: RingMem = GpadlRingMem> {
195    inner: WorkerInner,
196    rescan_notification: futures::channel::mpsc::Receiver<()>,
197    fast_select: FastSelect,
198    queue: Queue<T>,
199}
200
201struct Protocol {
202    state: RwLock<ProtocolState>,
203    /// Signaled when `state` transitions to `ProtocolState::Ready`.
204    ready: event_listener::Event,
205}
206
207struct WorkerInner {
208    protocol: Arc<Protocol>,
209    request_size: usize,
210    controller: Arc<ScsiControllerState>,
211    channel_index: u16,
212    scsi_queue: Arc<ScsiCommandQueue>,
213    scsi_requests: FuturesUnordered<ScsiRequest>,
214    scsi_requests_states: Slab<ScsiRequestState>,
215    full_request_pool: Vec<Arc<ScsiRequestAndRange>>,
216    future_pool: Vec<OversizedBox<(), ScsiOpStorage>>,
217    channel_control: ChannelControl,
218    max_io_queue_depth: usize,
219    stats: WorkerStats,
220}
221
222#[derive(Debug, Default, Inspect)]
223struct WorkerStats {
224    ios_submitted: Counter,
225    ios_completed: Counter,
226    wakes: Counter,
227    wakes_spurious: Counter,
228    per_wake_submissions: Histogram<10>,
229    per_wake_completions: Histogram<10>,
230}
231
232#[repr(u16)]
233#[derive(Copy, Clone, Debug, Inspect, PartialEq, Eq, PartialOrd, Ord)]
234enum Version {
235    Win6 = storvsp_protocol::VERSION_WIN6,
236    Win7 = storvsp_protocol::VERSION_WIN7,
237    Win8 = storvsp_protocol::VERSION_WIN8,
238    Blue = storvsp_protocol::VERSION_BLUE,
239}
240
241#[derive(Debug, Error)]
242#[error("protocol version {0:#x} not supported")]
243struct UnsupportedVersion(u16);
244
245impl Version {
246    fn parse(major_minor: u16) -> Result<Self, UnsupportedVersion> {
247        let version = match major_minor {
248            storvsp_protocol::VERSION_WIN6 => Self::Win6,
249            storvsp_protocol::VERSION_WIN7 => Self::Win7,
250            storvsp_protocol::VERSION_WIN8 => Self::Win8,
251            storvsp_protocol::VERSION_BLUE => Self::Blue,
252            version => return Err(UnsupportedVersion(version)),
253        };
254        assert_eq!(version as u16, major_minor);
255        Ok(version)
256    }
257
258    fn max_request_size(&self) -> usize {
259        match self {
260            Version::Win8 | Version::Blue => storvsp_protocol::SCSI_REQUEST_LEN_V2,
261            Version::Win6 | Version::Win7 => storvsp_protocol::SCSI_REQUEST_LEN_V1,
262        }
263    }
264}
265
266#[derive(Copy, Clone)]
267enum ProtocolState {
268    Init(InitState),
269    Ready {
270        version: Version,
271        subchannel_count: u16,
272    },
273}
274
275#[derive(Copy, Clone, Debug)]
276enum InitState {
277    Begin,
278    QueryVersion,
279    QueryProperties {
280        version: Version,
281    },
282    EndInitialization {
283        version: Version,
284        subchannel_count: Option<u16>,
285    },
286}
287
288/// The internal SCSI operation future type.
289///
290/// This is a boxed future of a large pre-determined size. The box is reused
291/// after a SCSI request completes to avoid allocations in the hot path.
292///
293/// An Option type is used so that the future can be efficiently dropped (via
294/// `Pin::set(x, None)`) before it is stashed away for reuse.
295type ScsiOpStorage = [u64; SCSI_REQUEST_STACK_SIZE / 8];
296type ScsiOpFuture = Pin<OversizedBox<dyn Future<Output = ScsiResult> + Send, ScsiOpStorage>>;
297
298/// The amount of space reserved for a ScsiOpFuture.
299///
300/// This was chosen by running `cargo test -p storvsp -- --no-capture` and looking at the required
301/// size that was given in the failure message
302const SCSI_REQUEST_STACK_SIZE: usize = scsi_core::ASYNC_SCSI_DISK_STACK_SIZE + 272;
303
304struct ScsiRequest {
305    request_id: usize,
306    future: Option<ScsiOpFuture>,
307}
308
309impl ScsiRequest {
310    fn new(
311        request_id: usize,
312        future: OversizedBox<dyn Future<Output = ScsiResult> + Send, ScsiOpStorage>,
313    ) -> Self {
314        Self {
315            request_id,
316            future: Some(future.into()),
317        }
318    }
319}
320
321impl Future for ScsiRequest {
322    type Output = (usize, ScsiResult, OversizedBox<(), ScsiOpStorage>);
323
324    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
325        let this = self.get_mut();
326        let future = this.future.as_mut().unwrap().as_mut();
327        let result = std::task::ready!(future.poll(cx));
328        // Return the future so that its storage can be reused.
329        let future = this.future.take().unwrap();
330        Poll::Ready((this.request_id, result, OversizedBox::empty_pinned(future)))
331    }
332}
333
334#[derive(Debug, Error)]
335enum WorkerError {
336    #[error("packet error")]
337    PacketError(#[source] PacketError),
338    #[error("queue error")]
339    Queue(#[source] queue::Error),
340    #[error("queue should have enough space but no longer does")]
341    NotEnoughSpace,
342}
343
344#[derive(Debug, Error)]
345enum PacketError {
346    #[error("Not transactional")]
347    NotTransactional,
348    #[error("Unrecognized operation {0:?}")]
349    UnrecognizedOperation(storvsp_protocol::Operation),
350    #[error("Invalid packet type")]
351    InvalidPacketType,
352    #[error("Invalid data transfer length")]
353    InvalidDataTransferLength,
354    #[error("Access error: {0}")]
355    Access(#[source] AccessError),
356    #[error("Range error")]
357    Range(#[source] ExternalDataError),
358}
359
360#[derive(Debug, Default, Clone)]
361struct Range {
362    len: usize,
363    is_write: bool,
364}
365
366impl Range {
367    fn new(buf: &MultiPagedRangeBuf, request: &storvsp_protocol::ScsiRequest) -> Option<Self> {
368        let len = request.data_transfer_length as usize;
369        let is_write = request.data_in != 0;
370        // Ensure there is exactly one range and it's large enough, or there are
371        // zero ranges and there is no associated SCSI buffer.
372        if buf.range_count() > 1 || (len > 0 && buf.first()?.len() < len) {
373            return None;
374        }
375        Some(Self { len, is_write })
376    }
377
378    fn buffer<'a>(
379        &'a self,
380        buf: &'a MultiPagedRangeBuf,
381        guest_memory: &'a GuestMemory,
382    ) -> RequestBuffers<'a> {
383        let mut range = buf.first().unwrap_or_else(PagedRange::empty);
384        range.truncate(self.len);
385        RequestBuffers::new(guest_memory, range, self.is_write)
386    }
387}
388
389#[derive(Debug)]
390struct Packet {
391    data: PacketData,
392    transaction_id: u64,
393    request_size: usize,
394}
395
396#[derive(Debug)]
397enum PacketData {
398    BeginInitialization,
399    EndInitialization,
400    QueryProtocolVersion(u16),
401    QueryProperties,
402    CreateSubChannels(u16),
403    ExecuteScsi(Arc<ScsiRequestAndRange>),
404    ResetBus,
405    ResetAdapter,
406    ResetLun,
407}
408
409fn parse_packet<T: RingMem>(
410    packet: &IncomingPacket<'_, T>,
411    pool: &mut Vec<Arc<ScsiRequestAndRange>>,
412) -> Result<Packet, PacketError> {
413    let packet = match packet {
414        IncomingPacket::Completion(_) => return Err(PacketError::InvalidPacketType),
415        IncomingPacket::Data(packet) => packet,
416    };
417    let transaction_id = packet
418        .transaction_id()
419        .ok_or(PacketError::NotTransactional)?;
420
421    let mut reader = packet.reader();
422    let header: storvsp_protocol::Packet = reader.read_plain().map_err(PacketError::Access)?;
423    // You would expect that this should be limited to the current protocol
424    // version's maximum packet size, but this is not what Hyper-V does, and
425    // Linux 6.1 relies on this behavior during protocol initialization.
426    let request_size = reader.len().min(storvsp_protocol::SCSI_REQUEST_LEN_MAX);
427    let data = match header.operation {
428        storvsp_protocol::Operation::BEGIN_INITIALIZATION => PacketData::BeginInitialization,
429        storvsp_protocol::Operation::END_INITIALIZATION => PacketData::EndInitialization,
430        storvsp_protocol::Operation::QUERY_PROTOCOL_VERSION => {
431            let mut version = storvsp_protocol::ProtocolVersion::new_zeroed();
432            reader
433                .read(version.as_mut_bytes())
434                .map_err(PacketError::Access)?;
435            PacketData::QueryProtocolVersion(version.major_minor)
436        }
437        storvsp_protocol::Operation::QUERY_PROPERTIES => PacketData::QueryProperties,
438        storvsp_protocol::Operation::EXECUTE_SRB => {
439            let mut full_request = pool.pop().unwrap_or_else(|| {
440                Arc::new(ScsiRequestAndRange {
441                    external_data: Range::default(),
442                    external_data_buf: MultiPagedRangeBuf::new(),
443                    request: storvsp_protocol::ScsiRequest::new_zeroed(),
444                    request_size,
445                })
446            });
447
448            {
449                let full_request = Arc::get_mut(&mut full_request).unwrap();
450                let request_buf = &mut full_request.request.as_mut_bytes()[..request_size];
451                reader.read(request_buf).map_err(PacketError::Access)?;
452
453                full_request.external_data_buf.clear();
454                packet
455                    .read_external_ranges(&mut full_request.external_data_buf)
456                    .map_err(PacketError::Range)?;
457
458                full_request.external_data =
459                    Range::new(&full_request.external_data_buf, &full_request.request)
460                        .ok_or(PacketError::InvalidDataTransferLength)?;
461            }
462
463            PacketData::ExecuteScsi(full_request)
464        }
465        storvsp_protocol::Operation::RESET_LUN => PacketData::ResetLun,
466        storvsp_protocol::Operation::RESET_ADAPTER => PacketData::ResetAdapter,
467        storvsp_protocol::Operation::RESET_BUS => PacketData::ResetBus,
468        storvsp_protocol::Operation::CREATE_SUB_CHANNELS => {
469            let mut sub_channel_count: u16 = 0;
470            reader
471                .read(sub_channel_count.as_mut_bytes())
472                .map_err(PacketError::Access)?;
473            PacketData::CreateSubChannels(sub_channel_count)
474        }
475        _ => return Err(PacketError::UnrecognizedOperation(header.operation)),
476    };
477
478    if let PacketData::ExecuteScsi(_) = data {
479        tracing::trace!(transaction_id, ?data, "parse_packet");
480    } else {
481        tracing::debug!(transaction_id, ?data, "parse_packet");
482    }
483
484    Ok(Packet {
485        data,
486        request_size,
487        transaction_id,
488    })
489}
490
491impl WorkerInner {
492    fn send_vmbus_packet<M: RingMem>(
493        &mut self,
494        writer: &mut queue::WriteBatch<'_, M>,
495        packet_type: OutgoingPacketType<'_>,
496        request_size: usize,
497        transaction_id: u64,
498        operation: storvsp_protocol::Operation,
499        status: storvsp_protocol::NtStatus,
500        payload: &[u8],
501    ) -> Result<(), WorkerError> {
502        let header = storvsp_protocol::Packet {
503            operation,
504            flags: 0,
505            status,
506        };
507
508        let packet_size = size_of_val(&header) + request_size;
509
510        // Zero pad or truncate the payload to the queue's packet size. This is
511        // necessary because Windows guests check that each packet's size is
512        // exactly the largest possible packet size for the negotiated protocol
513        // version.
514        let len = size_of_val(&header) + size_of_val(payload);
515        let padding = [0; storvsp_protocol::SCSI_REQUEST_LEN_MAX];
516        let (payload_bytes, padding_bytes) = if len > packet_size {
517            (&payload[..packet_size - size_of_val(&header)], &[][..])
518        } else {
519            (payload, &padding[..packet_size - len])
520        };
521        assert_eq!(
522            size_of_val(&header) + payload_bytes.len() + padding_bytes.len(),
523            packet_size
524        );
525        writer
526            .try_write(&OutgoingPacket {
527                transaction_id,
528                packet_type,
529                payload: &[header.as_bytes(), payload_bytes, padding_bytes],
530            })
531            .map_err(|err| match err {
532                queue::TryWriteError::Full(_) => WorkerError::NotEnoughSpace,
533                queue::TryWriteError::Queue(err) => WorkerError::Queue(err),
534            })
535    }
536
537    fn send_packet<M: RingMem, P: IntoBytes + Immutable + KnownLayout>(
538        &mut self,
539        writer: &mut queue::WriteHalf<'_, M>,
540        operation: storvsp_protocol::Operation,
541        status: storvsp_protocol::NtStatus,
542        payload: &P,
543    ) -> Result<(), WorkerError> {
544        self.send_vmbus_packet(
545            &mut writer.batched(),
546            OutgoingPacketType::InBandNoCompletion,
547            self.request_size,
548            0,
549            operation,
550            status,
551            payload.as_bytes(),
552        )
553    }
554
555    fn send_completion<M: RingMem, P: IntoBytes + Immutable + KnownLayout>(
556        &mut self,
557        writer: &mut queue::WriteHalf<'_, M>,
558        packet: &Packet,
559        status: storvsp_protocol::NtStatus,
560        payload: &P,
561    ) -> Result<(), WorkerError> {
562        self.send_vmbus_packet(
563            &mut writer.batched(),
564            OutgoingPacketType::Completion,
565            packet.request_size,
566            packet.transaction_id,
567            storvsp_protocol::Operation::COMPLETE_IO,
568            status,
569            payload.as_bytes(),
570        )
571    }
572}
573
574struct ScsiCommandQueue {
575    controller: Arc<ScsiControllerState>,
576    mem: GuestMemory,
577    force_path_id: Option<u8>,
578}
579
580impl ScsiCommandQueue {
581    async fn execute_scsi(&self, full_request: &ScsiRequestAndRange) -> ScsiResult {
582        let request = &full_request.request;
583        let op = ScsiOp(request.payload[0]);
584        let external_data = full_request
585            .external_data
586            .buffer(&full_request.external_data_buf, &self.mem);
587
588        tracing::trace!(
589            path_id = request.path_id,
590            target_id = request.target_id,
591            lun = request.lun,
592            op = ?op,
593            "execute_scsi start...",
594        );
595
596        let path_id = self.force_path_id.unwrap_or(request.path_id);
597
598        let controller_disk = self
599            .controller
600            .disks
601            .read()
602            .get(&ScsiPath {
603                path: path_id,
604                target: request.target_id,
605                lun: request.lun,
606            })
607            .cloned();
608
609        let result = match op {
610            ScsiOp::REPORT_LUNS => {
611                const HEADER_SIZE: usize = size_of::<scsi::LunList>();
612                let mut luns: Vec<u8> = self
613                    .controller
614                    .disks
615                    .read()
616                    .keys()
617                    .flat_map(|path| {
618                        // Use the original path ID and not the forced one to
619                        // match Hyper-V storvsp behavior.
620                        if request.path_id == path.path && request.target_id == path.target {
621                            Some(path.lun)
622                        } else {
623                            None
624                        }
625                    })
626                    .collect();
627                luns.sort_unstable();
628                let mut data: Vec<u64> = vec![0; luns.len() + 1];
629                let header = scsi::LunList {
630                    length: (luns.len() as u32 * 8).into(),
631                    reserved: [0; 4],
632                };
633                data.as_mut_bytes()[..HEADER_SIZE].copy_from_slice(header.as_bytes());
634                for (i, lun) in luns.iter().enumerate() {
635                    data[i + 1].as_mut_bytes()[..2].copy_from_slice(&(*lun as u16).to_be_bytes());
636                }
637                if external_data.len() >= HEADER_SIZE {
638                    let tx = std::cmp::min(external_data.len(), data.as_bytes().len());
639                    external_data.writer().write(&data.as_bytes()[..tx]).map_or(
640                        ScsiResult {
641                            scsi_status: ScsiStatus::CHECK_CONDITION,
642                            srb_status: SrbStatus::INVALID_REQUEST,
643                            tx: 0,
644                            sense_data: Some(illegal_request_sense(
645                                AdditionalSenseCode::INVALID_CDB,
646                            )),
647                        },
648                        |_| ScsiResult {
649                            scsi_status: ScsiStatus::GOOD,
650                            srb_status: SrbStatus::SUCCESS,
651                            tx,
652                            sense_data: None,
653                        },
654                    )
655                } else {
656                    ScsiResult {
657                        scsi_status: ScsiStatus::GOOD,
658                        srb_status: SrbStatus::SUCCESS,
659                        tx: 0,
660                        sense_data: None,
661                    }
662                }
663            }
664            _ if controller_disk.is_some() => {
665                let mut cdb = [0; 16];
666                cdb.copy_from_slice(&request.payload[0..storvsp_protocol::CDB16GENERIC_LENGTH]);
667                controller_disk
668                    .unwrap()
669                    .disk
670                    .execute_scsi(
671                        &external_data,
672                        &Request {
673                            cdb,
674                            srb_flags: request.srb_flags,
675                        },
676                    )
677                    .await
678            }
679            ScsiOp::INQUIRY => {
680                let cdb = scsi::CdbInquiry::ref_from_prefix(&request.payload)
681                    .unwrap()
682                    .0; // TODO: zerocopy: ref-from-prefix: use-rest-of-range (https://github.com/microsoft/openvmm/issues/759)
683                if external_data.len() < cdb.allocation_length.get() as usize
684                    || request.data_in != storvsp_protocol::SCSI_IOCTL_DATA_IN
685                    || (cdb.allocation_length.get() as usize) < size_of::<scsi::InquiryDataHeader>()
686                {
687                    ScsiResult {
688                        scsi_status: ScsiStatus::CHECK_CONDITION,
689                        srb_status: SrbStatus::INVALID_REQUEST,
690                        tx: 0,
691                        sense_data: Some(illegal_request_sense(AdditionalSenseCode::INVALID_CDB)),
692                    }
693                } else {
694                    let enable_vpd = cdb.flags.vpd();
695                    if enable_vpd || cdb.page_code != 0 {
696                        // cannot support VPD inquiry for non-existing device (lun).
697                        ScsiResult {
698                            scsi_status: ScsiStatus::CHECK_CONDITION,
699                            srb_status: SrbStatus::INVALID_REQUEST,
700                            tx: 0,
701                            sense_data: Some(illegal_request_sense(
702                                AdditionalSenseCode::INVALID_CDB,
703                            )),
704                        }
705                    } else {
706                        const LOGICAL_UNIT_NOT_PRESENT_DEVICE: u8 = 0x7F;
707                        let mut data = scsidisk::INQUIRY_DATA_TEMPLATE;
708                        data.header.device_type = LOGICAL_UNIT_NOT_PRESENT_DEVICE;
709
710                        if request.lun != 0 {
711                            // Below fields are only set for lun0 inquiry so zero out here.
712                            data.vendor_id = [0; 8];
713                            data.product_id = [0; 16];
714                            data.product_revision_level = [0; 4];
715                        }
716
717                        let datab = data.as_bytes();
718                        let tx = std::cmp::min(
719                            cdb.allocation_length.get() as usize,
720                            size_of::<scsi::InquiryData>(),
721                        );
722                        external_data.writer().write(&datab[..tx]).map_or(
723                            ScsiResult {
724                                scsi_status: ScsiStatus::CHECK_CONDITION,
725                                srb_status: SrbStatus::INVALID_REQUEST,
726                                tx: 0,
727                                sense_data: Some(illegal_request_sense(
728                                    AdditionalSenseCode::INVALID_CDB,
729                                )),
730                            },
731                            |_| ScsiResult {
732                                scsi_status: ScsiStatus::GOOD,
733                                srb_status: SrbStatus::SUCCESS,
734                                tx,
735                                sense_data: None,
736                            },
737                        )
738                    }
739                }
740            }
741            _ => ScsiResult {
742                scsi_status: ScsiStatus::CHECK_CONDITION,
743                srb_status: SrbStatus::INVALID_LUN,
744                tx: 0,
745                sense_data: None,
746            },
747        };
748
749        tracing::trace!(
750            path_id = request.path_id,
751            target_id = request.target_id,
752            lun = request.lun,
753            op = ?op,
754            result = ?result,
755            "execute_scsi completed.",
756        );
757        result
758    }
759}
760
761impl<T: RingMem + 'static> Worker<T> {
762    fn new(
763        controller: Arc<ScsiControllerState>,
764        channel: RawAsyncChannel<T>,
765        channel_index: u16,
766        mem: GuestMemory,
767        channel_control: ChannelControl,
768        io_queue_depth: u32,
769        protocol: Arc<Protocol>,
770        force_path_id: Option<u8>,
771    ) -> anyhow::Result<Self> {
772        let queue = Queue::new(channel)?;
773        #[expect(clippy::disallowed_methods)] // TODO
774        let (source, target) = futures::channel::mpsc::channel(1);
775        controller.add_rescan_notification_source(source);
776
777        let max_io_queue_depth = io_queue_depth.max(1) as usize;
778        Ok(Self {
779            inner: WorkerInner {
780                protocol,
781                request_size: storvsp_protocol::SCSI_REQUEST_LEN_V1,
782                controller: controller.clone(),
783                channel_index,
784                scsi_queue: Arc::new(ScsiCommandQueue {
785                    controller,
786                    mem,
787                    force_path_id,
788                }),
789                scsi_requests: FuturesUnordered::new(),
790                scsi_requests_states: Slab::with_capacity(max_io_queue_depth),
791                channel_control,
792                max_io_queue_depth,
793                future_pool: Vec::new(),
794                full_request_pool: Vec::new(),
795                stats: Default::default(),
796            },
797            queue,
798            rescan_notification: target,
799            fast_select: FastSelect::new(),
800        })
801    }
802
803    async fn wait_for_scsi_requests_complete(&mut self) {
804        tracing::debug!(
805            channel_index = self.inner.channel_index,
806            "wait for IOs completed..."
807        );
808        while let Some((id, _, _)) = self.inner.scsi_requests.next().await {
809            self.inner.scsi_requests_states.remove(id);
810        }
811    }
812}
813
814impl InspectTask<Worker> for WorkerState {
815    fn inspect(&self, req: inspect::Request<'_>, worker: Option<&Worker>) {
816        if let Some(worker) = worker {
817            let mut resp = req.respond();
818            if worker.inner.channel_index == 0 {
819                let (state, version, subchannel_count) = match *worker.inner.protocol.state.read() {
820                    ProtocolState::Init(state) => match state {
821                        InitState::Begin => ("begin_init", None, None),
822                        InitState::QueryVersion => ("query_version", None, None),
823                        InitState::QueryProperties { version } => {
824                            ("query_properties", Some(version), None)
825                        }
826                        InitState::EndInitialization {
827                            version,
828                            subchannel_count,
829                        } => ("end_init", Some(version), subchannel_count),
830                    },
831                    ProtocolState::Ready {
832                        version,
833                        subchannel_count,
834                    } => ("ready", Some(version), Some(subchannel_count)),
835                };
836                resp.field("state", state)
837                    .field("version", version)
838                    .field("subchannel_count", subchannel_count);
839            }
840            resp.field("pending_packets", worker.inner.scsi_requests_states.len())
841                .fields("io", worker.inner.scsi_requests_states.iter())
842                .field("stats", &worker.inner.stats)
843                .field("ring", &worker.queue)
844                .field("max_io_queue_depth", worker.inner.max_io_queue_depth);
845        }
846    }
847}
848
849impl<T: 'static + Send + Sync + RingMem> AsyncRun<Worker<T>> for WorkerState {
850    async fn run(
851        &mut self,
852        stop: &mut StopTask<'_>,
853        worker: &mut Worker<T>,
854    ) -> Result<(), task_control::Cancelled> {
855        let fut = async {
856            if worker.inner.channel_index == 0 {
857                worker.process_primary().await
858            } else {
859                // Wait for initialization to end before processing any
860                // subchannel packets.
861                let protocol_version = loop {
862                    let listener = worker.inner.protocol.ready.listen();
863                    if let ProtocolState::Ready { version, .. } =
864                        *worker.inner.protocol.state.read()
865                    {
866                        break version;
867                    }
868                    tracing::debug!("subchannel waiting for initialization to end");
869                    listener.await
870                };
871                worker
872                    .inner
873                    .process_ready(&mut worker.queue, protocol_version)
874                    .await
875            }
876        };
877
878        match stop.until_stopped(fut).await? {
879            Ok(_) => {}
880            Err(e) => tracing::error!(error = e.as_error(), "process_packets error"),
881        }
882        Ok(())
883    }
884}
885
886impl WorkerInner {
887    /// Awaits the next incoming packet, without checking for any other events (device add/remove notifications or available completions).
888    /// Increments the count of outstanding packets when returning `Ok(Packet)`.
889    async fn next_packet<'a, M: RingMem>(
890        &mut self,
891        reader: &'a mut queue::ReadHalf<'a, M>,
892    ) -> Result<Packet, WorkerError> {
893        let packet = reader.read().await.map_err(WorkerError::Queue)?;
894        let stor_packet =
895            parse_packet(&packet, &mut self.full_request_pool).map_err(WorkerError::PacketError)?;
896        Ok(stor_packet)
897    }
898
899    /// Polls for enough ring space in the outgoing ring to send a packet.
900    ///
901    /// This is used to ensure there is enough space in the ring before
902    /// committing to sending a packet. This avoids the need to save pending
903    /// packets on the side if queue processing is interrupted while the ring is
904    /// full.
905    fn poll_for_ring_space<M: RingMem>(
906        &mut self,
907        cx: &mut Context<'_>,
908        writer: &mut queue::WriteHalf<'_, M>,
909    ) -> Poll<Result<(), WorkerError>> {
910        writer
911            .poll_ready(cx, MAX_VMBUS_PACKET_SIZE)
912            .map_err(WorkerError::Queue)
913    }
914}
915
916const MAX_VMBUS_PACKET_SIZE: usize = ring::PacketSize::in_band(
917    size_of::<storvsp_protocol::Packet>() + storvsp_protocol::SCSI_REQUEST_LEN_MAX,
918);
919
920impl<T: RingMem> Worker<T> {
921    /// Processes the protocol state machine.
922    async fn process_primary(&mut self) -> Result<(), WorkerError> {
923        loop {
924            let current_state = *self.inner.protocol.state.read();
925            match current_state {
926                ProtocolState::Ready { version, .. } => {
927                    break loop {
928                        select_biased! {
929                            r = self.inner.process_ready(&mut self.queue, version).fuse() => break r,
930                            _ = self.fast_select.select((self.rescan_notification.select_next_some(),)).fuse() => {
931                                if version >= Version::Win7
932                                {
933                                    tracing::debug!("rescan notification received, sending ENUMERATE_BUS");
934                                    self.inner.send_packet(&mut self.queue.split().1, storvsp_protocol::Operation::ENUMERATE_BUS, storvsp_protocol::NtStatus::SUCCESS, &())?;
935                                }
936                            }
937                        }
938                    };
939                }
940                ProtocolState::Init(state) => {
941                    let (mut reader, mut writer) = self.queue.split();
942
943                    // Ensure that subsequent calls to `send_completion` won't
944                    // fail due to lack of ring space, to avoid keeping (and saving/restoring) interim states.
945                    poll_fn(|cx| self.inner.poll_for_ring_space(cx, &mut writer)).await?;
946
947                    tracing::debug!(?state, "process_primary");
948                    match state {
949                        InitState::Begin => {
950                            let packet = self.inner.next_packet(&mut reader).await?;
951                            if let PacketData::BeginInitialization = packet.data {
952                                self.inner.send_completion(
953                                    &mut writer,
954                                    &packet,
955                                    storvsp_protocol::NtStatus::SUCCESS,
956                                    &(),
957                                )?;
958                                *self.inner.protocol.state.write() =
959                                    ProtocolState::Init(InitState::QueryVersion);
960                            } else {
961                                tracelimit::warn_ratelimited!(?state, data = ?packet.data, "unexpected packet order");
962                                self.inner.send_completion(
963                                    &mut writer,
964                                    &packet,
965                                    storvsp_protocol::NtStatus::INVALID_DEVICE_STATE,
966                                    &(),
967                                )?;
968                            }
969                        }
970                        InitState::QueryVersion => {
971                            let packet = self.inner.next_packet(&mut reader).await?;
972                            if let PacketData::QueryProtocolVersion(major_minor) = packet.data {
973                                if let Ok(version) = Version::parse(major_minor) {
974                                    self.inner.send_completion(
975                                        &mut writer,
976                                        &packet,
977                                        storvsp_protocol::NtStatus::SUCCESS,
978                                        &storvsp_protocol::ProtocolVersion {
979                                            major_minor,
980                                            reserved: 0,
981                                        },
982                                    )?;
983                                    self.inner.request_size = version.max_request_size();
984                                    *self.inner.protocol.state.write() =
985                                        ProtocolState::Init(InitState::QueryProperties { version });
986
987                                    tracelimit::info_ratelimited!(
988                                        ?version,
989                                        "scsi version negotiated"
990                                    );
991                                } else {
992                                    self.inner.send_completion(
993                                        &mut writer,
994                                        &packet,
995                                        storvsp_protocol::NtStatus::REVISION_MISMATCH,
996                                        &storvsp_protocol::ProtocolVersion {
997                                            major_minor,
998                                            reserved: 0,
999                                        },
1000                                    )?;
1001                                    *self.inner.protocol.state.write() =
1002                                        ProtocolState::Init(InitState::QueryVersion);
1003                                }
1004                            } else {
1005                                tracelimit::warn_ratelimited!(?state, data = ?packet.data, "unexpected packet order");
1006                                self.inner.send_completion(
1007                                    &mut writer,
1008                                    &packet,
1009                                    storvsp_protocol::NtStatus::INVALID_DEVICE_STATE,
1010                                    &(),
1011                                )?;
1012                            }
1013                        }
1014                        InitState::QueryProperties { version } => {
1015                            let packet = self.inner.next_packet(&mut reader).await?;
1016                            if let PacketData::QueryProperties = packet.data {
1017                                let multi_channel_supported = version >= Version::Win8;
1018
1019                                self.inner.send_completion(
1020                                    &mut writer,
1021                                    &packet,
1022                                    storvsp_protocol::NtStatus::SUCCESS,
1023                                    &storvsp_protocol::ChannelProperties {
1024                                        max_transfer_bytes: 0x40000, // 256KB
1025                                        flags: {
1026                                            if multi_channel_supported {
1027                                                storvsp_protocol::STORAGE_CHANNEL_SUPPORTS_MULTI_CHANNEL
1028                                            } else {
1029                                                0
1030                                            }
1031                                        },
1032                                        maximum_sub_channel_count: if multi_channel_supported {
1033                                            self.inner.channel_control.max_subchannels()
1034                                        } else {
1035                                            0
1036                                        },
1037                                        reserved: 0,
1038                                        reserved2: 0,
1039                                        reserved3: [0, 0],
1040                                    },
1041                                )?;
1042                                *self.inner.protocol.state.write() =
1043                                    ProtocolState::Init(InitState::EndInitialization {
1044                                        version,
1045                                        subchannel_count: if multi_channel_supported {
1046                                            None
1047                                        } else {
1048                                            Some(0)
1049                                        },
1050                                    });
1051                            } else {
1052                                tracelimit::warn_ratelimited!(?state, data = ?packet.data, "unexpected packet order");
1053                                self.inner.send_completion(
1054                                    &mut writer,
1055                                    &packet,
1056                                    storvsp_protocol::NtStatus::INVALID_DEVICE_STATE,
1057                                    &(),
1058                                )?;
1059                            }
1060                        }
1061                        InitState::EndInitialization {
1062                            version,
1063                            subchannel_count,
1064                        } => {
1065                            let packet = self.inner.next_packet(&mut reader).await?;
1066                            match packet.data {
1067                                PacketData::CreateSubChannels(sub_channel_count)
1068                                    if subchannel_count.is_none() =>
1069                                {
1070                                    if let Err(err) = self
1071                                        .inner
1072                                        .channel_control
1073                                        .enable_subchannels(sub_channel_count)
1074                                    {
1075                                        tracelimit::warn_ratelimited!(
1076                                            ?err,
1077                                            "cannot enable subchannels"
1078                                        );
1079                                        self.inner.send_completion(
1080                                            &mut writer,
1081                                            &packet,
1082                                            storvsp_protocol::NtStatus::INVALID_PARAMETER,
1083                                            &(),
1084                                        )?;
1085                                    } else {
1086                                        self.inner.send_completion(
1087                                            &mut writer,
1088                                            &packet,
1089                                            storvsp_protocol::NtStatus::SUCCESS,
1090                                            &(),
1091                                        )?;
1092                                        *self.inner.protocol.state.write() =
1093                                            ProtocolState::Init(InitState::EndInitialization {
1094                                                version,
1095                                                subchannel_count: Some(sub_channel_count),
1096                                            });
1097                                    }
1098                                }
1099                                PacketData::EndInitialization => {
1100                                    self.inner.send_completion(
1101                                        &mut writer,
1102                                        &packet,
1103                                        storvsp_protocol::NtStatus::SUCCESS,
1104                                        &(),
1105                                    )?;
1106                                    // Reset the rescan notification event now, before the guest has a
1107                                    // chance to send any SCSI requests to scan the bus.
1108                                    self.rescan_notification.try_next().ok();
1109                                    *self.inner.protocol.state.write() = ProtocolState::Ready {
1110                                        version,
1111                                        subchannel_count: subchannel_count.unwrap_or(0),
1112                                    };
1113                                    // Wake up subchannels waiting for the
1114                                    // protocol state to become ready.
1115                                    self.inner.protocol.ready.notify(usize::MAX);
1116                                }
1117                                _ => {
1118                                    tracelimit::warn_ratelimited!(?state, data = ?packet.data, "unexpected packet order");
1119                                    self.inner.send_completion(
1120                                        &mut writer,
1121                                        &packet,
1122                                        storvsp_protocol::NtStatus::INVALID_DEVICE_STATE,
1123                                        &(),
1124                                    )?;
1125                                }
1126                            }
1127                        }
1128                    }
1129                }
1130            }
1131        }
1132    }
1133}
1134
1135fn convert_srb_status_to_nt_status(srb_status: SrbStatus) -> storvsp_protocol::NtStatus {
1136    match srb_status {
1137        SrbStatus::BUSY => storvsp_protocol::NtStatus::DEVICE_BUSY,
1138        SrbStatus::SUCCESS => storvsp_protocol::NtStatus::SUCCESS,
1139        SrbStatus::INVALID_LUN
1140        | SrbStatus::INVALID_TARGET_ID
1141        | SrbStatus::NO_DEVICE
1142        | SrbStatus::NO_HBA => storvsp_protocol::NtStatus::DEVICE_DOES_NOT_EXIST,
1143        SrbStatus::COMMAND_TIMEOUT | SrbStatus::TIMEOUT => storvsp_protocol::NtStatus::IO_TIMEOUT,
1144        SrbStatus::SELECTION_TIMEOUT => storvsp_protocol::NtStatus::DEVICE_NOT_CONNECTED,
1145        SrbStatus::BAD_FUNCTION | SrbStatus::BAD_SRB_BLOCK_LENGTH => {
1146            storvsp_protocol::NtStatus::INVALID_DEVICE_REQUEST
1147        }
1148        SrbStatus::DATA_OVERRUN => storvsp_protocol::NtStatus::BUFFER_OVERFLOW,
1149        SrbStatus::REQUEST_FLUSHED => storvsp_protocol::NtStatus::UNSUCCESSFUL,
1150        SrbStatus::ABORTED => storvsp_protocol::NtStatus::CANCELLED,
1151        _ => storvsp_protocol::NtStatus::IO_DEVICE_ERROR,
1152    }
1153}
1154
1155impl WorkerInner {
1156    /// Processes packets and SCSI completions after protocol negotiation has finished.
1157    async fn process_ready<M: RingMem>(
1158        &mut self,
1159        queue: &mut Queue<M>,
1160        protocol_version: Version,
1161    ) -> Result<(), WorkerError> {
1162        self.request_size = protocol_version.max_request_size();
1163        poll_fn(|cx| self.poll_process_ready(cx, queue)).await
1164    }
1165
1166    /// Processes packets and SCSI completions after protocol negotiation has finished.
1167    fn poll_process_ready<M: RingMem>(
1168        &mut self,
1169        cx: &mut Context<'_>,
1170        queue: &mut Queue<M>,
1171    ) -> Poll<Result<(), WorkerError>> {
1172        self.stats.wakes.increment();
1173
1174        let (mut reader, mut writer) = queue.split();
1175        let mut total_completions = 0;
1176        let mut total_submissions = 0;
1177        let poll_mode_queue_depth = self.controller.poll_mode_queue_depth.load(Relaxed) as usize;
1178
1179        loop {
1180            // Drive IOs forward and collect completions.
1181            'outer: while !self.scsi_requests_states.is_empty() {
1182                {
1183                    let mut batch = writer.batched();
1184                    loop {
1185                        // Ensure there is room for the completion before consuming
1186                        // the IO so that we don't have to track completed IOs whose
1187                        // completions haven't been sent.
1188                        if !batch
1189                            .can_write(MAX_VMBUS_PACKET_SIZE)
1190                            .map_err(WorkerError::Queue)?
1191                        {
1192                            // This batch is full but there may still be more completions.
1193                            break;
1194                        }
1195                        if let Poll::Ready(Some((request_id, result, future))) =
1196                            self.scsi_requests.poll_next_unpin(cx)
1197                        {
1198                            self.future_pool.push(future);
1199                            self.handle_completion(&mut batch, request_id, result)?;
1200                            total_completions += 1;
1201                        } else {
1202                            tracing::trace!("out of completions");
1203                            break 'outer;
1204                        }
1205                    }
1206                }
1207
1208                // Wait for enough space for any completion packets.
1209                if self.poll_for_ring_space(cx, &mut writer).is_pending() {
1210                    tracing::trace!("out of ring space");
1211                    break;
1212                }
1213            }
1214
1215            let mut submissions = 0;
1216            // Process new requests.
1217            'outer: loop {
1218                if self.scsi_requests_states.len() >= self.max_io_queue_depth {
1219                    break;
1220                }
1221                let mut batch = if self.scsi_requests_states.len() < poll_mode_queue_depth {
1222                    if let Poll::Ready(batch) = reader.poll_read_batch(cx) {
1223                        batch.map_err(WorkerError::Queue)?
1224                    } else {
1225                        tracing::trace!("out of incoming packets");
1226                        break;
1227                    }
1228                } else {
1229                    match reader.try_read_batch() {
1230                        Ok(batch) => batch,
1231                        Err(queue::TryReadError::Empty) => {
1232                            tracing::trace!(
1233                                pending_io_count = self.scsi_requests_states.len(),
1234                                "out of incoming packets, keeping interrupts masked"
1235                            );
1236                            break;
1237                        }
1238                        Err(queue::TryReadError::Queue(err)) => Err(WorkerError::Queue(err))?,
1239                    }
1240                };
1241
1242                let mut packets = batch.packets();
1243                loop {
1244                    if self.scsi_requests_states.len() >= self.max_io_queue_depth {
1245                        break 'outer;
1246                    }
1247                    // Wait for enough space for any completion packets that
1248                    // `handle_packet` may need to send, so that it isn't necessary
1249                    // to track pending completions.
1250                    if self.poll_for_ring_space(cx, &mut writer).is_pending() {
1251                        tracing::trace!("out of ring space");
1252                        break 'outer;
1253                    }
1254
1255                    let packet = if let Some(packet) = packets.next() {
1256                        packet.map_err(WorkerError::Queue)?
1257                    } else {
1258                        break;
1259                    };
1260
1261                    if self.handle_packet(&mut writer, &packet)? {
1262                        submissions += 1;
1263                    }
1264                }
1265            }
1266
1267            // Loop around to poll the IOs if any new IOs were submitted.
1268            if submissions == 0 {
1269                // No need to poll again.
1270                break;
1271            }
1272            total_submissions += submissions;
1273        }
1274
1275        if total_submissions != 0 || total_completions != 0 {
1276            self.stats.ios_submitted.add(total_submissions);
1277            self.stats
1278                .per_wake_submissions
1279                .add_sample(total_submissions);
1280            self.stats
1281                .per_wake_completions
1282                .add_sample(total_completions);
1283            self.stats.ios_completed.add(total_completions);
1284        } else {
1285            self.stats.wakes_spurious.increment();
1286        }
1287
1288        Poll::Pending
1289    }
1290
1291    fn handle_completion<M: RingMem>(
1292        &mut self,
1293        writer: &mut queue::WriteBatch<'_, M>,
1294        request_id: usize,
1295        result: ScsiResult,
1296    ) -> Result<(), WorkerError> {
1297        let state = self.scsi_requests_states.remove(request_id);
1298        let request_size = state.request.request_size;
1299
1300        // Push the request into the pool to avoid reallocating later.
1301        assert_eq!(
1302            Arc::strong_count(&state.request) + Arc::weak_count(&state.request),
1303            1
1304        );
1305        self.full_request_pool.push(state.request);
1306
1307        let status = convert_srb_status_to_nt_status(result.srb_status);
1308        let mut payload = [0; 0x14];
1309        if let Some(sense) = result.sense_data {
1310            payload[..size_of_val(&sense)].copy_from_slice(sense.as_bytes());
1311            tracing::trace!(sense_info = ?payload, sense_key = payload[2], asc = payload[12], "execute_scsi");
1312        };
1313        let response = storvsp_protocol::ScsiRequest {
1314            length: size_of::<storvsp_protocol::ScsiRequest>() as u16,
1315            scsi_status: result.scsi_status,
1316            srb_status: SrbStatusAndFlags::new()
1317                .with_status(result.srb_status)
1318                .with_autosense_valid(result.sense_data.is_some()),
1319            data_transfer_length: result.tx as u32,
1320            cdb_length: storvsp_protocol::CDB16GENERIC_LENGTH as u8,
1321            sense_info_ex_length: storvsp_protocol::VMSCSI_SENSE_BUFFER_SIZE as u8,
1322            payload,
1323            ..storvsp_protocol::ScsiRequest::new_zeroed()
1324        };
1325        self.send_vmbus_packet(
1326            writer,
1327            OutgoingPacketType::Completion,
1328            request_size,
1329            state.transaction_id,
1330            storvsp_protocol::Operation::COMPLETE_IO,
1331            status,
1332            response.as_bytes(),
1333        )?;
1334        Ok(())
1335    }
1336
1337    fn handle_packet<M: RingMem>(
1338        &mut self,
1339        writer: &mut queue::WriteHalf<'_, M>,
1340        packet: &IncomingPacket<'_, M>,
1341    ) -> Result<bool, WorkerError> {
1342        let packet =
1343            parse_packet(packet, &mut self.full_request_pool).map_err(WorkerError::PacketError)?;
1344        let submitted_io = match packet.data {
1345            PacketData::ExecuteScsi(request) => {
1346                self.push_scsi_request(packet.transaction_id, request);
1347                true
1348            }
1349            PacketData::ResetAdapter | PacketData::ResetBus | PacketData::ResetLun => {
1350                // These operations have always been no-ops.
1351                self.send_completion(writer, &packet, storvsp_protocol::NtStatus::SUCCESS, &())?;
1352                false
1353            }
1354            PacketData::CreateSubChannels(new_subchannel_count) if self.channel_index == 0 => {
1355                if let Err(err) = self
1356                    .channel_control
1357                    .enable_subchannels(new_subchannel_count)
1358                {
1359                    tracelimit::warn_ratelimited!(?err, "cannot create subchannels");
1360                    self.send_completion(
1361                        writer,
1362                        &packet,
1363                        storvsp_protocol::NtStatus::INVALID_PARAMETER,
1364                        &(),
1365                    )?;
1366                    false
1367                } else {
1368                    // Update the subchannel count in the protocol state for save.
1369                    if let ProtocolState::Ready {
1370                        subchannel_count, ..
1371                    } = &mut *self.protocol.state.write()
1372                    {
1373                        *subchannel_count = new_subchannel_count;
1374                    } else {
1375                        unreachable!()
1376                    }
1377
1378                    self.send_completion(
1379                        writer,
1380                        &packet,
1381                        storvsp_protocol::NtStatus::SUCCESS,
1382                        &(),
1383                    )?;
1384                    false
1385                }
1386            }
1387            _ => {
1388                tracelimit::warn_ratelimited!(data = ?packet.data, "unexpected packet on ready");
1389                self.send_completion(
1390                    writer,
1391                    &packet,
1392                    storvsp_protocol::NtStatus::INVALID_DEVICE_STATE,
1393                    &(),
1394                )?;
1395                false
1396            }
1397        };
1398        Ok(submitted_io)
1399    }
1400
1401    fn push_scsi_request(&mut self, transaction_id: u64, full_request: Arc<ScsiRequestAndRange>) {
1402        let scsi_queue = self.scsi_queue.clone();
1403        let scsi_request_state = ScsiRequestState {
1404            transaction_id,
1405            request: full_request.clone(),
1406        };
1407        let request_id = self.scsi_requests_states.insert(scsi_request_state);
1408        let future = self
1409            .future_pool
1410            .pop()
1411            .unwrap_or_else(|| OversizedBox::new(()));
1412        let future = OversizedBox::refill(future, async move {
1413            scsi_queue.execute_scsi(full_request.as_ref()).await
1414        });
1415        let request = ScsiRequest::new(request_id, oversized_box::coerce!(future));
1416        self.scsi_requests.push(request);
1417    }
1418}
1419
1420impl<T: RingMem> Drop for Worker<T> {
1421    fn drop(&mut self) {
1422        self.inner
1423            .controller
1424            .remove_rescan_notification_source(&self.rescan_notification);
1425    }
1426}
1427
1428#[derive(Debug, Error)]
1429#[error("SCSI path {}:{}:{} is already in use", self.0.path, self.0.target, self.0.lun)]
1430pub struct ScsiPathInUse(pub ScsiPath);
1431
1432#[derive(Debug, Error)]
1433#[error("SCSI path {}:{}:{} is not in use", self.0.path, self.0.target, self.0.lun)]
1434pub struct ScsiPathNotInUse(ScsiPath);
1435
1436#[derive(Clone)]
1437struct ScsiRequestState {
1438    transaction_id: u64,
1439    request: Arc<ScsiRequestAndRange>,
1440}
1441
1442#[derive(Debug)]
1443struct ScsiRequestAndRange {
1444    external_data: Range,
1445    external_data_buf: MultiPagedRangeBuf,
1446    request: storvsp_protocol::ScsiRequest,
1447    request_size: usize,
1448}
1449
1450impl Inspect for ScsiRequestState {
1451    fn inspect(&self, req: inspect::Request<'_>) {
1452        req.respond()
1453            .field("transaction_id", self.transaction_id)
1454            .display(
1455                "address",
1456                &ScsiPath {
1457                    path: self.request.request.path_id,
1458                    target: self.request.request.target_id,
1459                    lun: self.request.request.lun,
1460                },
1461            )
1462            .display_debug("operation", &ScsiOp(self.request.request.payload[0]));
1463    }
1464}
1465
1466impl StorageDevice {
1467    /// Returns a new SCSI device.
1468    pub fn build_scsi(
1469        driver_source: &VmTaskDriverSource,
1470        controller: &ScsiController,
1471        instance_id: Guid,
1472        max_sub_channel_count: u16,
1473        io_queue_depth: u32,
1474    ) -> Self {
1475        Self::build_inner(
1476            driver_source,
1477            controller,
1478            instance_id,
1479            None,
1480            max_sub_channel_count,
1481            io_queue_depth,
1482        )
1483    }
1484
1485    /// Returns a new SCSI device for implementing an IDE accelerator channel
1486    /// for IDE device `device_id` on channel `channel_id`.
1487    pub fn build_ide(
1488        driver_source: &VmTaskDriverSource,
1489        channel_id: u8,
1490        device_id: u8,
1491        disk: ScsiControllerDisk,
1492        io_queue_depth: u32,
1493    ) -> Self {
1494        let path = ScsiPath {
1495            path: channel_id,
1496            target: device_id,
1497            lun: 0,
1498        };
1499
1500        let controller = ScsiController::new();
1501        controller.attach(path, disk).unwrap();
1502
1503        // Construct the specific GUID that drivers in the guest expect for this
1504        // IDE device.
1505        let instance_id = Guid {
1506            data1: channel_id.into(),
1507            data2: device_id.into(),
1508            data3: 0x8899,
1509            data4: [0; 8],
1510        };
1511        Self::build_inner(
1512            driver_source,
1513            &controller,
1514            instance_id,
1515            Some(path),
1516            0,
1517            io_queue_depth,
1518        )
1519    }
1520
1521    fn build_inner(
1522        driver_source: &VmTaskDriverSource,
1523        controller: &ScsiController,
1524        instance_id: Guid,
1525        ide_path: Option<ScsiPath>,
1526        max_sub_channel_count: u16,
1527        io_queue_depth: u32,
1528    ) -> Self {
1529        let workers = (0..max_sub_channel_count + 1)
1530            .map(|channel_index| WorkerAndDriver {
1531                worker: TaskControl::new(WorkerState),
1532                driver: driver_source
1533                    .builder()
1534                    .target_vp(0)
1535                    .run_on_target(true)
1536                    .build(format!("storvsp-{}-{}", instance_id, channel_index)),
1537            })
1538            .collect();
1539
1540        Self {
1541            instance_id,
1542            ide_path,
1543            workers,
1544            controller: controller.state.clone(),
1545            resources: Default::default(),
1546            max_sub_channel_count,
1547            driver_source: driver_source.clone(),
1548            protocol: Arc::new(Protocol {
1549                state: RwLock::new(ProtocolState::Init(InitState::Begin)),
1550                ready: Default::default(),
1551            }),
1552            io_queue_depth,
1553        }
1554    }
1555
1556    fn new_worker(
1557        &mut self,
1558        open_request: &OpenRequest,
1559        channel_index: u16,
1560    ) -> anyhow::Result<&mut Worker> {
1561        let controller = self.controller.clone();
1562
1563        // VMBus doesn't provide a target VP if the channel is not using interrupts. Run on VP 0 in
1564        // that case.
1565        let target_vp = open_request.open_data.target_vp.unwrap_or_default();
1566        let driver = self
1567            .driver_source
1568            .builder()
1569            .target_vp(target_vp)
1570            .run_on_target(true)
1571            .build(format!("storvsp-{}-{}", self.instance_id, channel_index));
1572
1573        let channel = gpadl_channel(&driver, &self.resources, open_request, channel_index)
1574            .context("failed to create vmbus channel")?;
1575
1576        let channel_control = self.resources.channel_control.clone();
1577
1578        tracing::debug!(
1579            target_vp = open_request.open_data.target_vp,
1580            channel_index,
1581            "packet processing starting...",
1582        );
1583
1584        // Force the path ID on incoming SCSI requests to match the IDE
1585        // channel ID, since guests do not reliably set the path ID
1586        // correctly.
1587        let force_path_id = self.ide_path.map(|p| p.path);
1588
1589        let worker = Worker::new(
1590            controller,
1591            channel,
1592            channel_index,
1593            self.resources
1594                .offer_resources
1595                .guest_memory(open_request)
1596                .clone(),
1597            channel_control,
1598            self.io_queue_depth,
1599            self.protocol.clone(),
1600            force_path_id,
1601        )
1602        .map_err(RestoreError::Other)?;
1603
1604        self.workers[channel_index as usize]
1605            .driver
1606            .retarget_vp(target_vp);
1607
1608        Ok(self.workers[channel_index as usize].worker.insert(
1609            &driver,
1610            format!("storvsp worker {}-{}", self.instance_id, channel_index),
1611            worker,
1612        ))
1613    }
1614}
1615
1616/// A disk that can be added to a SCSI controller.
1617#[derive(Clone)]
1618pub struct ScsiControllerDisk {
1619    disk: Arc<dyn AsyncScsiDisk>,
1620}
1621
1622impl ScsiControllerDisk {
1623    /// Creates a new controller disk from an async SCSI disk.
1624    pub fn new(disk: Arc<dyn AsyncScsiDisk>) -> Self {
1625        Self { disk }
1626    }
1627}
1628
1629struct ScsiControllerState {
1630    disks: RwLock<HashMap<ScsiPath, ScsiControllerDisk>>,
1631    rescan_notification_source: Mutex<Vec<futures::channel::mpsc::Sender<()>>>,
1632    poll_mode_queue_depth: AtomicU32,
1633}
1634
1635pub struct ScsiController {
1636    state: Arc<ScsiControllerState>,
1637}
1638
1639impl ScsiController {
1640    pub fn new() -> Self {
1641        Self::new_with_poll_mode_queue_depth(None)
1642    }
1643
1644    pub fn new_with_poll_mode_queue_depth(poll_mode_queue_depth: Option<u32>) -> Self {
1645        Self {
1646            state: Arc::new(ScsiControllerState {
1647                disks: Default::default(),
1648                rescan_notification_source: Mutex::new(Vec::new()),
1649                poll_mode_queue_depth: AtomicU32::new(
1650                    poll_mode_queue_depth.unwrap_or(DEFAULT_POLL_MODE_QUEUE_DEPTH),
1651                ),
1652            }),
1653        }
1654    }
1655
1656    pub fn attach(&self, path: ScsiPath, disk: ScsiControllerDisk) -> Result<(), ScsiPathInUse> {
1657        match self.state.disks.write().entry(path) {
1658            Entry::Occupied(_) => return Err(ScsiPathInUse(path)),
1659            Entry::Vacant(entry) => entry.insert(disk),
1660        };
1661        for source in self.state.rescan_notification_source.lock().iter_mut() {
1662            // Ok to ignore errors here. If the channel is full a previous notification has not yet
1663            // been processed by the primary channel worker.
1664            source.try_send(()).ok();
1665        }
1666        Ok(())
1667    }
1668
1669    pub fn remove(&self, path: ScsiPath) -> Result<(), ScsiPathNotInUse> {
1670        match self.state.disks.write().entry(path) {
1671            Entry::Vacant(_) => return Err(ScsiPathNotInUse(path)),
1672            Entry::Occupied(entry) => {
1673                entry.remove();
1674            }
1675        }
1676        for source in self.state.rescan_notification_source.lock().iter_mut() {
1677            // Ok to ignore errors here. If the channel is full a previous notification has not yet
1678            // been processed by the primary channel worker.
1679            source.try_send(()).ok();
1680        }
1681        Ok(())
1682    }
1683}
1684
1685impl ScsiControllerState {
1686    fn add_rescan_notification_source(&self, source: futures::channel::mpsc::Sender<()>) {
1687        self.rescan_notification_source.lock().push(source);
1688    }
1689
1690    fn remove_rescan_notification_source(&self, target: &futures::channel::mpsc::Receiver<()>) {
1691        let mut sources = self.rescan_notification_source.lock();
1692        if let Some(index) = sources
1693            .iter()
1694            .position(|source| source.is_connected_to(target))
1695        {
1696            sources.remove(index);
1697        }
1698    }
1699}
1700
1701#[async_trait]
1702impl VmbusDevice for StorageDevice {
1703    fn offer(&self) -> OfferParams {
1704        if let Some(path) = self.ide_path {
1705            let offer_properties = storvsp_protocol::OfferProperties {
1706                path_id: path.path,
1707                target_id: path.target,
1708                flags: storvsp_protocol::OFFER_PROPERTIES_FLAG_IDE_DEVICE,
1709                ..FromZeros::new_zeroed()
1710            };
1711            let mut user_defined = UserDefinedData::new_zeroed();
1712            offer_properties
1713                .write_to_prefix(&mut user_defined[..])
1714                .unwrap();
1715            OfferParams {
1716                interface_name: "ide-accel".to_owned(),
1717                instance_id: self.instance_id,
1718                interface_id: storvsp_protocol::IDE_ACCELERATOR_INTERFACE_ID,
1719                channel_type: ChannelType::Interface { user_defined },
1720                ..Default::default()
1721            }
1722        } else {
1723            OfferParams {
1724                interface_name: "scsi".to_owned(),
1725                instance_id: self.instance_id,
1726                interface_id: storvsp_protocol::SCSI_INTERFACE_ID,
1727                ..Default::default()
1728            }
1729        }
1730    }
1731
1732    fn max_subchannels(&self) -> u16 {
1733        self.max_sub_channel_count
1734    }
1735
1736    fn install(&mut self, resources: DeviceResources) {
1737        self.resources = resources;
1738    }
1739
1740    async fn open(
1741        &mut self,
1742        channel_index: u16,
1743        open_request: &OpenRequest,
1744    ) -> Result<(), ChannelOpenError> {
1745        tracing::debug!(channel_index, "scsi open channel");
1746        self.new_worker(open_request, channel_index)?;
1747        self.workers[channel_index as usize].worker.start();
1748        Ok(())
1749    }
1750
1751    async fn close(&mut self, channel_index: u16) {
1752        tracing::debug!(channel_index, "scsi close channel");
1753        let worker = &mut self.workers[channel_index as usize].worker;
1754        worker.stop().await;
1755        if worker.state_mut().is_some() {
1756            worker
1757                .state_mut()
1758                .unwrap()
1759                .wait_for_scsi_requests_complete()
1760                .await;
1761            worker.remove();
1762        }
1763        if channel_index == 0 {
1764            *self.protocol.state.write() = ProtocolState::Init(InitState::Begin);
1765        }
1766    }
1767
1768    async fn retarget_vp(&mut self, channel_index: u16, target_vp: u32) {
1769        self.workers[channel_index as usize]
1770            .driver
1771            .retarget_vp(target_vp);
1772    }
1773
1774    fn start(&mut self) {
1775        for task in self
1776            .workers
1777            .iter_mut()
1778            .filter(|task| task.worker.has_state() && !task.worker.is_running())
1779        {
1780            task.worker.start();
1781        }
1782    }
1783
1784    async fn stop(&mut self) {
1785        tracing::debug!(instance_id = ?self.instance_id, "StorageDevice stopping...");
1786        for task in self
1787            .workers
1788            .iter_mut()
1789            .filter(|task| task.worker.has_state() && task.worker.is_running())
1790        {
1791            task.worker.stop().await;
1792        }
1793    }
1794
1795    fn supports_save_restore(&mut self) -> Option<&mut dyn SaveRestoreVmbusDevice> {
1796        Some(self)
1797    }
1798}
1799
1800#[async_trait]
1801impl SaveRestoreVmbusDevice for StorageDevice {
1802    async fn save(&mut self) -> Result<SavedStateBlob, SaveError> {
1803        Ok(SavedStateBlob::new(self.save()?))
1804    }
1805
1806    async fn restore(
1807        &mut self,
1808        control: RestoreControl<'_>,
1809        state: SavedStateBlob,
1810    ) -> Result<(), RestoreError> {
1811        self.restore(control, state.parse()?).await
1812    }
1813}
1814
1815#[cfg(test)]
1816mod tests {
1817    use super::*;
1818    use crate::test_helpers::TestWorker;
1819    use crate::test_helpers::parse_guest_completion;
1820    use crate::test_helpers::parse_guest_completion_check_flags_status;
1821    use pal_async::DefaultDriver;
1822    use pal_async::async_test;
1823    use scsi::srb::SrbStatus;
1824    use test_with_tracing::test;
1825    use vmbus_channel::connected_async_channels;
1826
1827    // Discourage `Clone` for `ScsiController` outside the crate, but it is
1828    // necessary for testing. The fuzzer also uses `TestWorker`, which needs
1829    // a `clone` of the inner state, but is not in this crate.
1830    impl Clone for ScsiController {
1831        fn clone(&self) -> Self {
1832            ScsiController {
1833                state: self.state.clone(),
1834            }
1835        }
1836    }
1837
1838    #[async_test]
1839    async fn test_channel_working(driver: DefaultDriver) {
1840        // set up the channels and worker
1841        let (host, guest) = connected_async_channels(16 * 1024);
1842        let guest_queue = Queue::new(guest).unwrap();
1843
1844        let test_guest_mem = GuestMemory::allocate(16384);
1845        let controller = ScsiController::new();
1846        let disk = scsidisk::SimpleScsiDisk::new(
1847            disklayer_ram::ram_disk(10 * 1024 * 1024, false).unwrap(),
1848            Default::default(),
1849        );
1850        controller
1851            .attach(
1852                ScsiPath {
1853                    path: 0,
1854                    target: 0,
1855                    lun: 0,
1856                },
1857                ScsiControllerDisk::new(Arc::new(disk)),
1858            )
1859            .unwrap();
1860
1861        let test_worker = TestWorker::start(
1862            controller.clone(),
1863            driver.clone(),
1864            test_guest_mem.clone(),
1865            host,
1866            None,
1867        );
1868
1869        let mut guest = test_helpers::TestGuest {
1870            queue: guest_queue,
1871            transaction_id: 0,
1872        };
1873
1874        guest.perform_protocol_negotiation().await;
1875
1876        // Set up the buffer for a write request
1877        const IO_LEN: usize = 4 * 1024;
1878        let write_buf = [7u8; IO_LEN];
1879        let write_gpa = 4 * 1024u64;
1880        test_guest_mem.write_at(write_gpa, &write_buf).unwrap();
1881        guest
1882            .send_write_packet(ScsiPath::default(), write_gpa, 1, IO_LEN)
1883            .await;
1884
1885        guest
1886            .verify_completion(|p| test_helpers::parse_guest_completed_io(p, SrbStatus::SUCCESS))
1887            .await;
1888
1889        let read_gpa = 8 * 1024u64;
1890        guest
1891            .send_read_packet(ScsiPath::default(), read_gpa, 1, IO_LEN)
1892            .await;
1893
1894        guest
1895            .verify_completion(|p| test_helpers::parse_guest_completed_io(p, SrbStatus::SUCCESS))
1896            .await;
1897        let mut read_buf = [0u8; IO_LEN];
1898        test_guest_mem.read_at(read_gpa, &mut read_buf).unwrap();
1899        for (b1, b2) in read_buf.iter().zip(write_buf.iter()) {
1900            assert_eq!(b1, b2);
1901        }
1902
1903        // stop everything
1904        guest.verify_graceful_close(test_worker).await;
1905    }
1906
1907    #[async_test]
1908    async fn test_packet_sizes(driver: DefaultDriver) {
1909        // set up the channels and worker
1910        let (host, guest) = connected_async_channels(16384);
1911        let guest_queue = Queue::new(guest).unwrap();
1912
1913        let test_guest_mem = GuestMemory::allocate(1024);
1914        let controller = ScsiController::new();
1915
1916        let _worker = TestWorker::start(
1917            controller.clone(),
1918            driver.clone(),
1919            test_guest_mem,
1920            host,
1921            None,
1922        );
1923
1924        let mut guest = test_helpers::TestGuest {
1925            queue: guest_queue,
1926            transaction_id: 0,
1927        };
1928
1929        let negotiate_packet = storvsp_protocol::Packet {
1930            operation: storvsp_protocol::Operation::BEGIN_INITIALIZATION,
1931            flags: 0,
1932            status: storvsp_protocol::NtStatus::SUCCESS,
1933        };
1934        guest
1935            .send_data_packet_sync(&[negotiate_packet.as_bytes()])
1936            .await;
1937
1938        guest.verify_completion(parse_guest_completion).await;
1939
1940        let header = storvsp_protocol::Packet {
1941            operation: storvsp_protocol::Operation::QUERY_PROTOCOL_VERSION,
1942            flags: 0,
1943            status: storvsp_protocol::NtStatus::SUCCESS,
1944        };
1945
1946        let mut buf = [0u8; 128];
1947        storvsp_protocol::ProtocolVersion {
1948            major_minor: !0,
1949            reserved: 0,
1950        }
1951        .write_to_prefix(&mut buf[..])
1952        .unwrap(); // PANIC: Infallable since `ProtcolVersion` is less than 128 bytes
1953
1954        for &(len, resp_len) in &[(48, 48), (50, 56), (56, 56), (64, 64), (72, 64)] {
1955            guest
1956                .send_data_packet_sync(&[header.as_bytes(), &buf[..len - size_of_val(&header)]])
1957                .await;
1958
1959            guest
1960                .verify_completion(|packet| {
1961                    let IncomingPacket::Completion(packet) = packet else {
1962                        unreachable!()
1963                    };
1964                    assert_eq!(packet.reader().len(), resp_len);
1965                    assert_eq!(
1966                        packet
1967                            .reader()
1968                            .read_plain::<storvsp_protocol::Packet>()
1969                            .unwrap()
1970                            .status,
1971                        storvsp_protocol::NtStatus::REVISION_MISMATCH
1972                    );
1973                    Ok(())
1974                })
1975                .await;
1976        }
1977    }
1978
1979    #[async_test]
1980    async fn test_wrong_first_packet(driver: DefaultDriver) {
1981        // set up the channels and worker
1982        let (host, guest) = connected_async_channels(16384);
1983        let guest_queue = Queue::new(guest).unwrap();
1984
1985        let test_guest_mem = GuestMemory::allocate(1024);
1986        let controller = ScsiController::new();
1987
1988        let _worker = TestWorker::start(
1989            controller.clone(),
1990            driver.clone(),
1991            test_guest_mem,
1992            host,
1993            None,
1994        );
1995
1996        let mut guest = test_helpers::TestGuest {
1997            queue: guest_queue,
1998            transaction_id: 0,
1999        };
2000
2001        // Protocol negotiation done out of order
2002        let negotiate_packet = storvsp_protocol::Packet {
2003            operation: storvsp_protocol::Operation::END_INITIALIZATION,
2004            flags: 0,
2005            status: storvsp_protocol::NtStatus::SUCCESS,
2006        };
2007        guest
2008            .send_data_packet_sync(&[negotiate_packet.as_bytes()])
2009            .await;
2010
2011        guest
2012            .verify_completion(|packet| {
2013                let IncomingPacket::Completion(packet) = packet else {
2014                    unreachable!()
2015                };
2016                assert_eq!(
2017                    packet
2018                        .reader()
2019                        .read_plain::<storvsp_protocol::Packet>()
2020                        .unwrap()
2021                        .status,
2022                    storvsp_protocol::NtStatus::INVALID_DEVICE_STATE
2023                );
2024                Ok(())
2025            })
2026            .await;
2027    }
2028
2029    #[async_test]
2030    async fn test_unrecognized_operation(driver: DefaultDriver) {
2031        // set up the channels and worker
2032        let (host, guest) = connected_async_channels(16384);
2033        let guest_queue = Queue::new(guest).unwrap();
2034
2035        let test_guest_mem = GuestMemory::allocate(1024);
2036        let controller = ScsiController::new();
2037
2038        let worker = TestWorker::start(
2039            controller.clone(),
2040            driver.clone(),
2041            test_guest_mem,
2042            host,
2043            None,
2044        );
2045
2046        let mut guest = test_helpers::TestGuest {
2047            queue: guest_queue,
2048            transaction_id: 0,
2049        };
2050
2051        // Send packet with unrecognized operation
2052        let negotiate_packet = storvsp_protocol::Packet {
2053            operation: storvsp_protocol::Operation::REMOVE_DEVICE,
2054            flags: 0,
2055            status: storvsp_protocol::NtStatus::SUCCESS,
2056        };
2057        guest
2058            .send_data_packet_sync(&[negotiate_packet.as_bytes()])
2059            .await;
2060
2061        match worker.teardown().await {
2062            Err(WorkerError::PacketError(PacketError::UnrecognizedOperation(
2063                storvsp_protocol::Operation::REMOVE_DEVICE,
2064            ))) => {}
2065            result => panic!("Worker failed with unexpected result {:?}!", result),
2066        }
2067    }
2068
2069    #[async_test]
2070    async fn test_too_many_subchannels(driver: DefaultDriver) {
2071        // set up the channels and worker
2072        let (host, guest) = connected_async_channels(16384);
2073        let guest_queue = Queue::new(guest).unwrap();
2074
2075        let test_guest_mem = GuestMemory::allocate(1024);
2076        let controller = ScsiController::new();
2077
2078        let _worker = TestWorker::start(
2079            controller.clone(),
2080            driver.clone(),
2081            test_guest_mem,
2082            host,
2083            None,
2084        );
2085
2086        let mut guest = test_helpers::TestGuest {
2087            queue: guest_queue,
2088            transaction_id: 0,
2089        };
2090
2091        let negotiate_packet = storvsp_protocol::Packet {
2092            operation: storvsp_protocol::Operation::BEGIN_INITIALIZATION,
2093            flags: 0,
2094            status: storvsp_protocol::NtStatus::SUCCESS,
2095        };
2096        guest
2097            .send_data_packet_sync(&[negotiate_packet.as_bytes()])
2098            .await;
2099        guest.verify_completion(parse_guest_completion).await;
2100
2101        let version_packet = storvsp_protocol::Packet {
2102            operation: storvsp_protocol::Operation::QUERY_PROTOCOL_VERSION,
2103            flags: 0,
2104            status: storvsp_protocol::NtStatus::SUCCESS,
2105        };
2106        let version = storvsp_protocol::ProtocolVersion {
2107            major_minor: storvsp_protocol::VERSION_BLUE,
2108            reserved: 0,
2109        };
2110        guest
2111            .send_data_packet_sync(&[version_packet.as_bytes(), version.as_bytes()])
2112            .await;
2113        guest.verify_completion(parse_guest_completion).await;
2114
2115        let properties_packet = storvsp_protocol::Packet {
2116            operation: storvsp_protocol::Operation::QUERY_PROPERTIES,
2117            flags: 0,
2118            status: storvsp_protocol::NtStatus::SUCCESS,
2119        };
2120        guest
2121            .send_data_packet_sync(&[properties_packet.as_bytes()])
2122            .await;
2123
2124        guest.verify_completion(parse_guest_completion).await;
2125
2126        let negotiate_packet = storvsp_protocol::Packet {
2127            operation: storvsp_protocol::Operation::CREATE_SUB_CHANNELS,
2128            flags: 0,
2129            status: storvsp_protocol::NtStatus::SUCCESS,
2130        };
2131        // Create sub channels more than maximum_sub_channel_count
2132        guest
2133            .send_data_packet_sync(&[negotiate_packet.as_bytes(), 1_u16.as_bytes()])
2134            .await;
2135
2136        guest
2137            .verify_completion(|packet| {
2138                let IncomingPacket::Completion(packet) = packet else {
2139                    unreachable!()
2140                };
2141                assert_eq!(
2142                    packet
2143                        .reader()
2144                        .read_plain::<storvsp_protocol::Packet>()
2145                        .unwrap()
2146                        .status,
2147                    storvsp_protocol::NtStatus::INVALID_PARAMETER
2148                );
2149                Ok(())
2150            })
2151            .await;
2152    }
2153
2154    #[async_test]
2155    async fn test_begin_init_on_ready(driver: DefaultDriver) {
2156        // set up the channels and worker
2157        let (host, guest) = connected_async_channels(16384);
2158        let guest_queue = Queue::new(guest).unwrap();
2159
2160        let test_guest_mem = GuestMemory::allocate(1024);
2161        let controller = ScsiController::new();
2162
2163        let _worker = TestWorker::start(
2164            controller.clone(),
2165            driver.clone(),
2166            test_guest_mem,
2167            host,
2168            None,
2169        );
2170
2171        let mut guest = test_helpers::TestGuest {
2172            queue: guest_queue,
2173            transaction_id: 0,
2174        };
2175
2176        guest.perform_protocol_negotiation().await;
2177
2178        // Protocol negotiation done out of order
2179        let negotiate_packet = storvsp_protocol::Packet {
2180            operation: storvsp_protocol::Operation::BEGIN_INITIALIZATION,
2181            flags: 0,
2182            status: storvsp_protocol::NtStatus::SUCCESS,
2183        };
2184        guest
2185            .send_data_packet_sync(&[negotiate_packet.as_bytes()])
2186            .await;
2187
2188        guest
2189            .verify_completion(|p| {
2190                parse_guest_completion_check_flags_status(
2191                    p,
2192                    0,
2193                    storvsp_protocol::NtStatus::INVALID_DEVICE_STATE,
2194                )
2195            })
2196            .await;
2197    }
2198
2199    #[async_test]
2200    async fn test_hot_add_remove(driver: DefaultDriver) {
2201        // set up channels and worker.
2202        let (host, guest) = connected_async_channels(16 * 1024);
2203        let guest_queue = Queue::new(guest).unwrap();
2204
2205        let test_guest_mem = GuestMemory::allocate(16384);
2206        // create a controller with no disk yet.
2207        let controller = ScsiController::new();
2208
2209        let test_worker = TestWorker::start(
2210            controller.clone(),
2211            driver.clone(),
2212            test_guest_mem.clone(),
2213            host,
2214            None,
2215        );
2216
2217        let mut guest = test_helpers::TestGuest {
2218            queue: guest_queue,
2219            transaction_id: 0,
2220        };
2221
2222        guest.perform_protocol_negotiation().await;
2223
2224        // Verify no LUNs are reported initially.
2225        let mut lun_list_buffer: [u8; 256] = [0; 256];
2226        let mut disk_count = 0;
2227        guest
2228            .send_report_luns_packet(ScsiPath::default(), 0, lun_list_buffer.len())
2229            .await;
2230        guest
2231            .verify_completion(|p| {
2232                test_helpers::parse_guest_completed_io_check_tx_len(p, SrbStatus::SUCCESS, Some(8))
2233            })
2234            .await;
2235        test_guest_mem.read_at(0, &mut lun_list_buffer).unwrap();
2236        let lun_list_size = u32::from_be_bytes(lun_list_buffer[0..4].try_into().unwrap());
2237        assert_eq!(lun_list_size, disk_count as u32 * 8);
2238
2239        // Set up a buffer for writes.
2240        const IO_LEN: usize = 4 * 1024;
2241        let write_buf = [7u8; IO_LEN];
2242        let write_gpa = 4 * 1024u64;
2243        test_guest_mem.write_at(write_gpa, &write_buf).unwrap();
2244
2245        guest
2246            .send_write_packet(ScsiPath::default(), write_gpa, 1, IO_LEN)
2247            .await;
2248        guest
2249            .verify_completion(|p| {
2250                test_helpers::parse_guest_completed_io(p, SrbStatus::INVALID_LUN)
2251            })
2252            .await;
2253
2254        // Add some disks while the guest is running.
2255        for lun in 0..4 {
2256            let disk = scsidisk::SimpleScsiDisk::new(
2257                disklayer_ram::ram_disk(10 * 1024 * 1024, false).unwrap(),
2258                Default::default(),
2259            );
2260            controller
2261                .attach(
2262                    ScsiPath {
2263                        path: 0,
2264                        target: 0,
2265                        lun,
2266                    },
2267                    ScsiControllerDisk::new(Arc::new(disk)),
2268                )
2269                .unwrap();
2270            guest
2271                .verify_completion(test_helpers::parse_guest_enumerate_bus)
2272                .await;
2273
2274            disk_count += 1;
2275            guest
2276                .send_report_luns_packet(ScsiPath::default(), 0, 256)
2277                .await;
2278            guest
2279                .verify_completion(|p| {
2280                    test_helpers::parse_guest_completed_io_check_tx_len(
2281                        p,
2282                        SrbStatus::SUCCESS,
2283                        Some((disk_count + 1) * 8),
2284                    )
2285                })
2286                .await;
2287            test_guest_mem.read_at(0, &mut lun_list_buffer).unwrap();
2288            let lun_list_size = u32::from_be_bytes(lun_list_buffer[0..4].try_into().unwrap());
2289            assert_eq!(lun_list_size, disk_count as u32 * 8);
2290
2291            guest
2292                .send_write_packet(
2293                    ScsiPath {
2294                        path: 0,
2295                        target: 0,
2296                        lun,
2297                    },
2298                    write_gpa,
2299                    1,
2300                    IO_LEN,
2301                )
2302                .await;
2303            guest
2304                .verify_completion(|p| {
2305                    test_helpers::parse_guest_completed_io(p, SrbStatus::SUCCESS)
2306                })
2307                .await;
2308        }
2309
2310        // Remove all disks while the guest is running.
2311        for lun in 0..4 {
2312            controller
2313                .remove(ScsiPath {
2314                    path: 0,
2315                    target: 0,
2316                    lun,
2317                })
2318                .unwrap();
2319            guest
2320                .verify_completion(test_helpers::parse_guest_enumerate_bus)
2321                .await;
2322
2323            disk_count -= 1;
2324            guest
2325                .send_report_luns_packet(ScsiPath::default(), 0, 4096)
2326                .await;
2327            guest
2328                .verify_completion(|p| {
2329                    test_helpers::parse_guest_completed_io_check_tx_len(
2330                        p,
2331                        SrbStatus::SUCCESS,
2332                        Some((disk_count + 1) * 8),
2333                    )
2334                })
2335                .await;
2336            test_guest_mem.read_at(0, &mut lun_list_buffer).unwrap();
2337            let lun_list_size = u32::from_be_bytes(lun_list_buffer[0..4].try_into().unwrap());
2338            assert_eq!(lun_list_size, disk_count as u32 * 8);
2339
2340            guest
2341                .send_write_packet(
2342                    ScsiPath {
2343                        path: 0,
2344                        target: 0,
2345                        lun,
2346                    },
2347                    write_gpa,
2348                    1,
2349                    IO_LEN,
2350                )
2351                .await;
2352            guest
2353                .verify_completion(|p| {
2354                    test_helpers::parse_guest_completed_io(p, SrbStatus::INVALID_LUN)
2355                })
2356                .await;
2357        }
2358
2359        guest.verify_graceful_close(test_worker).await;
2360    }
2361
2362    #[async_test]
2363    async fn test_async_disk(driver: DefaultDriver) {
2364        let device = disklayer_ram::ram_disk(64 * 1024, false).unwrap();
2365        let controller = ScsiController::new();
2366        let disk = ScsiControllerDisk::new(Arc::new(scsidisk::SimpleScsiDisk::new(
2367            device,
2368            Default::default(),
2369        )));
2370        controller
2371            .attach(
2372                ScsiPath {
2373                    path: 0,
2374                    target: 0,
2375                    lun: 0,
2376                },
2377                disk,
2378            )
2379            .unwrap();
2380
2381        let (host, guest) = connected_async_channels(16 * 1024);
2382        let guest_queue = Queue::new(guest).unwrap();
2383
2384        let mut guest = test_helpers::TestGuest {
2385            queue: guest_queue,
2386            transaction_id: 0,
2387        };
2388
2389        let test_guest_mem = GuestMemory::allocate(16384);
2390        let worker = TestWorker::start(
2391            controller.clone(),
2392            &driver,
2393            test_guest_mem.clone(),
2394            host,
2395            None,
2396        );
2397
2398        let negotiate_packet = storvsp_protocol::Packet {
2399            operation: storvsp_protocol::Operation::BEGIN_INITIALIZATION,
2400            flags: 0,
2401            status: storvsp_protocol::NtStatus::SUCCESS,
2402        };
2403        guest
2404            .send_data_packet_sync(&[negotiate_packet.as_bytes()])
2405            .await;
2406        guest.verify_completion(parse_guest_completion).await;
2407
2408        let version_packet = storvsp_protocol::Packet {
2409            operation: storvsp_protocol::Operation::QUERY_PROTOCOL_VERSION,
2410            flags: 0,
2411            status: storvsp_protocol::NtStatus::SUCCESS,
2412        };
2413        let version = storvsp_protocol::ProtocolVersion {
2414            major_minor: storvsp_protocol::VERSION_BLUE,
2415            reserved: 0,
2416        };
2417        guest
2418            .send_data_packet_sync(&[version_packet.as_bytes(), version.as_bytes()])
2419            .await;
2420        guest.verify_completion(parse_guest_completion).await;
2421
2422        let properties_packet = storvsp_protocol::Packet {
2423            operation: storvsp_protocol::Operation::QUERY_PROPERTIES,
2424            flags: 0,
2425            status: storvsp_protocol::NtStatus::SUCCESS,
2426        };
2427        guest
2428            .send_data_packet_sync(&[properties_packet.as_bytes()])
2429            .await;
2430        guest.verify_completion(parse_guest_completion).await;
2431
2432        let negotiate_packet = storvsp_protocol::Packet {
2433            operation: storvsp_protocol::Operation::END_INITIALIZATION,
2434            flags: 0,
2435            status: storvsp_protocol::NtStatus::SUCCESS,
2436        };
2437        guest
2438            .send_data_packet_sync(&[negotiate_packet.as_bytes()])
2439            .await;
2440        guest.verify_completion(parse_guest_completion).await;
2441
2442        const IO_LEN: usize = 4 * 1024;
2443        let write_buf = [7u8; IO_LEN];
2444        let write_gpa = 4 * 1024u64;
2445        test_guest_mem.write_at(write_gpa, &write_buf).unwrap();
2446        guest
2447            .send_write_packet(ScsiPath::default(), write_gpa, 1, IO_LEN)
2448            .await;
2449        guest
2450            .verify_completion(|p| test_helpers::parse_guest_completed_io(p, SrbStatus::SUCCESS))
2451            .await;
2452
2453        let read_gpa = 8 * 1024u64;
2454        guest
2455            .send_read_packet(ScsiPath::default(), read_gpa, 1, IO_LEN)
2456            .await;
2457        guest
2458            .verify_completion(|p| test_helpers::parse_guest_completed_io(p, SrbStatus::SUCCESS))
2459            .await;
2460        let mut read_buf = [0u8; IO_LEN];
2461        test_guest_mem.read_at(read_gpa, &mut read_buf).unwrap();
2462        for (b1, b2) in read_buf.iter().zip(write_buf.iter()) {
2463            assert_eq!(b1, b2);
2464        }
2465
2466        guest.verify_graceful_close(worker).await;
2467    }
2468}