1#![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
142const 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 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
288type ScsiOpStorage = [u64; SCSI_REQUEST_STACK_SIZE / 8];
296type ScsiOpFuture = Pin<OversizedBox<dyn Future<Output = ScsiResult> + Send, ScsiOpStorage>>;
297
298const 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 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 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 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 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 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; 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 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 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)] 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 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 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 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 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 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, 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 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 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 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 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 'outer: while !self.scsi_requests_states.is_empty() {
1182 {
1183 let mut batch = writer.batched();
1184 loop {
1185 if !batch
1189 .can_write(MAX_VMBUS_PACKET_SIZE)
1190 .map_err(WorkerError::Queue)?
1191 {
1192 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 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 '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 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 if submissions == 0 {
1269 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 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 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 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 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 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 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 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 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#[derive(Clone)]
1618pub struct ScsiControllerDisk {
1619 disk: Arc<dyn AsyncScsiDisk>,
1620}
1621
1622impl ScsiControllerDisk {
1623 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 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 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 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 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 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 guest.verify_graceful_close(test_worker).await;
1905 }
1906
1907 #[async_test]
1908 async fn test_packet_sizes(driver: DefaultDriver) {
1909 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(); 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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}